Compare commits
104 changed files with 7331 additions and 14088 deletions
1
.envrc
1
.envrc
|
|
@ -1 +0,0 @@
|
|||
use flake
|
||||
|
|
@ -7,12 +7,16 @@ on:
|
|||
|
||||
env:
|
||||
CARGO_TERM_COLOR: always
|
||||
NIX_CONFIG: experimental-features = nix-command flakes
|
||||
|
||||
jobs:
|
||||
checks:
|
||||
name: checks
|
||||
runs-on: nixos
|
||||
steps:
|
||||
- name: Install node
|
||||
run: nix profile add nixpkgs#nodejs_24
|
||||
|
||||
- name: Checkout
|
||||
uses: https://data.forgejo.org/actions/checkout@v7
|
||||
|
||||
|
|
@ -29,6 +33,7 @@ jobs:
|
|||
cargo machete
|
||||
|
||||
pnpm install --frozen-lockfile
|
||||
pnpm run dup
|
||||
|
||||
RUSTFLAGS="--cfg web_sys_unstable_apis" wasm-pack test --node wasm
|
||||
RUSTFLAGS="--cfg web_sys_unstable_apis" wasm-pack build wasm --target web
|
||||
|
|
|
|||
|
|
@ -14,10 +14,16 @@ on:
|
|||
required: true
|
||||
type: string
|
||||
|
||||
env:
|
||||
NIX_CONFIG: experimental-features = nix-command flakes
|
||||
|
||||
jobs:
|
||||
release:
|
||||
runs-on: nixos
|
||||
steps:
|
||||
- name: Install node & bun
|
||||
run: nix profile add nixpkgs#nodejs_24 nixpkgs#bun
|
||||
|
||||
- name: Check out repo
|
||||
uses: https://data.forgejo.org/actions/checkout@v7
|
||||
with:
|
||||
|
|
@ -26,6 +32,9 @@ jobs:
|
|||
- name: Install dependencies
|
||||
run: bun install
|
||||
|
||||
- name: Install cc linker, sed & jq
|
||||
run: nix profile add nixpkgs#stdenv.cc nixpkgs#gnused nixpkgs#jq
|
||||
|
||||
- name: Build all
|
||||
run: bun build:all
|
||||
|
||||
|
|
|
|||
1
.gitignore
vendored
1
.gitignore
vendored
|
|
@ -6,4 +6,3 @@ dist/
|
|||
*.tgz
|
||||
wasm/pkg/
|
||||
web_client/
|
||||
.direnv
|
||||
|
|
|
|||
163
Cargo.lock
generated
163
Cargo.lock
generated
|
|
@ -90,9 +90,9 @@ dependencies = [
|
|||
|
||||
[[package]]
|
||||
name = "async-trait"
|
||||
version = "0.1.92"
|
||||
version = "0.1.91"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "82f6aeea286b8eb4dd3431a1be1b59d290ace00f5bfd8e2a159bc2a05e2c1667"
|
||||
checksum = "ae36dc4177970ef04fde5178d3e2429882def40e57a451f919c098f72baa6cec"
|
||||
dependencies = [
|
||||
"proc-macro2",
|
||||
"quote",
|
||||
|
|
@ -221,9 +221,9 @@ checksum = "37b2a672a2cb129a2e41c10b1224bb368f9f37a2b16b612598138befd7b37eb5"
|
|||
|
||||
[[package]]
|
||||
name = "cc"
|
||||
version = "1.4.3"
|
||||
version = "1.4.2"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "509591b7bcd67f4ef775afad7662703b4935daaa6ec0e5605cfb1090b32a2b6d"
|
||||
checksum = "5d262e149917187838d5b42777c8253bcb64500067342904e7d429499a6f277e"
|
||||
dependencies = [
|
||||
"find-msvc-tools",
|
||||
"jobserver",
|
||||
|
|
@ -256,9 +256,9 @@ dependencies = [
|
|||
|
||||
[[package]]
|
||||
name = "chacha20"
|
||||
version = "0.10.2"
|
||||
version = "0.10.1"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "65c35e4b699c7e15ccbe7ee35c005e4fc0a278d22238a2857e6ce2dadeda1b06"
|
||||
checksum = "d524456ba66e72eb8b115ff89e01e497f8e6d11d78b70b1aa13c0fbd97540a81"
|
||||
dependencies = [
|
||||
"cfg-if",
|
||||
"cpufeatures 0.3.0",
|
||||
|
|
@ -596,9 +596,9 @@ checksum = "64cd1e32ddd350061ae6edb1b082d7c54915b5c672c389143b9a63403a109f24"
|
|||
|
||||
[[package]]
|
||||
name = "find-msvc-tools"
|
||||
version = "0.1.11"
|
||||
version = "0.1.10"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "d45db016d36b838f563236e9193d0ee6ce38f3f68b6c94e914b4929c96bbb890"
|
||||
checksum = "26b73573e6edcd2af0cdf47bd6cb58f0b3839491263c314eaad1ccf24430e1de"
|
||||
|
||||
[[package]]
|
||||
name = "fnv"
|
||||
|
|
@ -629,9 +629,9 @@ checksum = "42703706b716c37f96a77aea830392ad231f44c9e9a67872fa5548707e11b11c"
|
|||
|
||||
[[package]]
|
||||
name = "futures"
|
||||
version = "0.3.34"
|
||||
version = "0.3.33"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "9a31d2a3fbaaeb2af2368bbdd904aa8e812d3c04a1ee10d3171f52d556e5d0a3"
|
||||
checksum = "a88cf1f829d945f548cf8fec32c61b1f202b6d93b45848602fc02af4b12ad218"
|
||||
dependencies = [
|
||||
"futures-channel",
|
||||
"futures-core",
|
||||
|
|
@ -660,9 +660,9 @@ checksum = "92d699e522242e69e3003b94ecc1f960f3a5e015aa7c5d7486e65ad01dd94f5e"
|
|||
|
||||
[[package]]
|
||||
name = "futures-executor"
|
||||
version = "0.3.34"
|
||||
version = "0.3.33"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "031b47cf1a3c6cc8bc2fc76cd437f521619387907d469316e7c0bc278f1f5432"
|
||||
checksum = "6754879cc9f2c66f88c6e5c35344bb0bdb0708b0352b1201815667c7eabc7458"
|
||||
dependencies = [
|
||||
"futures-core",
|
||||
"futures-task",
|
||||
|
|
@ -764,9 +764,9 @@ dependencies = [
|
|||
|
||||
[[package]]
|
||||
name = "h2"
|
||||
version = "0.4.16"
|
||||
version = "0.4.15"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "a9f37a958b41b3b19ee2707c06439c0e9e547e847223eb791ecb0cb821c65e27"
|
||||
checksum = "6cb093c84e8bd9b188d4c4a8cb6579fc016968d14c99882163cd3ff402a4f155"
|
||||
dependencies = [
|
||||
"atomic-waker",
|
||||
"bytes",
|
||||
|
|
@ -895,9 +895,9 @@ dependencies = [
|
|||
|
||||
[[package]]
|
||||
name = "http-body-util"
|
||||
version = "0.1.5"
|
||||
version = "0.1.4"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "23169fe34a5fbcdd3f3862e78fb9b6fccd5f02a6dc6f732547005d45631ce71c"
|
||||
checksum = "e9f41fd6a08e4d4ec69df65976da761afd5ad5e58a9d4acb46bd1c953a9e3ff2"
|
||||
dependencies = [
|
||||
"bytes",
|
||||
"futures-core",
|
||||
|
|
@ -966,9 +966,9 @@ dependencies = [
|
|||
|
||||
[[package]]
|
||||
name = "icu_collections"
|
||||
version = "2.3.0"
|
||||
version = "2.2.0"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "fa68d21081c4a05d5a901a1c62add574c77048b6a1c67be3b50ce0b60d4ca513"
|
||||
checksum = "2984d1cd16c883d7935b9e07e44071dca8d917fd52ecc02c04d5fa0b5a3f191c"
|
||||
dependencies = [
|
||||
"displaydoc",
|
||||
"potential_utf",
|
||||
|
|
@ -980,9 +980,9 @@ dependencies = [
|
|||
|
||||
[[package]]
|
||||
name = "icu_locale_core"
|
||||
version = "2.3.0"
|
||||
version = "2.2.0"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "d56e28588da92eee5c3201a6eff33fabdd49b62269c8938d4ff050ce4d900deb"
|
||||
checksum = "92219b62b3e2b4d88ac5119f8904c10f8f61bf7e95b640d25ba3075e6cac2c29"
|
||||
dependencies = [
|
||||
"displaydoc",
|
||||
"litemap",
|
||||
|
|
@ -993,9 +993,9 @@ dependencies = [
|
|||
|
||||
[[package]]
|
||||
name = "icu_normalizer"
|
||||
version = "2.3.0"
|
||||
version = "2.2.0"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "12f9cf5f235641ed274641dd81c3f28d870e276763d0797aeeab72317b1c646f"
|
||||
checksum = "c56e5ee99d6e3d33bd91c5d85458b6005a22140021cc324cea84dd0e72cff3b4"
|
||||
dependencies = [
|
||||
"icu_collections",
|
||||
"icu_normalizer_data",
|
||||
|
|
@ -1007,17 +1007,16 @@ dependencies = [
|
|||
|
||||
[[package]]
|
||||
name = "icu_normalizer_data"
|
||||
version = "2.3.0"
|
||||
version = "2.2.0"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "1563da1ed3e0b3bf3d74c9b85917ac9c56464d2f57242270c09c9e752f8021a0"
|
||||
checksum = "da3be0ae77ea334f4da67c12f149704f19f81d1adf7c51cf482943e84a2bad38"
|
||||
|
||||
[[package]]
|
||||
name = "icu_properties"
|
||||
version = "2.3.0"
|
||||
version = "2.2.0"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "7e7ca276ad3145661a65914e6daf131ca5120cd3dcee8f8f3214b8875184a148"
|
||||
checksum = "bee3b67d0ea5c2cca5003417989af8996f8604e34fb9ddf96208a033901e70de"
|
||||
dependencies = [
|
||||
"displaydoc",
|
||||
"icu_collections",
|
||||
"icu_locale_core",
|
||||
"icu_properties_data",
|
||||
|
|
@ -1028,15 +1027,15 @@ dependencies = [
|
|||
|
||||
[[package]]
|
||||
name = "icu_properties_data"
|
||||
version = "2.3.0"
|
||||
version = "2.2.0"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "e590f038c1464a96894fd6d10127e90a8be4509f56ff7ecef851b15cee0b7caa"
|
||||
checksum = "8e2bbb201e0c04f7b4b3e14382af113e17ba4f63e2c9d2ee626b720cbce54a14"
|
||||
|
||||
[[package]]
|
||||
name = "icu_provider"
|
||||
version = "2.3.0"
|
||||
version = "2.2.0"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "92a7ed671a6aad807a8651a2e1782a6598fda9ce5185dd8158549e95a91c6428"
|
||||
checksum = "139c4cf31c8b5f33d7e199446eff9c1e02decfc2f0eec2c8d71f65befa45b421"
|
||||
dependencies = [
|
||||
"displaydoc",
|
||||
"icu_locale_core",
|
||||
|
|
@ -1154,9 +1153,9 @@ dependencies = [
|
|||
|
||||
[[package]]
|
||||
name = "js-sys"
|
||||
version = "0.3.104"
|
||||
version = "0.3.103"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "0e0c1080212aad755ea003d18543e8768dd432c48819efd73a7bf1e39b7a5a3a"
|
||||
checksum = "53b44bfcdb3f8d5837a46dae1ca9660a837176eee74a28b229bc626816589102"
|
||||
dependencies = [
|
||||
"cfg-if",
|
||||
"futures-util",
|
||||
|
|
@ -1174,9 +1173,9 @@ dependencies = [
|
|||
|
||||
[[package]]
|
||||
name = "keccak"
|
||||
version = "0.2.1"
|
||||
version = "0.2.0"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "ffd9697dc4a9a62e2da93389f34400b77a28f0287711263cabb203b3ccb9c0e4"
|
||||
checksum = "9e24a010dd405bd7ed803e5253182815b41bf2e6a80cc3bfc066658e03a198aa"
|
||||
dependencies = [
|
||||
"cfg-if",
|
||||
"cpufeatures 0.3.0",
|
||||
|
|
@ -1202,9 +1201,9 @@ checksum = "b6d2cec3eae94f9f509c767b45932f1ada8350c4bdb85af2fcab4a3c14807981"
|
|||
|
||||
[[package]]
|
||||
name = "litemap"
|
||||
version = "0.8.3"
|
||||
version = "0.8.2"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "47d9d19d1d6efa0109d2f65ff4c85cddd50bd572e5a00127ab10987290bcefae"
|
||||
checksum = "92daf443525c4cce67b150400bc2316076100ce0b3686209eb8cf3c31612e6f0"
|
||||
|
||||
[[package]]
|
||||
name = "lock_api"
|
||||
|
|
@ -1235,9 +1234,9 @@ checksum = "cf8baf1c55e62ffcace7a9f06f4bd9cd3f0c4beb022d3b367256b91b87513d98"
|
|||
|
||||
[[package]]
|
||||
name = "minicov"
|
||||
version = "0.3.9"
|
||||
version = "0.3.8"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "c3aa3aa12b448ac225b3102217d1ac5cc717908f02722926524b0599c933c7a0"
|
||||
checksum = "4869b6a491569605d66d3952bcdf03df789e5b536e5f0cf7758a7f08a55ae24d"
|
||||
dependencies = [
|
||||
"cc",
|
||||
"walkdir",
|
||||
|
|
@ -1326,6 +1325,9 @@ dependencies = [
|
|||
"mtp-transport",
|
||||
"mtp-type-map",
|
||||
"mtp-webserver",
|
||||
"rand",
|
||||
"rcgen",
|
||||
"tokio",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
|
|
@ -1360,6 +1362,7 @@ name = "mtp-common"
|
|||
version = "0.3.0"
|
||||
dependencies = [
|
||||
"quinn",
|
||||
"rustls",
|
||||
"thiserror 2.0.20",
|
||||
"wtransport",
|
||||
]
|
||||
|
|
@ -1369,7 +1372,6 @@ name = "mtp-crypto"
|
|||
version = "0.3.0"
|
||||
dependencies = [
|
||||
"aes-gcm",
|
||||
"argon2",
|
||||
"base64 0.22.1",
|
||||
"chacha20poly1305",
|
||||
"ed25519-dalek",
|
||||
|
|
@ -1393,6 +1395,7 @@ dependencies = [
|
|||
name = "mtp-files"
|
||||
version = "0.3.0"
|
||||
dependencies = [
|
||||
"argon2",
|
||||
"mtp-crypto",
|
||||
"rand",
|
||||
"thiserror 2.0.20",
|
||||
|
|
@ -1408,7 +1411,6 @@ dependencies = [
|
|||
"mtp-crypto",
|
||||
"mtp-transport",
|
||||
"rand",
|
||||
"thiserror 2.0.20",
|
||||
"tokio",
|
||||
"tracing",
|
||||
"wtransport",
|
||||
|
|
@ -1662,9 +1664,9 @@ dependencies = [
|
|||
|
||||
[[package]]
|
||||
name = "pkg-config"
|
||||
version = "0.3.34"
|
||||
version = "0.3.33"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "f6b464fbc74e149a392436b17d523f769e057cb6877f6a5c4618bc6f11800548"
|
||||
checksum = "19f132c84eca552bf34cab8ec81f1c1dcc229b811638f9d283dceabe58c5569e"
|
||||
|
||||
[[package]]
|
||||
name = "poly1305"
|
||||
|
|
@ -1697,9 +1699,9 @@ checksum = "05c8b63e8d9609db387f0324918f81d68fe27748f084ef092fb35954d0539a85"
|
|||
|
||||
[[package]]
|
||||
name = "potential_utf"
|
||||
version = "0.1.6"
|
||||
version = "0.1.5"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "d83eb9bc6d8e5cf568e7a1101d60ee05e81ed50ea106026f3d18deeb046d7661"
|
||||
checksum = "0103b1cef7ec0cf76490e969665504990193874ea05c85ff9bab8b911d0a0564"
|
||||
dependencies = [
|
||||
"zerovec",
|
||||
]
|
||||
|
|
@ -1742,9 +1744,9 @@ dependencies = [
|
|||
|
||||
[[package]]
|
||||
name = "quinn-proto"
|
||||
version = "0.11.17"
|
||||
version = "0.11.16"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "04759210543be93709136e28212294a659ef5001836ff4eab4d663e4529bba83"
|
||||
checksum = "2f4bfc015262b9df63c8845072ce59068853ff5872180c2ce2f13038b970e560"
|
||||
dependencies = [
|
||||
"aws-lc-rs",
|
||||
"bytes",
|
||||
|
|
@ -1800,7 +1802,7 @@ version = "0.10.2"
|
|||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "c7f5fa3a058cd35567ef9bfa5e75732bee0f9e4c55fa90477bef2dfcdbc4be80"
|
||||
dependencies = [
|
||||
"chacha20 0.10.2",
|
||||
"chacha20 0.10.1",
|
||||
"getrandom 0.4.3",
|
||||
"rand_core 0.10.1",
|
||||
]
|
||||
|
|
@ -2117,7 +2119,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index"
|
|||
checksum = "09057cb2149ad4cbd2da1e26b351f9a4c354219421229c69c3063e6f61947c4a"
|
||||
dependencies = [
|
||||
"digest 0.11.3",
|
||||
"keccak 0.2.1",
|
||||
"keccak 0.2.0",
|
||||
"sponge-cursor",
|
||||
]
|
||||
|
||||
|
|
@ -2342,9 +2344,9 @@ dependencies = [
|
|||
|
||||
[[package]]
|
||||
name = "tinystr"
|
||||
version = "0.8.4"
|
||||
version = "0.8.3"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "b1e27c91459209c2986af3dcf603a5a74a4368754ce37414f59acc971167f643"
|
||||
checksum = "c8323304221c2a851516f22236c5722a72eaa19749016521d6dff0824447d96d"
|
||||
dependencies = [
|
||||
"displaydoc",
|
||||
"zerovec",
|
||||
|
|
@ -2405,9 +2407,9 @@ dependencies = [
|
|||
|
||||
[[package]]
|
||||
name = "tokio-stream"
|
||||
version = "0.1.19"
|
||||
version = "0.1.18"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "a3d06f0b082ba57c26b79407372e57cf2a1e28124f78e9479fe80322cf53420b"
|
||||
checksum = "32da49809aab5c3bc678af03902d4ccddea2a87d028d86392a4b1560c6906c70"
|
||||
dependencies = [
|
||||
"futures-core",
|
||||
"pin-project-lite",
|
||||
|
|
@ -2416,14 +2418,13 @@ dependencies = [
|
|||
|
||||
[[package]]
|
||||
name = "tokio-util"
|
||||
version = "0.7.19"
|
||||
version = "0.7.18"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "494815d09bf52b5548659851081238f0ca39ff638363907596da739561c62c52"
|
||||
checksum = "9ae9cec805b01e8fc3fd2fe289f89149a9b66dd16786abd8b19cfa7b48cb0098"
|
||||
dependencies = [
|
||||
"bytes",
|
||||
"futures-core",
|
||||
"futures-sink",
|
||||
"libc",
|
||||
"pin-project-lite",
|
||||
"tokio",
|
||||
]
|
||||
|
|
@ -2568,9 +2569,9 @@ checksum = "ccf3ec651a847eb01de73ccad15eb7d99f80485de043efb2f370cd654f4ea44b"
|
|||
|
||||
[[package]]
|
||||
name = "wasm-bindgen"
|
||||
version = "0.2.127"
|
||||
version = "0.2.126"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "1b70935747edd64d89de3efa29d73789b806c15798f8e7dca4d8ac356b50ce70"
|
||||
checksum = "4b067c0c11094aef6b7a801c1e34a26affafdf3d051dba08456b868789aaf9a4"
|
||||
dependencies = [
|
||||
"cfg-if",
|
||||
"once_cell",
|
||||
|
|
@ -2581,9 +2582,9 @@ dependencies = [
|
|||
|
||||
[[package]]
|
||||
name = "wasm-bindgen-futures"
|
||||
version = "0.4.77"
|
||||
version = "0.4.76"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "6b7777d5cc23d0e91404e53ce2d5e8ec7acae3026b16233dba62cd3246457950"
|
||||
checksum = "c62df1340f32221cb9c54d6a27b030e3dba64361d4a95bed55f9aacb44da291d"
|
||||
dependencies = [
|
||||
"js-sys",
|
||||
"wasm-bindgen",
|
||||
|
|
@ -2591,9 +2592,9 @@ dependencies = [
|
|||
|
||||
[[package]]
|
||||
name = "wasm-bindgen-macro"
|
||||
version = "0.2.127"
|
||||
version = "0.2.126"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "77775f8f3f7217702089053b94958f8f54061a3f663417df76e19cbdcca29bc1"
|
||||
checksum = "167ce5e579f6bcf889c4f7175a8a5a585de84e8ff93976ce393efa5f2837aab1"
|
||||
dependencies = [
|
||||
"quote",
|
||||
"wasm-bindgen-macro-support",
|
||||
|
|
@ -2601,9 +2602,9 @@ dependencies = [
|
|||
|
||||
[[package]]
|
||||
name = "wasm-bindgen-macro-support"
|
||||
version = "0.2.127"
|
||||
version = "0.2.126"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "e11d33f857dc2fb11b8bc75aee111aa9cbeb12cd9f25efd3d4c2a3dd4e235284"
|
||||
checksum = "f3997c7839262f4ef12cf90b818d6340c18e80f263f1a94bf157d0ec4420380e"
|
||||
dependencies = [
|
||||
"bumpalo",
|
||||
"proc-macro2",
|
||||
|
|
@ -2614,18 +2615,18 @@ dependencies = [
|
|||
|
||||
[[package]]
|
||||
name = "wasm-bindgen-shared"
|
||||
version = "0.2.127"
|
||||
version = "0.2.126"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "7ef64dbcc55df09c7e5a46182d181c2cfa3e925f3da937ea764728b4bbb9dcbf"
|
||||
checksum = "dc1b4cb0cc549fcf58d7dfc081778139b3d283a081644e833e84682ad71cea24"
|
||||
dependencies = [
|
||||
"unicode-ident",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "wasm-bindgen-test"
|
||||
version = "0.3.77"
|
||||
version = "0.3.76"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "895a2607575412a4eda1df892084a375ea10dfeadc4d7d2ab87b854e4ddc7ba1"
|
||||
checksum = "2a0d555ca874445df8d314f94f5c948a4e74e5418f332c89f660a3d8310a96f4"
|
||||
dependencies = [
|
||||
"async-trait",
|
||||
"cast",
|
||||
|
|
@ -2645,9 +2646,9 @@ dependencies = [
|
|||
|
||||
[[package]]
|
||||
name = "wasm-bindgen-test-macro"
|
||||
version = "0.3.77"
|
||||
version = "0.3.76"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "4288cb0ebe215033bf949ae1fd046726daa4c32a157f24b9dc6ac387a52aa759"
|
||||
checksum = "94eb68555b95bcea5e8cf4abe280b529049479fa995bfc23734af96a6aedc120"
|
||||
dependencies = [
|
||||
"proc-macro2",
|
||||
"quote",
|
||||
|
|
@ -2656,9 +2657,9 @@ dependencies = [
|
|||
|
||||
[[package]]
|
||||
name = "wasm-bindgen-test-shared"
|
||||
version = "0.2.127"
|
||||
version = "0.2.126"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "33ff1c1b360982e93b6d8ea9c04836f71dba0817a16f91e229cf3a51bdd9d987"
|
||||
checksum = "c31d56021e873866c968588ed85ccdf56db5c426e44afdb4618c39895104b920"
|
||||
|
||||
[[package]]
|
||||
name = "wasm-tracing"
|
||||
|
|
@ -2789,9 +2790,9 @@ checksum = "589f6da84c646204747d1270a2a5661ea66ed1cced2631d546fdfb155959f9ec"
|
|||
|
||||
[[package]]
|
||||
name = "writeable"
|
||||
version = "0.6.4"
|
||||
version = "0.6.3"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "3ad82d2a33cdc9674dc7465672f271e096168fcdbe0f799d9e6db8c5892679dc"
|
||||
checksum = "1ffae5123b2d3fc086436f8834ae3ab053a283cfac8fe0a0b8eaae044768a4c4"
|
||||
|
||||
[[package]]
|
||||
name = "wtransport"
|
||||
|
|
@ -2936,9 +2937,9 @@ dependencies = [
|
|||
|
||||
[[package]]
|
||||
name = "zerotrie"
|
||||
version = "0.2.5"
|
||||
version = "0.2.4"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "4ea269c3bd32f0a32c321907a2ae912ba6f4649bb0fc764a15627e99a7095a3f"
|
||||
checksum = "0f9152d31db0792fa83f70fb2f83148effb5c1f5b8c7686c3459e361d9bc20bf"
|
||||
dependencies = [
|
||||
"displaydoc",
|
||||
"yoke",
|
||||
|
|
@ -2947,9 +2948,9 @@ dependencies = [
|
|||
|
||||
[[package]]
|
||||
name = "zerovec"
|
||||
version = "0.11.7"
|
||||
version = "0.11.6"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "94b5c6b5976d66c1d703c4fd17d3f5e43c8cedaacf604961b171adc7130896d8"
|
||||
checksum = "90f911cbc359ab6af17377d242225f4d75119aec87ea711a880987b18cd7b239"
|
||||
dependencies = [
|
||||
"yoke",
|
||||
"zerofrom",
|
||||
|
|
@ -2958,13 +2959,13 @@ dependencies = [
|
|||
|
||||
[[package]]
|
||||
name = "zerovec-derive"
|
||||
version = "0.11.5"
|
||||
version = "0.11.3"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "9f212a141d820099d57ffafb9569be9617a6f27d3dc881fbee8fb56642f917a9"
|
||||
checksum = "625dc425cab0dca6dc3c3319506e6593dcb08a9f387ea3b284dbd52a92c40555"
|
||||
dependencies = [
|
||||
"proc-macro2",
|
||||
"quote",
|
||||
"syn 3.0.3",
|
||||
"syn 2.0.119",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
|
|
|
|||
|
|
@ -113,5 +113,10 @@ tls = ["crypto", "mtp-crypto?/tls"]
|
|||
# Requires MTP_INSECURE_TLS=1 at runtime.
|
||||
insecure-tls = ["dep:mtp-transport", "mtp-transport?/insecure-tls"]
|
||||
|
||||
[dev-dependencies]
|
||||
tokio = { version = "1", features = ["full"] }
|
||||
rcgen = "0.14"
|
||||
rand = "0.10.1"
|
||||
|
||||
[package.metadata.cargo-machete]
|
||||
ignored = ["mtp-transport"]
|
||||
|
|
|
|||
23
README.md
23
README.md
|
|
@ -47,29 +47,18 @@ Feature summary:
|
|||
|
||||
| Feature | Pulls in | Enables |
|
||||
| --- | --- | --- |
|
||||
| `serde` | Crypto serialization support | Serde implementations for crypto key types |
|
||||
| `crypto` | `mtp::crypto` | AEAD, signatures, KEM, KDF, hashing, and connection authentication support |
|
||||
| `host` | `mtp::host` | Native QUIC host and version negotiation |
|
||||
| `client` | `mtp::client` | Native QUIC client connections |
|
||||
| `transport` | `mtp-transport` dependency | Low-level transport support; enabled automatically by `host` and `client` |
|
||||
| `pipes` | Pipe support in transport, host, client, and web server | Raw and encrypted byte streams |
|
||||
| `files` | `mtp::files` | `.mk` keyrings and `.mpkb` public bundles; also enables `crypto` |
|
||||
| `raw` | Raw file APIs | Legacy plaintext keyring migration APIs |
|
||||
| `web-server` | `mtp::webserver` | HTTPS server with HTTP/1.1, HTTP/2, HTTP/3, and WebTransport MTP sessions |
|
||||
| `full-server` | Native host and web-server surface | `host`, `web-server`, `crypto`, and `pipes` together |
|
||||
| `tls` | `mtp::crypto::tls` | Development self-signed certificate generation |
|
||||
| `insecure-tls` | Lower-level transport | Development-only certificate verification bypass, gated by `MTP_INSECURE_TLS=1` |
|
||||
| `crypto` | `mtp::crypto` | AEAD, signatures, KEM, KDF, hashing |
|
||||
| `host` | `mtp::host`, codec registry | QUIC host and version negotiation |
|
||||
| `client` | `mtp::client` | QUIC client connections |
|
||||
| `webserver` | `mtp::webserver` | HTTPS server with HTTP/1.1, HTTP/2, HTTP/3, and WebTransport MTP sessions |
|
||||
|
||||
The core modules always available from the facade are `codec`, `common`, and
|
||||
`type_map`. Native `client` and `host` modules re-export the transport policy
|
||||
types; the low-level transport crate is not exposed as `mtp::transport`. See the [native client](./docs/NATIVE-CLIENT.md) and [native host](./docs/NATIVE-HOST.md)
|
||||
The core crates are always available: `codec`, `transport`, `common`, and `type_map`. See the [native client](./docs/NATIVE-CLIENT.md) and [native host](./docs/NATIVE-HOST.md)
|
||||
guides for configuration and usage. See [Security](./docs/SECURITY.md) for security boundaries.
|
||||
|
||||
## Sub-crates
|
||||
|
||||
The `mtp` facade re-exports the following modules:
|
||||
`mtp::codec`, `mtp::common`, `mtp::type_map`, `mtp::crypto`, `mtp::host`,
|
||||
`mtp::client`, `mtp::files`, and `mtp::webserver` when their features are enabled.
|
||||
`mtp::codec`, `mtp::transport`, `mtp::common`, `mtp::type_map`, `mtp::crypto`, `mtp::host`, and `mtp::client`.
|
||||
|
||||
### Codec
|
||||
|
||||
|
|
|
|||
|
|
@ -13,10 +13,6 @@ use crate::error::AuthState;
|
|||
use crate::ping::{PingSession, start_ping_session};
|
||||
#[cfg(feature = "pipes")]
|
||||
use crate::pipe::PipeRequest;
|
||||
#[cfg(feature = "pipes")]
|
||||
use crate::pipe::is_expired_creation;
|
||||
#[cfg(feature = "pipes")]
|
||||
use crate::pipe::{PendingCreation, PendingCreationGuard};
|
||||
use crate::pipe::{PendingRequest, PipeDispatcher, run_dispatcher};
|
||||
|
||||
pub struct MTPConnection {
|
||||
|
|
@ -138,33 +134,17 @@ impl MTPConnection {
|
|||
description: &str,
|
||||
) -> Result<crate::pipe::PipeHandle, mtp_common::PipeError> {
|
||||
let (tx, rx) = tokio::sync::oneshot::channel();
|
||||
let token = Arc::new(());
|
||||
let pipe_id = {
|
||||
let mut pending = self
|
||||
.pipe_dispatcher
|
||||
.pending_creations
|
||||
.lock()
|
||||
.map_err(|_| mtp_common::PipeError::ConnectionClosed)?;
|
||||
let mut pending = self.pipe_dispatcher.pending_creations.lock().await;
|
||||
let pipe_id = loop {
|
||||
let candidate = rand::random::<u32>();
|
||||
if candidate != 0
|
||||
&& !pending.contains_key(&candidate)
|
||||
&& !is_expired_creation(&self.pipe_dispatcher, candidate)
|
||||
{
|
||||
if candidate != 0 && !pending.contains_key(&candidate) {
|
||||
break candidate;
|
||||
}
|
||||
};
|
||||
pending.insert(
|
||||
pipe_id,
|
||||
PendingCreation {
|
||||
token: token.clone(),
|
||||
sender: tx,
|
||||
},
|
||||
);
|
||||
pending.insert(pipe_id, tx);
|
||||
pipe_id
|
||||
};
|
||||
let mut creation_guard =
|
||||
PendingCreationGuard::new(self.pipe_dispatcher.clone(), pipe_id, token.clone());
|
||||
|
||||
let request = CommunicationValue::new_with_type_map(
|
||||
mtp_codec::CommunicationType::PipeRequest,
|
||||
|
|
@ -174,17 +154,19 @@ impl MTPConnection {
|
|||
.add_typed_default(DataType::Description, DataValue::Str(description.into()));
|
||||
|
||||
if let Err(error) = self.sender.send(&request).await {
|
||||
self.pipe_dispatcher
|
||||
.pending_creations
|
||||
.lock()
|
||||
.await
|
||||
.remove(&pipe_id);
|
||||
return Err(mtp_common::PipeError::from(error));
|
||||
}
|
||||
|
||||
creation_guard.disarm();
|
||||
Ok(crate::pipe::PipeHandle {
|
||||
pipe_id,
|
||||
description: description.to_string(),
|
||||
sender: self.sender.clone(),
|
||||
response_rx: rx,
|
||||
dispatcher: self.pipe_dispatcher.clone(),
|
||||
token,
|
||||
})
|
||||
}
|
||||
|
||||
|
|
@ -225,19 +207,18 @@ pub(crate) async fn connection_from_parts(
|
|||
|
||||
#[cfg(feature = "pipes")]
|
||||
{
|
||||
let receiver_queue_capacity = config.policy.receiver_queue_capacity.max(1);
|
||||
let (app_tx, app_rx) = mpsc::channel::<Result<CommunicationValue, CommunicationError>>(
|
||||
receiver_queue_capacity,
|
||||
config.policy.receiver_queue_capacity,
|
||||
);
|
||||
let (pipe_req_tx, pipe_req_rx) = mpsc::channel::<PipeRequest>(receiver_queue_capacity);
|
||||
let (pipe_req_tx, pipe_req_rx) =
|
||||
mpsc::channel::<PipeRequest>(config.policy.receiver_queue_capacity);
|
||||
|
||||
let dispatcher = Arc::new(PipeDispatcher {
|
||||
pending_requests: Mutex::new(std::collections::HashMap::new()),
|
||||
expired_requests: Mutex::new(std::collections::HashMap::new()),
|
||||
#[cfg(feature = "pipes")]
|
||||
type_map: type_map.clone(),
|
||||
pending_creations: std::sync::Mutex::new(std::collections::HashMap::new()),
|
||||
expired_creations: std::sync::Mutex::new(std::collections::HashMap::new()),
|
||||
pending_creations: Mutex::new(std::collections::HashMap::new()),
|
||||
pending_pipes: Mutex::new(std::collections::HashMap::new()),
|
||||
policy: Arc::new(config.policy),
|
||||
});
|
||||
|
|
@ -274,9 +255,8 @@ pub(crate) async fn connection_from_parts(
|
|||
|
||||
#[cfg(not(feature = "pipes"))]
|
||||
{
|
||||
let receiver_queue_capacity = config.policy.receiver_queue_capacity.max(1);
|
||||
let (app_tx, app_rx) = mpsc::channel::<Result<CommunicationValue, CommunicationError>>(
|
||||
receiver_queue_capacity,
|
||||
config.policy.receiver_queue_capacity,
|
||||
);
|
||||
let dispatcher = Arc::new(PipeDispatcher {
|
||||
pending_requests: Mutex::new(std::collections::HashMap::new()),
|
||||
|
|
|
|||
|
|
@ -222,10 +222,6 @@ impl MTPClient {
|
|||
sender.set_type_map(&tm).await;
|
||||
receiver.set_type_map(&tm).await;
|
||||
let version_str = format!("{}", PROTOCOL_VERSION);
|
||||
let public_key_bytes = keys
|
||||
.public_key_bundle()
|
||||
.try_as_bytes()
|
||||
.map_err(|error| CommunicationError::ParseError(error.to_string()))?;
|
||||
|
||||
let mut ident =
|
||||
CommunicationValue::new_with_type_map(CommunicationType::Identification, &tm)
|
||||
|
|
@ -237,7 +233,10 @@ impl MTPClient {
|
|||
// This capability marker lets a non-crypto host reject an
|
||||
// authentication attempt instead of treating it as a plain
|
||||
// unauthenticated connection.
|
||||
.add_typed_default(DataType::PublicKeys, DataValue::Bytes(public_key_bytes));
|
||||
.add_typed_default(
|
||||
DataType::PublicKeys,
|
||||
DataValue::Bytes(keys.public_key_bundle().as_bytes()),
|
||||
);
|
||||
if let Some(desc) = &config.description {
|
||||
ident = ident.add_typed_default(DataType::Description, DataValue::Str(desc.clone()));
|
||||
}
|
||||
|
|
@ -416,9 +415,7 @@ impl MTPClient {
|
|||
receiver.set_type_map(&tm).await;
|
||||
let version_str = format!("{}", PROTOCOL_VERSION);
|
||||
let pk_bundle = keys.public_key_bundle();
|
||||
let pk_bytes = pk_bundle
|
||||
.try_as_bytes()
|
||||
.map_err(|error| CommunicationError::ParseError(error.to_string()))?;
|
||||
let pk_bytes = pk_bundle.as_bytes();
|
||||
|
||||
let mut register = CommunicationValue::new_with_type_map(CommunicationType::Register, &tm)
|
||||
.add_typed_default(DataType::Version, DataValue::Str(version_str.clone()))
|
||||
|
|
@ -604,9 +601,7 @@ mod tests {
|
|||
#[cfg(feature = "pipes")]
|
||||
type_map: mtp_codec::TypeMap::latest(),
|
||||
#[cfg(feature = "pipes")]
|
||||
pending_creations: std::sync::Mutex::new(HashMap::new()),
|
||||
#[cfg(feature = "pipes")]
|
||||
expired_creations: std::sync::Mutex::new(HashMap::new()),
|
||||
pending_creations: Mutex::new(HashMap::new()),
|
||||
#[cfg(feature = "pipes")]
|
||||
pending_pipes: Mutex::new(HashMap::new()),
|
||||
#[cfg(feature = "pipes")]
|
||||
|
|
@ -644,9 +639,7 @@ mod tests {
|
|||
#[cfg(feature = "pipes")]
|
||||
type_map: mtp_codec::TypeMap::latest(),
|
||||
#[cfg(feature = "pipes")]
|
||||
pending_creations: std::sync::Mutex::new(HashMap::new()),
|
||||
#[cfg(feature = "pipes")]
|
||||
expired_creations: std::sync::Mutex::new(HashMap::new()),
|
||||
pending_creations: Mutex::new(HashMap::new()),
|
||||
#[cfg(feature = "pipes")]
|
||||
pending_pipes: Mutex::new(HashMap::new()),
|
||||
#[cfg(feature = "pipes")]
|
||||
|
|
|
|||
|
|
@ -66,8 +66,8 @@ pub(crate) async fn start_ping_session(
|
|||
return None;
|
||||
}
|
||||
|
||||
let (pong_tx, mut pong_rx) = mpsc::channel(1);
|
||||
receiver.observe_pongs_bounded(pong_tx).await;
|
||||
let (pong_tx, mut pong_rx) = mpsc::unbounded_channel();
|
||||
receiver.observe_pongs(pong_tx).await;
|
||||
let last_ping = Arc::new(Mutex::new(None));
|
||||
let ping_state = last_ping.clone();
|
||||
let interval = config.ping_interval;
|
||||
|
|
@ -75,7 +75,6 @@ pub(crate) async fn start_ping_session(
|
|||
let max_missed_pings = config.max_missed_pings;
|
||||
let ping_timestamp = config.ping_timestamp;
|
||||
let type_map = type_map.clone();
|
||||
let ping_receiver = receiver.clone();
|
||||
let mut close_rx = receiver.handle().subscribe_close();
|
||||
|
||||
let task = tokio::spawn(async move {
|
||||
|
|
@ -92,7 +91,6 @@ pub(crate) async fn start_ping_session(
|
|||
}
|
||||
_ = ticker.tick() => {
|
||||
let missed_pings = tracker.begin_round();
|
||||
ping_receiver.set_expected_pong_id(None).await;
|
||||
if max_missed_pings > 0 && missed_pings >= max_missed_pings {
|
||||
sender.close().await;
|
||||
break;
|
||||
|
|
@ -123,9 +121,7 @@ pub(crate) async fn start_ping_session(
|
|||
sender.close().await;
|
||||
break;
|
||||
};
|
||||
ping_receiver.set_expected_pong_id(Some(id)).await;
|
||||
if sender.send(&ping).await.is_err() {
|
||||
ping_receiver.set_expected_pong_id(None).await;
|
||||
sender.close().await;
|
||||
break;
|
||||
}
|
||||
|
|
|
|||
|
|
@ -5,8 +5,6 @@ use mtp_common::CommunicationError;
|
|||
use mtp_transport::Receiver;
|
||||
use std::collections::HashMap;
|
||||
use std::sync::Arc;
|
||||
#[cfg(feature = "pipes")]
|
||||
use std::sync::Mutex as StdMutex;
|
||||
use tokio::sync::{Mutex, mpsc};
|
||||
use tokio::time::{Duration, Instant};
|
||||
|
||||
|
|
@ -23,8 +21,6 @@ pub struct PipeHandle {
|
|||
pub(crate) description: String,
|
||||
pub(crate) sender: Sender,
|
||||
pub(crate) response_rx: tokio::sync::oneshot::Receiver<Result<bool, PipeError>>,
|
||||
pub(crate) dispatcher: Arc<PipeDispatcher>,
|
||||
pub(crate) token: Arc<()>,
|
||||
}
|
||||
|
||||
#[cfg(feature = "pipes")]
|
||||
|
|
@ -37,11 +33,9 @@ impl PipeHandle {
|
|||
&self.description
|
||||
}
|
||||
|
||||
pub async fn wait(mut self) -> Result<Option<mtp_transport::PipeWriter>, PipeError> {
|
||||
let response =
|
||||
tokio::time::timeout(self.dispatcher.policy.read_timeout, &mut self.response_rx).await;
|
||||
match response {
|
||||
Ok(Ok(Ok(true))) => {
|
||||
pub async fn wait(self) -> Result<Option<mtp_transport::PipeWriter>, PipeError> {
|
||||
match self.response_rx.await {
|
||||
Ok(Ok(true)) => {
|
||||
let writer = self
|
||||
.sender
|
||||
.open_pipe(self.pipe_id, &self.description)
|
||||
|
|
@ -49,27 +43,10 @@ impl PipeHandle {
|
|||
.map_err(PipeError::from)?;
|
||||
Ok(Some(writer))
|
||||
}
|
||||
Ok(Ok(Ok(false))) => Ok(None),
|
||||
Ok(Ok(Err(error))) => {
|
||||
expire_pending_creation(&self.dispatcher, self.pipe_id, &self.token);
|
||||
Err(error)
|
||||
Ok(Ok(false)) => Ok(None),
|
||||
Ok(Err(e)) => Err(e),
|
||||
Err(_) => Err(PipeError::StreamClosed),
|
||||
}
|
||||
Ok(Err(_)) => {
|
||||
expire_pending_creation(&self.dispatcher, self.pipe_id, &self.token);
|
||||
Err(PipeError::StreamClosed)
|
||||
}
|
||||
Err(_) => {
|
||||
expire_pending_creation(&self.dispatcher, self.pipe_id, &self.token);
|
||||
Err(PipeError::HandshakeTimeout)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(feature = "pipes")]
|
||||
impl Drop for PipeHandle {
|
||||
fn drop(&mut self) {
|
||||
expire_pending_creation(&self.dispatcher, self.pipe_id, &self.token);
|
||||
}
|
||||
}
|
||||
|
||||
|
|
@ -78,41 +55,9 @@ pub struct PipeRequest {
|
|||
pub(crate) pipe_id: u32,
|
||||
pub(crate) description: String,
|
||||
pub(crate) sender: Sender,
|
||||
pub(crate) receiver: Receiver,
|
||||
pub(crate) dispatcher: Arc<PipeDispatcher>,
|
||||
}
|
||||
|
||||
#[cfg(feature = "pipes")]
|
||||
struct ExpectedPipeGuard {
|
||||
receiver: Receiver,
|
||||
pipe_id: u32,
|
||||
armed: bool,
|
||||
}
|
||||
|
||||
#[cfg(feature = "pipes")]
|
||||
impl ExpectedPipeGuard {
|
||||
fn new(receiver: Receiver, pipe_id: u32) -> Self {
|
||||
Self {
|
||||
receiver,
|
||||
pipe_id,
|
||||
armed: true,
|
||||
}
|
||||
}
|
||||
|
||||
fn disarm(&mut self) {
|
||||
self.armed = false;
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(feature = "pipes")]
|
||||
impl Drop for ExpectedPipeGuard {
|
||||
fn drop(&mut self) {
|
||||
if self.armed {
|
||||
self.receiver.cancel_expected_pipe(self.pipe_id);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(feature = "pipes")]
|
||||
impl PipeRequest {
|
||||
pub fn id(&self) -> u32 {
|
||||
|
|
@ -124,10 +69,6 @@ impl PipeRequest {
|
|||
}
|
||||
|
||||
pub async fn accept(self) -> Result<mtp_transport::PipeReader, PipeError> {
|
||||
self.receiver
|
||||
.expect_pipe(self.pipe_id)
|
||||
.map_err(PipeError::from)?;
|
||||
let mut expected_pipe = ExpectedPipeGuard::new(self.receiver.clone(), self.pipe_id);
|
||||
let (pipe_tx, pipe_rx) = tokio::sync::oneshot::channel();
|
||||
{
|
||||
let mut pending = self.dispatcher.pending_pipes.lock().await;
|
||||
|
|
@ -151,10 +92,7 @@ impl PipeRequest {
|
|||
|
||||
let timeout = self.dispatcher.policy.read_timeout;
|
||||
match tokio::time::timeout(timeout, pipe_rx).await {
|
||||
Ok(Ok(reader)) => {
|
||||
expected_pipe.disarm();
|
||||
Ok(reader)
|
||||
}
|
||||
Ok(Ok(reader)) => Ok(reader),
|
||||
Ok(Err(_)) => {
|
||||
self.dispatcher
|
||||
.pending_pipes
|
||||
|
|
@ -191,54 +129,14 @@ pub(crate) struct PendingRequest {
|
|||
pub(crate) sender: tokio::sync::oneshot::Sender<Result<CommunicationValue, CommunicationError>>,
|
||||
}
|
||||
|
||||
#[cfg(feature = "pipes")]
|
||||
pub(crate) struct PendingCreation {
|
||||
pub(crate) token: Arc<()>,
|
||||
pub(crate) sender: tokio::sync::oneshot::Sender<Result<bool, PipeError>>,
|
||||
}
|
||||
|
||||
#[cfg(feature = "pipes")]
|
||||
pub(crate) struct PendingCreationGuard {
|
||||
dispatcher: Arc<PipeDispatcher>,
|
||||
pipe_id: u32,
|
||||
token: Arc<()>,
|
||||
armed: bool,
|
||||
}
|
||||
|
||||
#[cfg(feature = "pipes")]
|
||||
impl PendingCreationGuard {
|
||||
pub(crate) fn new(dispatcher: Arc<PipeDispatcher>, pipe_id: u32, token: Arc<()>) -> Self {
|
||||
Self {
|
||||
dispatcher,
|
||||
pipe_id,
|
||||
token,
|
||||
armed: true,
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) fn disarm(&mut self) {
|
||||
self.armed = false;
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(feature = "pipes")]
|
||||
impl Drop for PendingCreationGuard {
|
||||
fn drop(&mut self) {
|
||||
if self.armed {
|
||||
expire_pending_creation(&self.dispatcher, self.pipe_id, &self.token);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) struct PipeDispatcher {
|
||||
pub(crate) pending_requests: Mutex<HashMap<u32, PendingRequest>>,
|
||||
pub(crate) expired_requests: Mutex<HashMap<u32, Instant>>,
|
||||
#[cfg(feature = "pipes")]
|
||||
pub(crate) type_map: TypeMap,
|
||||
#[cfg(feature = "pipes")]
|
||||
pub(crate) pending_creations: StdMutex<HashMap<u32, PendingCreation>>,
|
||||
#[cfg(feature = "pipes")]
|
||||
pub(crate) expired_creations: StdMutex<HashMap<u32, Instant>>,
|
||||
pub(crate) pending_creations:
|
||||
Mutex<HashMap<u32, tokio::sync::oneshot::Sender<Result<bool, PipeError>>>>,
|
||||
#[cfg(feature = "pipes")]
|
||||
pub(crate) pending_pipes:
|
||||
Mutex<HashMap<u32, tokio::sync::oneshot::Sender<mtp_transport::PipeReader>>>,
|
||||
|
|
@ -246,90 +144,6 @@ pub(crate) struct PipeDispatcher {
|
|||
pub(crate) policy: Arc<Policy>,
|
||||
}
|
||||
|
||||
#[cfg(feature = "pipes")]
|
||||
const EXPIRED_CREATION_TOMBSTONE_TTL: Duration = Duration::from_secs(60);
|
||||
#[cfg(feature = "pipes")]
|
||||
const MAX_EXPIRED_CREATION_TOMBSTONES: usize = 1024;
|
||||
|
||||
#[cfg(feature = "pipes")]
|
||||
pub(crate) fn expire_pending_creation(dispatcher: &PipeDispatcher, pipe_id: u32, token: &Arc<()>) {
|
||||
let removed = dispatcher
|
||||
.pending_creations
|
||||
.lock()
|
||||
.ok()
|
||||
.and_then(|mut pending| {
|
||||
if pending
|
||||
.get(&pipe_id)
|
||||
.is_some_and(|entry| Arc::ptr_eq(&entry.token, token))
|
||||
{
|
||||
pending.remove(&pipe_id);
|
||||
Some(())
|
||||
} else {
|
||||
None
|
||||
}
|
||||
});
|
||||
if removed.is_none() {
|
||||
return;
|
||||
}
|
||||
let Ok(mut expired) = dispatcher.expired_creations.lock() else {
|
||||
return;
|
||||
};
|
||||
let now = Instant::now();
|
||||
expired.retain(|_, expires_at| *expires_at > now);
|
||||
if expired.len() >= MAX_EXPIRED_CREATION_TOMBSTONES
|
||||
&& let Some(oldest) = expired
|
||||
.iter()
|
||||
.min_by_key(|(_, expires_at)| **expires_at)
|
||||
.map(|(id, _)| *id)
|
||||
{
|
||||
expired.remove(&oldest);
|
||||
}
|
||||
expired.insert(pipe_id, now + EXPIRED_CREATION_TOMBSTONE_TTL);
|
||||
}
|
||||
|
||||
#[cfg(feature = "pipes")]
|
||||
fn consume_expired_creation(dispatcher: &PipeDispatcher, pipe_id: u32) -> bool {
|
||||
let Ok(mut expired) = dispatcher.expired_creations.lock() else {
|
||||
return false;
|
||||
};
|
||||
let now = Instant::now();
|
||||
expired.retain(|_, expires_at| *expires_at > now);
|
||||
expired.remove(&pipe_id).is_some()
|
||||
}
|
||||
|
||||
#[cfg(feature = "pipes")]
|
||||
pub(crate) fn is_expired_creation(dispatcher: &PipeDispatcher, pipe_id: u32) -> bool {
|
||||
let Ok(mut expired) = dispatcher.expired_creations.lock() else {
|
||||
return true;
|
||||
};
|
||||
let now = Instant::now();
|
||||
expired.retain(|_, expires_at| *expires_at > now);
|
||||
expired.contains_key(&pipe_id)
|
||||
}
|
||||
|
||||
#[cfg(feature = "pipes")]
|
||||
pub(crate) fn fail_pending_creations(dispatcher: &PipeDispatcher, error: &CommunicationError) {
|
||||
let pending = dispatcher
|
||||
.pending_creations
|
||||
.lock()
|
||||
.ok()
|
||||
.map(|mut pending| std::mem::take(&mut *pending));
|
||||
if let Some(pending) = pending {
|
||||
let error = PipeError::from(error.clone());
|
||||
for (_, pending) in pending {
|
||||
let _ = pending.sender.send(Err(error.clone()));
|
||||
}
|
||||
}
|
||||
if let Ok(mut expired) = dispatcher.expired_creations.lock() {
|
||||
expired.clear();
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(feature = "pipes")]
|
||||
pub(crate) async fn fail_pending_pipes(dispatcher: &PipeDispatcher) {
|
||||
dispatcher.pending_pipes.lock().await.clear();
|
||||
}
|
||||
|
||||
pub(crate) async fn route_message(
|
||||
msg: CommunicationValue,
|
||||
app_tx: &mpsc::Sender<Result<CommunicationValue, CommunicationError>>,
|
||||
|
|
@ -452,7 +266,6 @@ pub(crate) async fn run_dispatcher(
|
|||
pipe_id,
|
||||
description,
|
||||
sender: sender.clone(),
|
||||
receiver: receiver.clone(),
|
||||
dispatcher: dispatcher.clone(),
|
||||
};
|
||||
let _ = pipe_req_tx.send(req).await;
|
||||
|
|
@ -470,15 +283,9 @@ pub(crate) async fn run_dispatcher(
|
|||
continue;
|
||||
};
|
||||
let accepted = msg.get_bool(DataType::Accepted).unwrap_or(false);
|
||||
let pending = dispatcher
|
||||
.pending_creations
|
||||
.lock()
|
||||
.ok()
|
||||
.and_then(|mut pending| pending.remove(&pipe_id));
|
||||
if let Some(entry) = pending {
|
||||
let _ = entry.sender.send(Ok(accepted));
|
||||
} else {
|
||||
let _ = consume_expired_creation(&dispatcher, pipe_id);
|
||||
let mut pending = dispatcher.pending_creations.lock().await;
|
||||
if let Some(tx) = pending.remove(&pipe_id) {
|
||||
let _ = tx.send(Ok(accepted));
|
||||
}
|
||||
continue;
|
||||
}
|
||||
|
|
@ -496,10 +303,6 @@ pub(crate) async fn run_dispatcher(
|
|||
}
|
||||
Err(e) => {
|
||||
fail_pending_requests(&dispatcher, e.clone()).await;
|
||||
#[cfg(feature = "pipes")]
|
||||
fail_pending_creations(&dispatcher, &e);
|
||||
#[cfg(feature = "pipes")]
|
||||
fail_pending_pipes(&dispatcher).await;
|
||||
let _ = app_tx.send(Err(e)).await;
|
||||
break;
|
||||
}
|
||||
|
|
@ -522,10 +325,6 @@ pub(crate) async fn run_dispatcher(
|
|||
}
|
||||
Err(e) => {
|
||||
fail_pending_requests(&dispatcher, e.clone()).await;
|
||||
#[cfg(feature = "pipes")]
|
||||
fail_pending_creations(&dispatcher, &e);
|
||||
#[cfg(feature = "pipes")]
|
||||
fail_pending_pipes(&dispatcher).await;
|
||||
let _ = app_tx.send(Err(e)).await;
|
||||
break;
|
||||
}
|
||||
|
|
|
|||
6
codec/Cargo.lock
generated
6
codec/Cargo.lock
generated
|
|
@ -187,9 +187,9 @@ dependencies = [
|
|||
|
||||
[[package]]
|
||||
name = "chacha20"
|
||||
version = "0.10.2"
|
||||
version = "0.10.1"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "65c35e4b699c7e15ccbe7ee35c005e4fc0a278d22238a2857e6ce2dadeda1b06"
|
||||
checksum = "d524456ba66e72eb8b115ff89e01e497f8e6d11d78b70b1aa13c0fbd97540a81"
|
||||
dependencies = [
|
||||
"cfg-if",
|
||||
"cpufeatures 0.3.0",
|
||||
|
|
@ -1151,7 +1151,7 @@ version = "0.10.2"
|
|||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "c7f5fa3a058cd35567ef9bfa5e75732bee0f9e4c55fa90477bef2dfcdbc4be80"
|
||||
dependencies = [
|
||||
"chacha20 0.10.2",
|
||||
"chacha20 0.10.1",
|
||||
"getrandom 0.4.3",
|
||||
"rand_core 0.10.1",
|
||||
]
|
||||
|
|
|
|||
|
|
@ -2,7 +2,7 @@ use byteorder::{BigEndian, ReadBytesExt, WriteBytesExt};
|
|||
use std::fmt;
|
||||
use std::io::Cursor;
|
||||
|
||||
use crate::data_value::{DataKind, DataValue, DecodeError, DecodeLimits, EncodeLimits};
|
||||
use crate::data_value::{DataKind, DataValue, DecodeLimits};
|
||||
use crate::rand_u32;
|
||||
use mtp_common::CodecError;
|
||||
use mtp_type_map::{
|
||||
|
|
@ -260,50 +260,23 @@ impl CommunicationValue {
|
|||
|
||||
#[must_use]
|
||||
pub fn reply_to(&self, comm_type: CommunicationType) -> Self {
|
||||
let type_map = self
|
||||
.type_map
|
||||
.as_ref()
|
||||
.cloned()
|
||||
.unwrap_or_else(TypeMap::latest);
|
||||
let mut response = Self::new_with_type_map(comm_type, &type_map);
|
||||
let mut response = Self::new(comm_type);
|
||||
response.sender = self.receiver;
|
||||
response.receiver = self.sender;
|
||||
response
|
||||
}
|
||||
|
||||
/// Merge clear container fields after confirming both values use the same
|
||||
/// negotiated type map.
|
||||
pub fn try_merge(&mut self, other: &Self) -> Result<(), CodecError> {
|
||||
if let Some(error) = &self.mapping_error {
|
||||
return Err(error.clone());
|
||||
}
|
||||
let left = self.type_map().ok_or(CodecError::MissingTypeMap)?;
|
||||
let right = other.type_map().ok_or(CodecError::MissingTypeMap)?;
|
||||
if left.version != right.version {
|
||||
return Err(CodecError::TypeMapMismatch {
|
||||
expected: left.version.to_string(),
|
||||
actual: right.version.to_string(),
|
||||
});
|
||||
}
|
||||
if let Some(error) = &other.mapping_error {
|
||||
return Err(error.clone());
|
||||
}
|
||||
let other_entries = other
|
||||
.payload
|
||||
.container_entries()
|
||||
.ok_or(CodecError::InvalidEncoding)?;
|
||||
for (id, value) in other_entries {
|
||||
self.insert_data(*id, value.clone())?;
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
// Migrate to `try_merge` so a map mismatch cannot be silently recorded in
|
||||
// a frame that is later sent over the wire.
|
||||
#[deprecated(note = "migrate to try_merge to handle negotiated type-map mismatches")]
|
||||
pub fn merge(&mut self, other: &Self) {
|
||||
if let Err(error) = self.try_merge(other) {
|
||||
self.mapping_error.get_or_insert(error);
|
||||
if self.mapping_error.is_none() {
|
||||
self.mapping_error.clone_from(&other.mapping_error);
|
||||
}
|
||||
let Some(other_entries) = other.payload.container_entries() else {
|
||||
self.mapping_error
|
||||
.get_or_insert(CodecError::InvalidEncoding);
|
||||
return;
|
||||
};
|
||||
for (id, value) in other_entries {
|
||||
let _ = self.insert_data(*id, value.clone());
|
||||
}
|
||||
}
|
||||
|
||||
|
|
@ -342,22 +315,9 @@ impl CommunicationValue {
|
|||
}
|
||||
|
||||
pub fn to_bytes(&self) -> Result<Vec<u8>, CodecError> {
|
||||
self.to_bytes_with_limits(EncodeLimits::default())
|
||||
}
|
||||
|
||||
pub fn to_bytes_with_limits(&self, limits: EncodeLimits) -> Result<Vec<u8>, CodecError> {
|
||||
if let Some(error) = &self.mapping_error {
|
||||
return Err(error.clone());
|
||||
}
|
||||
let header_len = self.frame_header_len();
|
||||
let payload_limit = limits
|
||||
.max_output_size
|
||||
.checked_sub(header_len)
|
||||
.ok_or(CodecError::TooManyEntries)?;
|
||||
let payload = self.payload.to_bytes_with_limits(EncodeLimits {
|
||||
max_output_size: payload_limit,
|
||||
..limits
|
||||
})?;
|
||||
let mut body = Vec::new();
|
||||
body.write_u16::<BigEndian>(self.comm_type.0)
|
||||
.map_err(|_| CodecError::InvalidEncoding)?;
|
||||
|
|
@ -384,71 +344,44 @@ impl CommunicationValue {
|
|||
body.write_u64::<BigEndian>(receiver)
|
||||
.map_err(|_| CodecError::InvalidEncoding)?;
|
||||
}
|
||||
body.extend_from_slice(&payload);
|
||||
body.extend_from_slice(&self.payload.to_bytes()?);
|
||||
let length = u32::try_from(body.len()).map_err(|_| CodecError::TooManyEntries)?;
|
||||
let total_len = 4usize
|
||||
.checked_add(body.len())
|
||||
.ok_or(CodecError::TooManyEntries)?;
|
||||
if total_len > limits.max_output_size {
|
||||
return Err(CodecError::TooManyEntries);
|
||||
}
|
||||
let mut out = Vec::with_capacity(total_len);
|
||||
let mut out = Vec::with_capacity(4 + body.len());
|
||||
out.write_u32::<BigEndian>(length)
|
||||
.map_err(|_| CodecError::InvalidEncoding)?;
|
||||
out.extend_from_slice(&body);
|
||||
Ok(out)
|
||||
}
|
||||
|
||||
fn frame_header_len(&self) -> usize {
|
||||
4 + 2
|
||||
+ 1
|
||||
+ self.id.is_some() as usize * 4
|
||||
+ self.sender.is_some() as usize * 8
|
||||
+ self.receiver.is_some() as usize * 8
|
||||
}
|
||||
|
||||
pub fn from_bytes(bytes: &[u8]) -> Result<Self, CodecError> {
|
||||
Self::from_bytes_with_limits(bytes, DecodeLimits::default())
|
||||
}
|
||||
|
||||
pub fn from_bytes_with_limits(bytes: &[u8], limits: DecodeLimits) -> Result<Self, CodecError> {
|
||||
Self::try_from_bytes_with_limits(bytes, limits).map_err(|_| CodecError::InvalidEncoding)
|
||||
}
|
||||
|
||||
pub fn try_from_bytes(bytes: &[u8]) -> Result<Self, DecodeError> {
|
||||
Self::try_from_bytes_with_limits(bytes, DecodeLimits::default())
|
||||
}
|
||||
|
||||
pub fn try_from_bytes_with_limits(
|
||||
bytes: &[u8],
|
||||
limits: DecodeLimits,
|
||||
) -> Result<Self, DecodeError> {
|
||||
let mut cursor = Cursor::new(bytes);
|
||||
let length = cursor
|
||||
.read_u32::<BigEndian>()
|
||||
.map_err(|_| DecodeError::MalformedEncoding)? as usize;
|
||||
.map_err(|_| CodecError::InvalidEncoding)? as usize;
|
||||
let end = 4usize
|
||||
.checked_add(length)
|
||||
.ok_or(DecodeError::MalformedEncoding)?;
|
||||
.ok_or(CodecError::InvalidEncoding)?;
|
||||
if end != bytes.len() {
|
||||
return Err(DecodeError::MalformedEncoding);
|
||||
return Err(CodecError::InvalidEncoding);
|
||||
}
|
||||
let comm_type = CommunicationTypeId(
|
||||
cursor
|
||||
.read_u16::<BigEndian>()
|
||||
.map_err(|_| DecodeError::MalformedEncoding)?,
|
||||
.map_err(|_| CodecError::InvalidEncoding)?,
|
||||
);
|
||||
let flags = cursor
|
||||
.read_u8()
|
||||
.map_err(|_| DecodeError::MalformedEncoding)?;
|
||||
let flags = cursor.read_u8().map_err(|_| CodecError::InvalidEncoding)?;
|
||||
if flags & !FLAG_KNOWN != 0 {
|
||||
return Err(DecodeError::MalformedEncoding);
|
||||
return Err(CodecError::InvalidEncoding);
|
||||
}
|
||||
let id = if flags & FLAG_HAS_ID != 0 {
|
||||
Some(
|
||||
cursor
|
||||
.read_u32::<BigEndian>()
|
||||
.map_err(|_| DecodeError::MalformedEncoding)?,
|
||||
.map_err(|_| CodecError::InvalidEncoding)?,
|
||||
)
|
||||
} else {
|
||||
None
|
||||
|
|
@ -457,7 +390,7 @@ impl CommunicationValue {
|
|||
Some(
|
||||
cursor
|
||||
.read_u64::<BigEndian>()
|
||||
.map_err(|_| DecodeError::MalformedEncoding)?,
|
||||
.map_err(|_| CodecError::InvalidEncoding)?,
|
||||
)
|
||||
} else {
|
||||
None
|
||||
|
|
@ -466,14 +399,14 @@ impl CommunicationValue {
|
|||
Some(
|
||||
cursor
|
||||
.read_u64::<BigEndian>()
|
||||
.map_err(|_| DecodeError::MalformedEncoding)?,
|
||||
.map_err(|_| CodecError::InvalidEncoding)?,
|
||||
)
|
||||
} else {
|
||||
None
|
||||
};
|
||||
let payload = DataValue::read_from_with_diagnostics(&mut cursor, limits)?;
|
||||
let payload = DataValue::read_from_with_limits(&mut cursor, limits)?;
|
||||
if cursor.position() as usize != end {
|
||||
return Err(DecodeError::MalformedEncoding);
|
||||
return Err(CodecError::InvalidEncoding);
|
||||
}
|
||||
Ok(Self {
|
||||
id,
|
||||
|
|
@ -487,36 +420,13 @@ impl CommunicationValue {
|
|||
}
|
||||
|
||||
pub fn from_bytes_with(bytes: &[u8], type_map: &TypeMap) -> Result<Self, CodecError> {
|
||||
Self::try_from_bytes_with(bytes, type_map).map_err(|_| CodecError::InvalidEncoding)
|
||||
}
|
||||
|
||||
pub fn try_from_bytes_with(bytes: &[u8], type_map: &TypeMap) -> Result<Self, DecodeError> {
|
||||
Self::try_from_bytes_with_type_map_and_limits(bytes, type_map, DecodeLimits::default())
|
||||
}
|
||||
|
||||
pub fn try_from_bytes_with_type_map_and_limits(
|
||||
bytes: &[u8],
|
||||
type_map: &TypeMap,
|
||||
limits: DecodeLimits,
|
||||
) -> Result<Self, DecodeError> {
|
||||
let mut value = Self::try_from_bytes_with_limits(bytes, limits)?;
|
||||
let mut value = Self::from_bytes(bytes)?;
|
||||
value.set_type_map(type_map);
|
||||
Ok(value)
|
||||
}
|
||||
|
||||
#[cfg(feature = "registry")]
|
||||
pub fn migrate(&self, target: &TypeMap) -> Result<Self, CodecError> {
|
||||
self.migrate_with_limits(target, EncodeLimits::default())
|
||||
}
|
||||
|
||||
/// Migrate a clear frame while bounding the recursive traversal used to
|
||||
/// translate its type IDs.
|
||||
#[cfg(feature = "registry")]
|
||||
pub fn migrate_with_limits(
|
||||
&self,
|
||||
target: &TypeMap,
|
||||
limits: EncodeLimits,
|
||||
) -> Result<Self, CodecError> {
|
||||
if let Some(error) = &self.mapping_error {
|
||||
return Err(error.clone());
|
||||
}
|
||||
|
|
@ -531,8 +441,7 @@ impl CommunicationValue {
|
|||
.comm_id_enum(comm)
|
||||
.ok_or_else(|| CodecError::UnknownCommunicationType(comm_name.to_string()))?,
|
||||
);
|
||||
let mut context = MigrationContext::new(limits);
|
||||
let payload = migrate_data_value(&self.payload, source, target, &mut context)?;
|
||||
let payload = migrate_data_value(&self.payload, source, target)?;
|
||||
Ok(Self {
|
||||
id: self.id,
|
||||
comm_type,
|
||||
|
|
@ -545,63 +454,15 @@ impl CommunicationValue {
|
|||
}
|
||||
}
|
||||
|
||||
#[cfg(feature = "registry")]
|
||||
struct MigrationContext {
|
||||
limits: EncodeLimits,
|
||||
depth: usize,
|
||||
values: usize,
|
||||
}
|
||||
|
||||
#[cfg(feature = "registry")]
|
||||
impl MigrationContext {
|
||||
fn new(limits: EncodeLimits) -> Self {
|
||||
Self {
|
||||
limits,
|
||||
depth: 0,
|
||||
values: 0,
|
||||
}
|
||||
}
|
||||
|
||||
fn value(&mut self) -> Result<(), CodecError> {
|
||||
self.values = self
|
||||
.values
|
||||
.checked_add(1)
|
||||
.ok_or(CodecError::TooManyEntries)?;
|
||||
if self.values > self.limits.max_values {
|
||||
return Err(CodecError::TooManyEntries);
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn enter(&mut self) -> Result<(), CodecError> {
|
||||
self.depth = self
|
||||
.depth
|
||||
.checked_add(1)
|
||||
.ok_or(CodecError::TooManyEntries)?;
|
||||
if self.depth > self.limits.max_depth {
|
||||
return Err(CodecError::TooManyEntries);
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn leave(&mut self) {
|
||||
self.depth = self.depth.saturating_sub(1);
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(feature = "registry")]
|
||||
fn migrate_data_value(
|
||||
value: &DataValue,
|
||||
source: &TypeMap,
|
||||
target: &TypeMap,
|
||||
context: &mut MigrationContext,
|
||||
) -> Result<DataValue, CodecError> {
|
||||
context.value()?;
|
||||
match value {
|
||||
DataValue::Container(entries) => {
|
||||
context.enter()?;
|
||||
let count = u16::try_from(entries.len()).map_err(|_| CodecError::TooManyEntries)?;
|
||||
let mut migrated = Vec::with_capacity(usize::from(count));
|
||||
let mut migrated = Vec::with_capacity(entries.len());
|
||||
for (old_id, value) in entries {
|
||||
let name = source
|
||||
.data_type_name(old_id.0)
|
||||
|
|
@ -613,20 +474,16 @@ fn migrate_data_value(
|
|||
.data_id_enum(data)
|
||||
.ok_or_else(|| CodecError::UnknownDataType(name.to_string()))?,
|
||||
);
|
||||
migrated.push((new_id, migrate_data_value(value, source, target, context)?));
|
||||
migrated.push((new_id, migrate_data_value(value, source, target)?));
|
||||
}
|
||||
context.leave();
|
||||
Ok(DataValue::Container(migrated))
|
||||
}
|
||||
DataValue::Array(values) => {
|
||||
context.enter()?;
|
||||
let mut migrated = Vec::with_capacity(values.len());
|
||||
for value in values {
|
||||
migrated.push(migrate_data_value(value, source, target, context)?);
|
||||
}
|
||||
context.leave();
|
||||
Ok(DataValue::Array(migrated))
|
||||
}
|
||||
DataValue::Array(values) => Ok(DataValue::Array(
|
||||
values
|
||||
.iter()
|
||||
.map(|value| migrate_data_value(value, source, target))
|
||||
.collect::<Result<Vec<_>, _>>()?,
|
||||
)),
|
||||
#[cfg(feature = "crypto")]
|
||||
DataValue::Signed(_) | DataValue::Encrypted(_) => Err(CodecError::InvalidEncoding),
|
||||
scalar => Ok(scalar.clone()),
|
||||
|
|
@ -784,39 +641,6 @@ mod tests {
|
|||
assert_eq!(frame.get_data(DataType::Version), None);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn replies_retain_the_request_type_map() {
|
||||
let type_map = TypeMap::new(mtp_type_map::Version::new(3, 0));
|
||||
let request = CommunicationValue::new_with_type_map(CommunicationType::Ping, &type_map)
|
||||
.with_sender(7)
|
||||
.with_receiver(9);
|
||||
let reply = request.reply_to(CommunicationType::Pong);
|
||||
|
||||
assert_eq!(
|
||||
reply.type_map().map(|map| &map.version),
|
||||
Some(&type_map.version)
|
||||
);
|
||||
assert_eq!(reply.sender(), Some(9));
|
||||
assert_eq!(reply.receiver(), Some(7));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn try_merge_rejects_frames_from_different_type_maps() {
|
||||
let left_map = TypeMap::new(mtp_type_map::Version::new(3, 0));
|
||||
let right_map = TypeMap::new(mtp_type_map::Version::new(4, 0));
|
||||
let mut left = CommunicationValue::new_with_type_map(CommunicationType::Ping, &left_map);
|
||||
let right = CommunicationValue::new_with_type_map(CommunicationType::Ping, &right_map);
|
||||
|
||||
assert_eq!(
|
||||
left.try_merge(&right),
|
||||
Err(CodecError::TypeMapMismatch {
|
||||
expected: "3.0".into(),
|
||||
actual: "4.0".into(),
|
||||
})
|
||||
);
|
||||
assert_eq!(left.data_len(), 0);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn generic_payload_roundtrips_without_becoming_a_container() {
|
||||
let payload = DataValue::Array(vec![
|
||||
|
|
@ -916,19 +740,9 @@ mod tests {
|
|||
|
||||
let mut signer_public_keys = recipient.public_key_bundle();
|
||||
signer_public_keys.sig_cl_public_key = signer_public_key;
|
||||
signed.verify_with_policy(
|
||||
SENDER_ID,
|
||||
&signer_public_keys,
|
||||
ProtectionPurpose::from(1),
|
||||
crate::ProtectionPolicy::any_supported(),
|
||||
)?;
|
||||
signed.verify(SENDER_ID, &signer_public_keys, ProtectionPurpose::from(1))?;
|
||||
assert_eq!(
|
||||
signed.into_verified_with_policy(
|
||||
SENDER_ID,
|
||||
&signer_public_keys,
|
||||
ProtectionPurpose::from(1),
|
||||
crate::ProtectionPolicy::any_supported(),
|
||||
)?,
|
||||
signed.into_verified(SENDER_ID, &signer_public_keys, ProtectionPurpose::from(1))?,
|
||||
clear_payload
|
||||
);
|
||||
Ok(())
|
||||
|
|
|
|||
File diff suppressed because it is too large
Load diff
|
|
@ -11,33 +11,20 @@ pub use data_value::{
|
|||
ApplicationProtectionPurpose, EncryptedValue, MtpProtectionPurpose, ProtectionError,
|
||||
ProtectionPolicy, ProtectionPurpose, ProtectionPurposeError, SignaturePolicy, SignedValue,
|
||||
};
|
||||
pub use data_value::{
|
||||
DEFAULT_TRANSPORT_ALLOCATION_FACTOR, DataKind, DataValue, DecodeError, DecodeLimits,
|
||||
EncodeLimits,
|
||||
};
|
||||
pub use data_value::{DataKind, DataValue, DecodeLimits};
|
||||
pub use mtp_common::{CodecError, TimeError, unix_time_millis};
|
||||
#[cfg(feature = "crypto")]
|
||||
#[allow(deprecated)]
|
||||
pub use protected::{
|
||||
CURRENT_PROTECTED_VERSION, InMemoryReplayGuard, ProtectedError, ProtectedLimits,
|
||||
ProtectedMessageBuilder, ProtectedOpenOptions, ReplayError, ReplayGuard,
|
||||
VerifiedProtectedMessage, open_protected_checked, open_protected_with_checked,
|
||||
open_protected_with_keys_checked, open_protected_with_keys_without_replay,
|
||||
open_protected_with_without_replay, open_protected_without_replay, protected_claimed_signer_id,
|
||||
protected_claimed_signer_id_with_limits, protected_claimed_signer_id_with_options,
|
||||
CURRENT_PROTECTED_VERSION, InMemoryReplayGuard, ProtectedError, ProtectedMessageBuilder,
|
||||
ProtectedOpenOptions, ReplayError, ReplayGuard, VerifiedProtectedMessage, open_protected,
|
||||
open_protected_with, open_protected_with_keys, protected_claimed_signer_id,
|
||||
};
|
||||
#[cfg(feature = "crypto")]
|
||||
#[allow(deprecated)]
|
||||
pub use relay::{
|
||||
CURRENT_RELAY_VERSION, RelayError, RelayOpenOptions, SealedRelayBuilder, VerifiedRelayContent,
|
||||
CURRENT_RELAY_VERSION, RelayError, SealedRelayBuilder, VerifiedRelayContent,
|
||||
VerifiedRelayMetadata, forward_relay_frame, open_relay_content,
|
||||
open_relay_content_with_keyrings, open_relay_content_with_keyrings_and_limits,
|
||||
open_relay_content_with_keys, open_relay_content_with_limits,
|
||||
open_relay_content_with_limits_without_replay, open_relay_metadata_checked,
|
||||
open_relay_metadata_with_checked, open_relay_metadata_with_limits_checked,
|
||||
open_relay_metadata_with_limits_without_replay, open_relay_metadata_with_without_replay,
|
||||
open_relay_metadata_without_replay, relay_metadata_claimed_signer_id,
|
||||
relay_metadata_claimed_signer_id_with_limits, relay_metadata_claimed_signer_id_with_options,
|
||||
open_relay_content_with_keyrings, open_relay_content_with_keys, open_relay_metadata,
|
||||
open_relay_metadata_with, open_relay_metadata_with_keys, relay_metadata_claimed_signer_id,
|
||||
};
|
||||
|
||||
pub use mtp_type_map::{
|
||||
|
|
|
|||
|
|
@ -7,39 +7,15 @@
|
|||
#![cfg(feature = "crypto")]
|
||||
|
||||
use mtp_crypto::{Keyring, PublicKeyBundle, SignatureScheme};
|
||||
use mtp_type_map::{CommunicationType, DataType, DataTypeId, TypeMap};
|
||||
use std::collections::{HashSet, VecDeque};
|
||||
use mtp_type_map::{CommunicationType, DataType, TypeMap};
|
||||
use std::collections::HashSet;
|
||||
|
||||
use crate::{
|
||||
CommunicationValue, DataValue, DecodeLimits, EncodeLimits, ProtectionError, ProtectionPolicy,
|
||||
ProtectionPurpose,
|
||||
};
|
||||
use crate::{CommunicationValue, DataValue, ProtectionError, ProtectionPolicy, ProtectionPurpose};
|
||||
|
||||
/// The direct protected-message envelope schema version emitted by this
|
||||
/// codec.
|
||||
pub const CURRENT_PROTECTED_VERSION: u64 = 1;
|
||||
|
||||
/// Semantic limits for fields that are retained after a protected message is
|
||||
/// opened. These are intentionally separate from generic transport blobs.
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
|
||||
pub struct ProtectedLimits {
|
||||
pub max_message_id_bytes: usize,
|
||||
pub max_metadata_encoded_bytes: usize,
|
||||
pub max_signer_key_history: usize,
|
||||
pub max_decryption_key_history: usize,
|
||||
}
|
||||
|
||||
impl Default for ProtectedLimits {
|
||||
fn default() -> Self {
|
||||
Self {
|
||||
max_message_id_bytes: 256,
|
||||
max_metadata_encoded_bytes: 1024 * 1024,
|
||||
max_signer_key_history: 8,
|
||||
max_decryption_key_history: 8,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, thiserror::Error)]
|
||||
pub enum ProtectedError {
|
||||
#[error("value is not an application communication frame")]
|
||||
|
|
@ -70,8 +46,6 @@ pub enum ProtectedError {
|
|||
ReservedApplicationType(String),
|
||||
#[error("protected message was already accepted")]
|
||||
Replay,
|
||||
#[error("protected resource limit exceeded: {0}")]
|
||||
ResourceLimit(&'static str),
|
||||
#[error("protection error: {0}")]
|
||||
Protection(#[from] ProtectionError),
|
||||
#[error("replay guard error: {0}")]
|
||||
|
|
@ -104,39 +78,9 @@ pub enum ReplayError {
|
|||
|
||||
/// Small in-memory guard useful for tests and short-lived clients. Production
|
||||
/// consumers should implement [`ReplayGuard`] over persistent storage.
|
||||
#[derive(Debug)]
|
||||
#[derive(Debug, Default)]
|
||||
pub struct InMemoryReplayGuard {
|
||||
accepted: HashSet<(u64, String)>,
|
||||
order: VecDeque<(u64, String)>,
|
||||
capacity: usize,
|
||||
}
|
||||
|
||||
impl Default for InMemoryReplayGuard {
|
||||
fn default() -> Self {
|
||||
Self::with_capacity(10_000)
|
||||
}
|
||||
}
|
||||
|
||||
impl InMemoryReplayGuard {
|
||||
pub fn new(capacity: usize) -> Self {
|
||||
Self::with_capacity(capacity)
|
||||
}
|
||||
|
||||
pub fn with_capacity(capacity: usize) -> Self {
|
||||
Self {
|
||||
accepted: HashSet::new(),
|
||||
order: VecDeque::new(),
|
||||
capacity,
|
||||
}
|
||||
}
|
||||
|
||||
pub fn len(&self) -> usize {
|
||||
self.accepted.len()
|
||||
}
|
||||
|
||||
pub fn is_empty(&self) -> bool {
|
||||
self.accepted.is_empty()
|
||||
}
|
||||
}
|
||||
|
||||
impl ReplayGuard for InMemoryReplayGuard {
|
||||
|
|
@ -146,21 +90,7 @@ impl ReplayGuard for InMemoryReplayGuard {
|
|||
message_id: &str,
|
||||
_created_at: u64,
|
||||
) -> Result<bool, ReplayError> {
|
||||
let key = (signer_id, message_id.to_owned());
|
||||
if self.accepted.contains(&key) {
|
||||
return Ok(false);
|
||||
}
|
||||
if self.capacity == 0 {
|
||||
return Ok(false);
|
||||
}
|
||||
self.accepted.insert(key.clone());
|
||||
self.order.push_back(key);
|
||||
while self.accepted.len() > self.capacity {
|
||||
if let Some(oldest) = self.order.pop_front() {
|
||||
self.accepted.remove(&oldest);
|
||||
}
|
||||
}
|
||||
Ok(true)
|
||||
Ok(self.accepted.insert((signer_id, message_id.to_owned())))
|
||||
}
|
||||
}
|
||||
|
||||
|
|
@ -180,8 +110,6 @@ pub struct ProtectedMessageBuilder<'a> {
|
|||
type_map: Option<TypeMap>,
|
||||
frame_id: Option<u32>,
|
||||
expose_sender: bool,
|
||||
limits: ProtectedLimits,
|
||||
encode_limits: EncodeLimits,
|
||||
}
|
||||
|
||||
impl<'a> ProtectedMessageBuilder<'a> {
|
||||
|
|
@ -208,8 +136,6 @@ impl<'a> ProtectedMessageBuilder<'a> {
|
|||
type_map: None,
|
||||
frame_id: None,
|
||||
expose_sender: false,
|
||||
limits: ProtectedLimits::default(),
|
||||
encode_limits: EncodeLimits::default(),
|
||||
}
|
||||
}
|
||||
|
||||
|
|
@ -250,16 +176,6 @@ impl<'a> ProtectedMessageBuilder<'a> {
|
|||
self
|
||||
}
|
||||
|
||||
pub fn protected_limits(mut self, limits: ProtectedLimits) -> Self {
|
||||
self.limits = limits;
|
||||
self
|
||||
}
|
||||
|
||||
pub fn encode_limits(mut self, limits: EncodeLimits) -> Self {
|
||||
self.encode_limits = limits;
|
||||
self
|
||||
}
|
||||
|
||||
pub fn build(self) -> Result<CommunicationValue, ProtectedError> {
|
||||
let message_id = self.message_id.ok_or(ProtectedError::InvalidLayout(
|
||||
"protected builder requires a message ID",
|
||||
|
|
@ -272,9 +188,6 @@ impl<'a> ProtectedMessageBuilder<'a> {
|
|||
"protected identifiers must be non-empty",
|
||||
));
|
||||
}
|
||||
if message_id.len() > self.limits.max_message_id_bytes {
|
||||
return Err(ProtectedError::ResourceLimit("message ID"));
|
||||
}
|
||||
if self.recipients.is_empty() {
|
||||
return Err(ProtectedError::InvalidLayout(
|
||||
"protected builder requires at least one recipient",
|
||||
|
|
@ -304,17 +217,8 @@ impl<'a> ProtectedMessageBuilder<'a> {
|
|||
(created_at_id, DataValue::UnsignedNumber(created_at as u128)),
|
||||
(content_id, self.content),
|
||||
]);
|
||||
let signed = envelope.sign_with_limits(
|
||||
self.signer_id,
|
||||
self.signature_purpose,
|
||||
self.signer,
|
||||
self.encode_limits,
|
||||
)?;
|
||||
let encrypted = signed.encrypt_for_with_limits(
|
||||
&self.recipients,
|
||||
self.encryption_purpose,
|
||||
self.encode_limits,
|
||||
)?;
|
||||
let signed = envelope.sign(self.signer_id, self.signature_purpose, self.signer)?;
|
||||
let encrypted = signed.encrypt_for(&self.recipients, self.encryption_purpose)?;
|
||||
|
||||
let mut frame = CommunicationValue::new_with_type_map(application_type, &type_map)
|
||||
.with_receiver(self.final_recipient_id)
|
||||
|
|
@ -354,12 +258,6 @@ pub struct ProtectedOpenOptions {
|
|||
pub encryption_purpose: ProtectionPurpose,
|
||||
/// Signature algorithms accepted by the receiver.
|
||||
pub policy: ProtectionPolicy,
|
||||
/// Recursive and cumulative allocation policy used while opening.
|
||||
pub decode_limits: DecodeLimits,
|
||||
/// Bound used when reconstructing signed bytes for verification.
|
||||
pub encode_limits: EncodeLimits,
|
||||
/// Semantic limits for retained protected fields and key histories.
|
||||
pub protected_limits: ProtectedLimits,
|
||||
}
|
||||
|
||||
impl ProtectedOpenOptions {
|
||||
|
|
@ -374,41 +272,8 @@ impl ProtectedOpenOptions {
|
|||
signature_purpose,
|
||||
encryption_purpose,
|
||||
policy,
|
||||
decode_limits: DecodeLimits {
|
||||
max_depth: 64,
|
||||
max_values: 65_536,
|
||||
max_blob_size: 16 * 1024 * 1024,
|
||||
max_recipients: 64,
|
||||
max_allocated_bytes: 64 * 1024 * 1024,
|
||||
},
|
||||
encode_limits: EncodeLimits {
|
||||
max_depth: 64,
|
||||
max_values: 65_536,
|
||||
max_output_size: 16 * 1024 * 1024,
|
||||
},
|
||||
protected_limits: ProtectedLimits {
|
||||
max_message_id_bytes: 256,
|
||||
max_metadata_encoded_bytes: 1024 * 1024,
|
||||
max_signer_key_history: 8,
|
||||
max_decryption_key_history: 8,
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
pub const fn with_limits(
|
||||
mut self,
|
||||
decode_limits: DecodeLimits,
|
||||
protected_limits: ProtectedLimits,
|
||||
) -> Self {
|
||||
self.decode_limits = decode_limits;
|
||||
self.protected_limits = protected_limits;
|
||||
self
|
||||
}
|
||||
|
||||
pub const fn with_encode_limits(mut self, encode_limits: EncodeLimits) -> Self {
|
||||
self.encode_limits = encode_limits;
|
||||
self
|
||||
}
|
||||
}
|
||||
|
||||
fn protected_field_id(
|
||||
|
|
@ -463,26 +328,23 @@ fn validate_protected_frame(frame: &CommunicationValue) -> Result<TypeMap, Prote
|
|||
}
|
||||
|
||||
fn field<'a>(
|
||||
entries: &'a [(DataTypeId, DataValue)],
|
||||
value: &'a DataValue,
|
||||
data_type: DataType,
|
||||
type_map: &TypeMap,
|
||||
) -> Result<&'a DataValue, ProtectedError> {
|
||||
let field_id = protected_field_id(data_type, type_map)?;
|
||||
entries
|
||||
.iter()
|
||||
.find(|(id, _)| *id == field_id)
|
||||
.map(|(_, value)| value)
|
||||
value
|
||||
.get_field(protected_field_id(data_type, type_map)?)
|
||||
.ok_or(ProtectedError::InvalidLayout(
|
||||
"required protected field is missing",
|
||||
))
|
||||
}
|
||||
|
||||
fn unsigned_field(
|
||||
entries: &[(DataTypeId, DataValue)],
|
||||
value: &DataValue,
|
||||
data_type: DataType,
|
||||
type_map: &TypeMap,
|
||||
) -> Result<u128, ProtectedError> {
|
||||
field(entries, data_type, type_map)?
|
||||
field(value, data_type, type_map)?
|
||||
.as_unsigned_number()
|
||||
.ok_or(ProtectedError::InvalidLayout(
|
||||
"protected field is not unsigned",
|
||||
|
|
@ -490,11 +352,11 @@ fn unsigned_field(
|
|||
}
|
||||
|
||||
fn string_field(
|
||||
entries: &[(DataTypeId, DataValue)],
|
||||
value: &DataValue,
|
||||
data_type: DataType,
|
||||
type_map: &TypeMap,
|
||||
) -> Result<String, ProtectedError> {
|
||||
field(entries, data_type, type_map)?
|
||||
field(value, data_type, type_map)?
|
||||
.as_string()
|
||||
.filter(|value| !value.is_empty())
|
||||
.ok_or(ProtectedError::InvalidLayout(
|
||||
|
|
@ -502,28 +364,10 @@ fn string_field(
|
|||
))
|
||||
}
|
||||
|
||||
fn string_field_ref<'a>(
|
||||
entries: &'a [(DataTypeId, DataValue)],
|
||||
data_type: DataType,
|
||||
type_map: &TypeMap,
|
||||
) -> Result<&'a str, ProtectedError> {
|
||||
field(entries, data_type, type_map)?
|
||||
.as_str()
|
||||
.filter(|value| !value.is_empty())
|
||||
.ok_or(ProtectedError::InvalidLayout(
|
||||
"protected field is not a non-empty string",
|
||||
))
|
||||
}
|
||||
|
||||
fn protected_version(
|
||||
entries: &[(DataTypeId, DataValue)],
|
||||
type_map: &TypeMap,
|
||||
) -> Result<u64, ProtectedError> {
|
||||
fn protected_version(value: &DataValue, type_map: &TypeMap) -> Result<u64, ProtectedError> {
|
||||
let version_id = protected_field_id(DataType::ProtectedVersion, type_map)?;
|
||||
let version = entries
|
||||
.iter()
|
||||
.find(|(id, _)| *id == version_id)
|
||||
.map(|(_, value)| value)
|
||||
let version = value
|
||||
.get_field(version_id)
|
||||
.ok_or(ProtectedError::MissingProtectedVersion)?
|
||||
.as_unsigned_number()
|
||||
.ok_or(ProtectedError::InvalidLayout(
|
||||
|
|
@ -537,15 +381,10 @@ fn decrypt_protected_payload(
|
|||
frame: &CommunicationValue,
|
||||
keyrings: &[&Keyring],
|
||||
encryption_purpose: ProtectionPurpose,
|
||||
decode_limits: DecodeLimits,
|
||||
max_decryption_key_history: usize,
|
||||
) -> Result<DataValue, ProtectedError> {
|
||||
if keyrings.len() > max_decryption_key_history {
|
||||
return Err(ProtectedError::ResourceLimit("decryption key history"));
|
||||
}
|
||||
frame
|
||||
.payload()
|
||||
.decrypt_with_keyrings_and_limits(keyrings, encryption_purpose, decode_limits)
|
||||
.decrypt_with_keyrings(keyrings, encryption_purpose)
|
||||
.map_err(|error| match error {
|
||||
ProtectionError::NotEncrypted => ProtectedError::PayloadNotEncrypted,
|
||||
other => ProtectedError::Protection(other),
|
||||
|
|
@ -555,63 +394,22 @@ fn decrypt_protected_payload(
|
|||
/// Return the claimed signer ID after decryption, without verifying its
|
||||
/// signature. The value is untrusted and may only select the key history that
|
||||
/// is then bound to the same signer ID during the subsequent open.
|
||||
#[deprecated(note = "use protected_claimed_signer_id_with_limits; pass the receive DecodeLimits")]
|
||||
pub fn protected_claimed_signer_id(
|
||||
frame: &CommunicationValue,
|
||||
keyrings: &[&Keyring],
|
||||
encryption_purpose: ProtectionPurpose,
|
||||
) -> Result<u64, ProtectedError> {
|
||||
// Migrate to `protected_claimed_signer_id_with_limits` at receive boundaries.
|
||||
protected_claimed_signer_id_with_limits(
|
||||
frame,
|
||||
keyrings,
|
||||
encryption_purpose,
|
||||
DecodeLimits::default(),
|
||||
)
|
||||
}
|
||||
|
||||
pub fn protected_claimed_signer_id_with_limits(
|
||||
frame: &CommunicationValue,
|
||||
keyrings: &[&Keyring],
|
||||
encryption_purpose: ProtectionPurpose,
|
||||
decode_limits: DecodeLimits,
|
||||
) -> Result<u64, ProtectedError> {
|
||||
protected_claimed_signer_id_with_options(
|
||||
frame,
|
||||
keyrings,
|
||||
encryption_purpose,
|
||||
decode_limits,
|
||||
ProtectedLimits::default(),
|
||||
)
|
||||
}
|
||||
|
||||
/// Return the claimed signer ID while applying the complete receive policy.
|
||||
///
|
||||
/// This is deliberately separate from the compatibility decoder above: the
|
||||
/// claimed ID is used to select a signer-key history, so the decryption-key
|
||||
/// history bound must be the same bound used by the eventual open operation.
|
||||
pub fn protected_claimed_signer_id_with_options(
|
||||
frame: &CommunicationValue,
|
||||
keyrings: &[&Keyring],
|
||||
encryption_purpose: ProtectionPurpose,
|
||||
decode_limits: DecodeLimits,
|
||||
protected_limits: ProtectedLimits,
|
||||
) -> Result<u64, ProtectedError> {
|
||||
validate_protected_frame(frame)?;
|
||||
let decrypted = decrypt_protected_payload(
|
||||
frame,
|
||||
keyrings,
|
||||
encryption_purpose,
|
||||
decode_limits,
|
||||
protected_limits.max_decryption_key_history,
|
||||
)?;
|
||||
let decrypted = decrypt_protected_payload(frame, keyrings, encryption_purpose)?;
|
||||
let signed = decrypted
|
||||
.as_signed()
|
||||
.ok_or(ProtectedError::PayloadNotSigned)?;
|
||||
Ok(signed.signer_id)
|
||||
}
|
||||
|
||||
fn open_protected_with_impl<F>(
|
||||
/// Open a direct protected message using a resolver for trusted signer keys.
|
||||
/// The resolver receives a claimed, unverified signer ID only as a lookup key.
|
||||
pub fn open_protected_with<F>(
|
||||
frame: &CommunicationValue,
|
||||
keyrings: &[&Keyring],
|
||||
expected_signer_id: Option<u64>,
|
||||
|
|
@ -624,13 +422,7 @@ where
|
|||
{
|
||||
validate_protected_frame(frame)?;
|
||||
let type_map = frame.type_map().cloned().unwrap_or_else(TypeMap::latest);
|
||||
let decrypted = decrypt_protected_payload(
|
||||
frame,
|
||||
keyrings,
|
||||
options.encryption_purpose,
|
||||
options.decode_limits,
|
||||
options.protected_limits.max_decryption_key_history,
|
||||
)?;
|
||||
let decrypted = decrypt_protected_payload(frame, keyrings, options.encryption_purpose)?;
|
||||
let signed = decrypted
|
||||
.as_signed()
|
||||
.ok_or(ProtectedError::PayloadNotSigned)?;
|
||||
|
|
@ -645,54 +437,13 @@ where
|
|||
}
|
||||
let signer_keys = resolve_signer_keys(signed.signer_id)
|
||||
.ok_or(ProtectionError::SignerKeyNotFound(signed.signer_id))?;
|
||||
if signer_keys.len() > options.protected_limits.max_signer_key_history {
|
||||
return Err(ProtectedError::ResourceLimit("signer key history"));
|
||||
}
|
||||
open_decrypted_protected(frame, type_map, signed, &signer_keys, options, replay_guard)
|
||||
}
|
||||
|
||||
pub fn open_protected_with_checked<F>(
|
||||
frame: &CommunicationValue,
|
||||
keyrings: &[&Keyring],
|
||||
expected_signer_id: Option<u64>,
|
||||
resolve_signer_keys: F,
|
||||
options: ProtectedOpenOptions,
|
||||
replay_guard: &mut dyn ReplayGuard,
|
||||
) -> Result<VerifiedProtectedMessage, ProtectedError>
|
||||
where
|
||||
F: FnOnce(u64) -> Option<Vec<PublicKeyBundle>>,
|
||||
{
|
||||
open_protected_with_impl(
|
||||
frame,
|
||||
keyrings,
|
||||
expected_signer_id,
|
||||
resolve_signer_keys,
|
||||
options,
|
||||
Some(replay_guard),
|
||||
)
|
||||
}
|
||||
|
||||
pub fn open_protected_with_without_replay<F>(
|
||||
frame: &CommunicationValue,
|
||||
keyrings: &[&Keyring],
|
||||
expected_signer_id: Option<u64>,
|
||||
resolve_signer_keys: F,
|
||||
options: ProtectedOpenOptions,
|
||||
) -> Result<VerifiedProtectedMessage, ProtectedError>
|
||||
where
|
||||
F: FnOnce(u64) -> Option<Vec<PublicKeyBundle>>,
|
||||
{
|
||||
open_protected_with_impl(
|
||||
frame,
|
||||
keyrings,
|
||||
expected_signer_id,
|
||||
resolve_signer_keys,
|
||||
options,
|
||||
None,
|
||||
)
|
||||
}
|
||||
|
||||
fn open_protected_with_keys_impl(
|
||||
/// Open a direct protected message against already resolved trusted signer
|
||||
/// keys. The signer ID is mandatory so a key history cannot be applied to a
|
||||
/// different claimed identity.
|
||||
pub fn open_protected_with_keys(
|
||||
frame: &CommunicationValue,
|
||||
keyrings: &[&Keyring],
|
||||
expected_signer_id: u64,
|
||||
|
|
@ -701,13 +452,7 @@ fn open_protected_with_keys_impl(
|
|||
replay_guard: Option<&mut dyn ReplayGuard>,
|
||||
) -> Result<VerifiedProtectedMessage, ProtectedError> {
|
||||
let type_map = validate_protected_frame(frame)?;
|
||||
let decrypted = decrypt_protected_payload(
|
||||
frame,
|
||||
keyrings,
|
||||
options.encryption_purpose,
|
||||
options.decode_limits,
|
||||
options.protected_limits.max_decryption_key_history,
|
||||
)?;
|
||||
let decrypted = decrypt_protected_payload(frame, keyrings, options.encryption_purpose)?;
|
||||
let signed = decrypted
|
||||
.as_signed()
|
||||
.ok_or(ProtectedError::PayloadNotSigned)?;
|
||||
|
|
@ -728,73 +473,23 @@ fn open_protected_with_keys_impl(
|
|||
)
|
||||
}
|
||||
|
||||
pub fn open_protected_with_keys_checked(
|
||||
frame: &CommunicationValue,
|
||||
keyrings: &[&Keyring],
|
||||
expected_signer_id: u64,
|
||||
signer_public_keys: &[PublicKeyBundle],
|
||||
options: ProtectedOpenOptions,
|
||||
replay_guard: &mut dyn ReplayGuard,
|
||||
) -> Result<VerifiedProtectedMessage, ProtectedError> {
|
||||
open_protected_with_keys_impl(
|
||||
frame,
|
||||
keyrings,
|
||||
expected_signer_id,
|
||||
signer_public_keys,
|
||||
options,
|
||||
Some(replay_guard),
|
||||
)
|
||||
}
|
||||
|
||||
pub fn open_protected_with_keys_without_replay(
|
||||
frame: &CommunicationValue,
|
||||
keyrings: &[&Keyring],
|
||||
expected_signer_id: u64,
|
||||
signer_public_keys: &[PublicKeyBundle],
|
||||
options: ProtectedOpenOptions,
|
||||
) -> Result<VerifiedProtectedMessage, ProtectedError> {
|
||||
open_protected_with_keys_impl(
|
||||
frame,
|
||||
keyrings,
|
||||
expected_signer_id,
|
||||
signer_public_keys,
|
||||
options,
|
||||
None,
|
||||
)
|
||||
}
|
||||
|
||||
pub fn open_protected_checked(
|
||||
/// Open a direct protected message when the expected signer and one trusted
|
||||
/// public key are already known.
|
||||
pub fn open_protected(
|
||||
frame: &CommunicationValue,
|
||||
keyring: &Keyring,
|
||||
expected_signer_id: u64,
|
||||
signer_public_key: &PublicKeyBundle,
|
||||
options: ProtectedOpenOptions,
|
||||
replay_guard: &mut dyn ReplayGuard,
|
||||
replay_guard: Option<&mut dyn ReplayGuard>,
|
||||
) -> Result<VerifiedProtectedMessage, ProtectedError> {
|
||||
open_protected_with_keys_impl(
|
||||
open_protected_with_keys(
|
||||
frame,
|
||||
std::slice::from_ref(&keyring),
|
||||
expected_signer_id,
|
||||
std::slice::from_ref(signer_public_key),
|
||||
options,
|
||||
Some(replay_guard),
|
||||
)
|
||||
}
|
||||
|
||||
pub fn open_protected_without_replay(
|
||||
frame: &CommunicationValue,
|
||||
keyring: &Keyring,
|
||||
expected_signer_id: u64,
|
||||
signer_public_key: &PublicKeyBundle,
|
||||
options: ProtectedOpenOptions,
|
||||
) -> Result<VerifiedProtectedMessage, ProtectedError> {
|
||||
open_protected_with_keys_impl(
|
||||
frame,
|
||||
std::slice::from_ref(&keyring),
|
||||
expected_signer_id,
|
||||
std::slice::from_ref(signer_public_key),
|
||||
options,
|
||||
None,
|
||||
replay_guard,
|
||||
)
|
||||
}
|
||||
|
||||
|
|
@ -806,15 +501,11 @@ fn open_decrypted_protected(
|
|||
options: ProtectedOpenOptions,
|
||||
mut replay_guard: Option<&mut dyn ReplayGuard>,
|
||||
) -> Result<VerifiedProtectedMessage, ProtectedError> {
|
||||
if signer_public_keys.len() > options.protected_limits.max_signer_key_history {
|
||||
return Err(ProtectedError::ResourceLimit("signer key history"));
|
||||
}
|
||||
let matched_signer_key_index = signed.verify_with_key_history_index_and_limits(
|
||||
let matched_signer_key_index = signed.verify_with_key_history_index(
|
||||
signed.signer_id,
|
||||
signer_public_keys,
|
||||
options.signature_purpose,
|
||||
options.policy,
|
||||
options.encode_limits,
|
||||
)?;
|
||||
let receiver_id = frame.receiver().ok_or(ProtectedError::MissingReceiver)?;
|
||||
if options
|
||||
|
|
@ -829,23 +520,22 @@ fn open_decrypted_protected(
|
|||
{
|
||||
return Err(ProtectedError::SenderMismatch);
|
||||
}
|
||||
/* The authenticated value is already owned by the decoder. Keep this
|
||||
inspection borrowed so opening a large envelope does not clone it. */
|
||||
let envelope = signed
|
||||
.value
|
||||
.container_entries()
|
||||
.as_container()
|
||||
.ok_or(ProtectedError::MissingEnvelope)?;
|
||||
let version = protected_version(envelope, &type_map)?;
|
||||
let envelope = DataValue::Container(envelope);
|
||||
let version = protected_version(&envelope, &type_map)?;
|
||||
if version != CURRENT_PROTECTED_VERSION {
|
||||
return Err(ProtectedError::UnsupportedProtectedVersion(version));
|
||||
}
|
||||
let message_type = string_field(envelope, DataType::MessageType, &type_map)?;
|
||||
let message_type = string_field(&envelope, DataType::MessageType, &type_map)?;
|
||||
let application_type = validate_application_message_type(&message_type, &type_map)?;
|
||||
if frame.get_comm_type_enum() != Some(application_type) {
|
||||
return Err(ProtectedError::MessageTypeMismatch);
|
||||
}
|
||||
let final_recipient_id = u64::try_from(unsigned_field(
|
||||
envelope,
|
||||
&envelope,
|
||||
DataType::FinalRecipientId,
|
||||
&type_map,
|
||||
)?)
|
||||
|
|
@ -853,14 +543,10 @@ fn open_decrypted_protected(
|
|||
if final_recipient_id != receiver_id {
|
||||
return Err(ProtectedError::FinalRecipientMismatch);
|
||||
}
|
||||
let message_id = string_field_ref(envelope, DataType::MessageId, &type_map)?;
|
||||
if message_id.len() > options.protected_limits.max_message_id_bytes {
|
||||
return Err(ProtectedError::ResourceLimit("message ID"));
|
||||
}
|
||||
let message_id = message_id.to_owned();
|
||||
let created_at = u64::try_from(unsigned_field(envelope, DataType::CreatedAt, &type_map)?)
|
||||
let message_id = string_field(&envelope, DataType::MessageId, &type_map)?;
|
||||
let created_at = u64::try_from(unsigned_field(&envelope, DataType::CreatedAt, &type_map)?)
|
||||
.map_err(|_| ProtectedError::InvalidLayout("created-at value is out of range"))?;
|
||||
let content = field(envelope, DataType::Content, &type_map)?.clone();
|
||||
let content = field(&envelope, DataType::Content, &type_map)?.clone();
|
||||
|
||||
if let Some(guard) = replay_guard.as_mut()
|
||||
&& !guard.accept(signed.signer_id, &message_id, created_at)?
|
||||
|
|
@ -883,7 +569,7 @@ fn open_decrypted_protected(
|
|||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use mtp_crypto::{Ed25519Signer, Keyring, PublicKeyBundle};
|
||||
use mtp_crypto::{Ed25519Signer, Keyring};
|
||||
use mtp_type_map::{DataType, DataTypeId};
|
||||
|
||||
const SIGNATURE_PURPOSE: ProtectionPurpose = ProtectionPurpose(0x40);
|
||||
|
|
@ -898,98 +584,10 @@ mod tests {
|
|||
)
|
||||
}
|
||||
|
||||
// Keep the existing test cases concise while making the production API
|
||||
// choice explicit: every call below is routed to either the checked or
|
||||
// the named without-replay entry point.
|
||||
fn open_protected(
|
||||
frame: &CommunicationValue,
|
||||
keyring: &Keyring,
|
||||
expected_signer_id: u64,
|
||||
signer_public_key: &PublicKeyBundle,
|
||||
options: ProtectedOpenOptions,
|
||||
replay_guard: Option<&mut dyn ReplayGuard>,
|
||||
) -> Result<VerifiedProtectedMessage, ProtectedError> {
|
||||
match replay_guard {
|
||||
Some(replay_guard) => super::open_protected_checked(
|
||||
frame,
|
||||
keyring,
|
||||
expected_signer_id,
|
||||
signer_public_key,
|
||||
options,
|
||||
replay_guard,
|
||||
),
|
||||
None => super::open_protected_without_replay(
|
||||
frame,
|
||||
keyring,
|
||||
expected_signer_id,
|
||||
signer_public_key,
|
||||
options,
|
||||
),
|
||||
}
|
||||
}
|
||||
|
||||
fn open_protected_with<F>(
|
||||
frame: &CommunicationValue,
|
||||
keyrings: &[&Keyring],
|
||||
expected_signer_id: Option<u64>,
|
||||
resolve_signer_keys: F,
|
||||
options: ProtectedOpenOptions,
|
||||
replay_guard: Option<&mut dyn ReplayGuard>,
|
||||
) -> Result<VerifiedProtectedMessage, ProtectedError>
|
||||
where
|
||||
F: FnOnce(u64) -> Option<Vec<PublicKeyBundle>>,
|
||||
{
|
||||
match replay_guard {
|
||||
Some(replay_guard) => super::open_protected_with_checked(
|
||||
frame,
|
||||
keyrings,
|
||||
expected_signer_id,
|
||||
resolve_signer_keys,
|
||||
options,
|
||||
replay_guard,
|
||||
),
|
||||
None => super::open_protected_with_without_replay(
|
||||
frame,
|
||||
keyrings,
|
||||
expected_signer_id,
|
||||
resolve_signer_keys,
|
||||
options,
|
||||
),
|
||||
}
|
||||
}
|
||||
|
||||
fn open_protected_with_keys(
|
||||
frame: &CommunicationValue,
|
||||
keyrings: &[&Keyring],
|
||||
expected_signer_id: u64,
|
||||
signer_public_keys: &[PublicKeyBundle],
|
||||
options: ProtectedOpenOptions,
|
||||
replay_guard: Option<&mut dyn ReplayGuard>,
|
||||
) -> Result<VerifiedProtectedMessage, ProtectedError> {
|
||||
match replay_guard {
|
||||
Some(replay_guard) => super::open_protected_with_keys_checked(
|
||||
frame,
|
||||
keyrings,
|
||||
expected_signer_id,
|
||||
signer_public_keys,
|
||||
options,
|
||||
replay_guard,
|
||||
),
|
||||
None => super::open_protected_with_keys_without_replay(
|
||||
frame,
|
||||
keyrings,
|
||||
expected_signer_id,
|
||||
signer_public_keys,
|
||||
options,
|
||||
),
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Default)]
|
||||
struct RecordingReplayGuard {
|
||||
created_at: Option<u64>,
|
||||
accepted: bool,
|
||||
calls: usize,
|
||||
}
|
||||
|
||||
impl ReplayGuard for RecordingReplayGuard {
|
||||
|
|
@ -999,7 +597,6 @@ mod tests {
|
|||
_message_id: &str,
|
||||
created_at: u64,
|
||||
) -> Result<bool, ReplayError> {
|
||||
self.calls += 1;
|
||||
self.created_at = Some(created_at);
|
||||
if self.accepted {
|
||||
Ok(false)
|
||||
|
|
@ -1010,38 +607,6 @@ mod tests {
|
|||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn in_memory_replay_guard_is_bounded_and_deduplicates() {
|
||||
let mut guard = InMemoryReplayGuard::with_capacity(2);
|
||||
assert!(guard.accept(7, "first", 1).expect("first replay decision"));
|
||||
assert!(
|
||||
guard
|
||||
.accept(7, "second", 2)
|
||||
.expect("second replay decision")
|
||||
);
|
||||
assert!(
|
||||
!guard
|
||||
.accept(7, "first", 3)
|
||||
.expect("duplicate replay decision")
|
||||
);
|
||||
assert_eq!(guard.len(), 2);
|
||||
|
||||
assert!(guard.accept(7, "third", 4).expect("third replay decision"));
|
||||
assert_eq!(guard.len(), 2);
|
||||
assert!(
|
||||
guard
|
||||
.accept(7, "first", 5)
|
||||
.expect("evicted replay decision")
|
||||
);
|
||||
|
||||
let mut disabled = InMemoryReplayGuard::with_capacity(0);
|
||||
assert!(
|
||||
!disabled
|
||||
.accept(7, "disabled", 1)
|
||||
.expect("disabled replay decision")
|
||||
);
|
||||
}
|
||||
|
||||
fn protected_field(data_type: DataType, type_map: &TypeMap) -> DataTypeId {
|
||||
data_type
|
||||
.try_to_id(type_map)
|
||||
|
|
@ -1178,7 +743,6 @@ mod tests {
|
|||
assert_eq!(opened.message_type, "ProtectedMessage");
|
||||
assert_eq!(opened.message_id, "protected-test");
|
||||
assert_eq!(guard.created_at, Some(1_700_000_000_000));
|
||||
assert_eq!(guard.calls, 1);
|
||||
assert_eq!(opened.content, DataValue::Str("hello".into()));
|
||||
assert!(matches!(
|
||||
open_protected(
|
||||
|
|
@ -1193,29 +757,6 @@ mod tests {
|
|||
));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn oversized_message_id_is_rejected_before_replay_guard() {
|
||||
let sender = Keyring::generate();
|
||||
let recipient = Keyring::generate();
|
||||
let frame = valid_frame(&sender, &recipient, DataValue::Str("hello".into()));
|
||||
let mut options = open_options(Some(42));
|
||||
options.protected_limits.max_message_id_bytes = 3;
|
||||
let mut guard = RecordingReplayGuard::default();
|
||||
|
||||
assert!(matches!(
|
||||
open_protected_checked(
|
||||
&frame,
|
||||
&recipient,
|
||||
7,
|
||||
&sender.public_key_bundle(),
|
||||
options,
|
||||
&mut guard,
|
||||
),
|
||||
Err(ProtectedError::ResourceLimit("message ID"))
|
||||
));
|
||||
assert_eq!(guard.calls, 0);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn builder_owns_outer_sender_and_frame_id() {
|
||||
let sender = Keyring::generate();
|
||||
|
|
|
|||
|
|
@ -2,7 +2,6 @@ use mtp_common::CodecError;
|
|||
use mtp_type_map::{PROTOCOL_VERSION, TypeMap, Version};
|
||||
|
||||
use crate::CommunicationValue;
|
||||
use crate::EncodeLimits;
|
||||
|
||||
pub use mtp_type_map::Registry;
|
||||
|
||||
|
|
@ -43,41 +42,7 @@ impl VersionedCodec {
|
|||
|
||||
/// Encode a value using the codec's negotiated framing rules.
|
||||
pub fn encode(&self, value: &CommunicationValue) -> Result<Vec<u8>, CodecError> {
|
||||
self.encode_with_limits(value, EncodeLimits::default())
|
||||
}
|
||||
|
||||
/// Encode using an explicit output/resource limit after verifying the
|
||||
/// value belongs to this codec's negotiated type map.
|
||||
pub fn encode_with_limits(
|
||||
&self,
|
||||
value: &CommunicationValue,
|
||||
limits: EncodeLimits,
|
||||
) -> Result<Vec<u8>, CodecError> {
|
||||
let value_map = value.type_map().ok_or(CodecError::MissingTypeMap)?;
|
||||
if value_map.version != self.type_map.version {
|
||||
return Err(CodecError::TypeMapMismatch {
|
||||
expected: self.type_map.version.to_string(),
|
||||
actual: value_map.version.to_string(),
|
||||
});
|
||||
}
|
||||
value.to_bytes_with_limits(limits)
|
||||
}
|
||||
|
||||
/// Explicitly migrate a clear frame to this codec's negotiated type map
|
||||
/// before encoding it.
|
||||
pub fn encode_migrating(&self, value: &CommunicationValue) -> Result<Vec<u8>, CodecError> {
|
||||
self.encode_migrating_with_limits(value, EncodeLimits::default())
|
||||
}
|
||||
|
||||
/// Explicitly migrate and encode with bounded traversal/output.
|
||||
pub fn encode_migrating_with_limits(
|
||||
&self,
|
||||
value: &CommunicationValue,
|
||||
limits: EncodeLimits,
|
||||
) -> Result<Vec<u8>, CodecError> {
|
||||
value
|
||||
.migrate_with_limits(&self.type_map, limits)?
|
||||
.to_bytes_with_limits(limits)
|
||||
value.to_bytes()
|
||||
}
|
||||
|
||||
/// Decode a frame and retain the negotiated type map for typed access.
|
||||
|
|
@ -93,34 +58,3 @@ impl VersionedCodec {
|
|||
&self.registry
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use crate::DataValue;
|
||||
use mtp_type_map::{CommunicationType, Version};
|
||||
|
||||
#[test]
|
||||
fn encode_rejects_a_value_from_another_negotiated_map() {
|
||||
let mut registry = Registry::new();
|
||||
let version_a = Version::new(3, 0);
|
||||
let version_b = Version::new(4, 0);
|
||||
registry.register(TypeMap::new(version_a.clone()));
|
||||
registry.register(TypeMap::new(version_b.clone()));
|
||||
|
||||
let codec = VersionedCodec::for_version(registry, version_b).expect("codec version");
|
||||
let value = CommunicationValue::new_with_type_map(
|
||||
CommunicationType::Ping,
|
||||
&TypeMap::new(version_a.clone()),
|
||||
)
|
||||
.with_payload(DataValue::Null);
|
||||
|
||||
assert_eq!(
|
||||
codec.encode(&value),
|
||||
Err(CodecError::TypeMapMismatch {
|
||||
expected: "4.0".into(),
|
||||
actual: "3.0".into(),
|
||||
})
|
||||
);
|
||||
}
|
||||
}
|
||||
|
|
|
|||
File diff suppressed because it is too large
Load diff
4
common/Cargo.lock
generated
4
common/Cargo.lock
generated
|
|
@ -139,9 +139,9 @@ checksum = "f079e83a288787bcd14a6aea84cee5c87a67c5a3e660c30f557a3d24761b3527"
|
|||
|
||||
[[package]]
|
||||
name = "chacha20"
|
||||
version = "0.10.2"
|
||||
version = "0.10.1"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "65c35e4b699c7e15ccbe7ee35c005e4fc0a278d22238a2857e6ce2dadeda1b06"
|
||||
checksum = "d524456ba66e72eb8b115ff89e01e497f8e6d11d78b70b1aa13c0fbd97540a81"
|
||||
dependencies = [
|
||||
"cfg-if",
|
||||
"cpufeatures",
|
||||
|
|
|
|||
|
|
@ -15,6 +15,7 @@ wtransport = { version = "0.7.1", default-features = false, features = [
|
|||
"quinn",
|
||||
"self-signed",
|
||||
] }
|
||||
rustls = { version = "0.23.41" }
|
||||
quinn = { version = "0.11.11", default-features = false, features = [
|
||||
"rustls-aws-lc-rs",
|
||||
"rustls",
|
||||
|
|
|
|||
|
|
@ -41,10 +41,6 @@ pub enum CodecError {
|
|||
InvalidEncoding,
|
||||
#[error("Too many entries to encode")]
|
||||
TooManyEntries,
|
||||
#[error("Missing negotiated type map")]
|
||||
MissingTypeMap,
|
||||
#[error("Type-map mismatch: expected {expected}, actual {actual}")]
|
||||
TypeMapMismatch { expected: String, actual: String },
|
||||
#[error("Crypto failed: {0}")]
|
||||
CryptoFailed(String),
|
||||
#[error("Missing required field: {0}")]
|
||||
|
|
@ -164,9 +160,6 @@ pub enum CommunicationError {
|
|||
#[error("Stream Error")]
|
||||
StreamError,
|
||||
|
||||
#[error("Stream failed after delivery may have started")]
|
||||
DeliveryUnknown,
|
||||
|
||||
#[error("Stream Error: {0}")]
|
||||
#[cfg(not(target_arch = "wasm32"))]
|
||||
StreamWriteError(#[from] wtransport::error::StreamWriteError),
|
||||
|
|
@ -185,38 +178,6 @@ pub enum CommunicationError {
|
|||
Other(String),
|
||||
}
|
||||
|
||||
/// How the protocol layer should handle the first frame on a receive stream.
|
||||
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
|
||||
pub enum FirstFrameDisposition {
|
||||
Message,
|
||||
Pipe(u32),
|
||||
}
|
||||
|
||||
/// Classify a first frame without tying the decision to a WebTransport backend.
|
||||
///
|
||||
/// `PipeRequest` is used both as a control message and as the header of the raw
|
||||
/// stream opened after that request is accepted. Only the protocol layer knows
|
||||
/// which raw stream IDs are currently expected.
|
||||
pub fn classify_first_frame(
|
||||
is_pipe_request: bool,
|
||||
pipe_id: Option<u32>,
|
||||
pipe_is_expected: bool,
|
||||
) -> Result<FirstFrameDisposition, CommunicationError> {
|
||||
if !is_pipe_request {
|
||||
return Ok(FirstFrameDisposition::Message);
|
||||
}
|
||||
|
||||
let pipe_id = pipe_id.filter(|id| *id != 0).ok_or_else(|| {
|
||||
CommunicationError::Other("PipeRequest frame must contain a non-zero id".into())
|
||||
})?;
|
||||
|
||||
if pipe_is_expected {
|
||||
Ok(FirstFrameDisposition::Pipe(pipe_id))
|
||||
} else {
|
||||
Ok(FirstFrameDisposition::Message)
|
||||
}
|
||||
}
|
||||
|
||||
// ---- manual PartialEq (quinn / wtransport types don't impl PartialEq) ----
|
||||
|
||||
impl PartialEq for CommunicationError {
|
||||
|
|
@ -247,7 +208,6 @@ impl PartialEq for CommunicationError {
|
|||
(Self::ReadExactError(_), Self::ReadExactError(_)) => true,
|
||||
(Self::StreamClosed, Self::StreamClosed) => true,
|
||||
(Self::StreamError, Self::StreamError) => true,
|
||||
(Self::DeliveryUnknown, Self::DeliveryUnknown) => true,
|
||||
#[cfg(not(target_arch = "wasm32"))]
|
||||
(Self::StreamWriteError(_), Self::StreamWriteError(_)) => true,
|
||||
#[cfg(not(target_arch = "wasm32"))]
|
||||
|
|
|
|||
|
|
@ -1,243 +0,0 @@
|
|||
#!/usr/bin/env node
|
||||
|
||||
import { execFile, spawn } from "node:child_process";
|
||||
import { access, cp, mkdir, mkdtemp, readFile, rm, writeFile } from "node:fs/promises";
|
||||
import os from "node:os";
|
||||
import path from "node:path";
|
||||
import { promisify } from "node:util";
|
||||
import { fileURLToPath } from "node:url";
|
||||
|
||||
const execFileAsync = promisify(execFile);
|
||||
const repositoryRoot = path.resolve(path.dirname(fileURLToPath(import.meta.url)), ".");
|
||||
const packageJsonPath = path.join(repositoryRoot, "package.json");
|
||||
|
||||
function usage() {
|
||||
return `Usage: node create-web-release.mjs [options]
|
||||
|
||||
Build and pack the browser package using the version of the root Cargo package.
|
||||
|
||||
Options:
|
||||
--skip-build Pack the existing dist/ and wasm/pkg/ artifacts
|
||||
--output-dir <path> Write the archive to this directory (default: repository root)
|
||||
--help Show this help
|
||||
`;
|
||||
}
|
||||
|
||||
function parseArguments(arguments_) {
|
||||
const options = {
|
||||
outputDir: repositoryRoot,
|
||||
skipBuild: false,
|
||||
};
|
||||
|
||||
for (let index = 0; index < arguments_.length; index += 1) {
|
||||
const argument = arguments_[index];
|
||||
if (argument === "--help") {
|
||||
options.help = true;
|
||||
} else if (argument === "--skip-build") {
|
||||
options.skipBuild = true;
|
||||
} else if (argument === "--output-dir") {
|
||||
const outputDir = arguments_[index + 1];
|
||||
if (!outputDir || outputDir.startsWith("--")) {
|
||||
throw new Error("--output-dir requires a directory path");
|
||||
}
|
||||
options.outputDir = path.resolve(repositoryRoot, outputDir);
|
||||
index += 1;
|
||||
} else if (argument.startsWith("--output-dir=")) {
|
||||
const outputDir = argument.slice("--output-dir=".length);
|
||||
if (!outputDir) {
|
||||
throw new Error("--output-dir requires a directory path");
|
||||
}
|
||||
options.outputDir = path.resolve(repositoryRoot, outputDir);
|
||||
} else {
|
||||
throw new Error(`Unknown option: ${argument}`);
|
||||
}
|
||||
}
|
||||
|
||||
return options;
|
||||
}
|
||||
|
||||
async function readJson(filePath) {
|
||||
const source = await readFile(filePath, "utf8");
|
||||
try {
|
||||
return JSON.parse(source);
|
||||
} catch (error) {
|
||||
throw new Error(`Invalid JSON in ${path.relative(repositoryRoot, filePath)}`, {
|
||||
cause: error,
|
||||
});
|
||||
}
|
||||
}
|
||||
|
||||
async function run(command, arguments_, options = {}) {
|
||||
const renderedArguments = arguments_.map((argument) => JSON.stringify(argument)).join(" ");
|
||||
console.log(`\n> ${command}${renderedArguments ? ` ${renderedArguments}` : ""}`);
|
||||
|
||||
await new Promise((resolve, reject) => {
|
||||
const child = spawn(command, arguments_, {
|
||||
cwd: options.cwd ?? repositoryRoot,
|
||||
env: options.env ?? process.env,
|
||||
stdio: "inherit",
|
||||
});
|
||||
|
||||
child.once("error", (error) => {
|
||||
reject(new Error(`Failed to run ${command}: ${error.message}`, { cause: error }));
|
||||
});
|
||||
child.once("exit", (code, signal) => {
|
||||
if (code === 0) {
|
||||
resolve();
|
||||
return;
|
||||
}
|
||||
|
||||
const reason = signal ? `signal ${signal}` : `exit code ${code}`;
|
||||
reject(new Error(`${command} failed with ${reason}`));
|
||||
});
|
||||
});
|
||||
}
|
||||
|
||||
async function readCargoVersion() {
|
||||
let stdout;
|
||||
try {
|
||||
({ stdout } = await execFileAsync(
|
||||
"cargo",
|
||||
[
|
||||
"metadata",
|
||||
"--no-deps",
|
||||
"--format-version",
|
||||
"1",
|
||||
"--manifest-path",
|
||||
path.join(repositoryRoot, "Cargo.toml"),
|
||||
],
|
||||
{ cwd: repositoryRoot, maxBuffer: 1024 * 1024 },
|
||||
));
|
||||
} catch (error) {
|
||||
throw new Error(`Unable to read the root Cargo package version: ${error.message}`, {
|
||||
cause: error,
|
||||
});
|
||||
}
|
||||
|
||||
let metadata;
|
||||
try {
|
||||
metadata = JSON.parse(stdout);
|
||||
} catch (error) {
|
||||
throw new Error("cargo metadata returned invalid JSON", { cause: error });
|
||||
}
|
||||
|
||||
const rootPackage = metadata.packages?.find((packageMetadata) => packageMetadata.name === "mtp");
|
||||
if (!rootPackage || typeof rootPackage.version !== "string") {
|
||||
throw new Error("The root Cargo package named 'mtp' was not found");
|
||||
}
|
||||
|
||||
return rootPackage.version;
|
||||
}
|
||||
|
||||
function packageRelativePath(entry) {
|
||||
if (typeof entry !== "string" || entry.length === 0) {
|
||||
throw new Error("package.json files entries must be non-empty strings");
|
||||
}
|
||||
|
||||
const relativePath = entry.replace(/\/$/, "");
|
||||
if (
|
||||
!relativePath ||
|
||||
path.isAbsolute(relativePath) ||
|
||||
relativePath.split(/[\\/]/u).includes("..") ||
|
||||
relativePath.includes("*")
|
||||
) {
|
||||
throw new Error(`Unsupported package file entry: ${entry}`);
|
||||
}
|
||||
|
||||
return relativePath;
|
||||
}
|
||||
|
||||
async function copyPackageFiles(stageRoot, packageJson) {
|
||||
if (!Array.isArray(packageJson.files)) {
|
||||
throw new Error("package.json must declare a files array for Web releases");
|
||||
}
|
||||
|
||||
for (const entry of packageJson.files) {
|
||||
const relativePath = packageRelativePath(entry);
|
||||
const sourcePath = path.join(repositoryRoot, relativePath);
|
||||
const destinationPath = path.join(stageRoot, relativePath);
|
||||
|
||||
try {
|
||||
await access(sourcePath);
|
||||
} catch (error) {
|
||||
throw new Error(`Release file is missing: ${relativePath}`, { cause: error });
|
||||
}
|
||||
|
||||
await mkdir(path.dirname(destinationPath), { recursive: true });
|
||||
await cp(sourcePath, destinationPath, { recursive: true });
|
||||
}
|
||||
}
|
||||
|
||||
async function createRelease({ outputDir, packageJson, version }) {
|
||||
const stageRoot = await mkdtemp(path.join(os.tmpdir(), "mtp-web-release-"));
|
||||
const stagedPackageJson = {
|
||||
...packageJson,
|
||||
version,
|
||||
};
|
||||
|
||||
try {
|
||||
await writeFile(
|
||||
path.join(stageRoot, "package.json"),
|
||||
`${JSON.stringify(stagedPackageJson, null, 2)}\n`,
|
||||
);
|
||||
await copyPackageFiles(stageRoot, packageJson);
|
||||
|
||||
const stagedWasmPackagePath = path.join(stageRoot, "wasm", "pkg", "package.json");
|
||||
const stagedWasmPackageJson = await readJson(stagedWasmPackagePath);
|
||||
stagedWasmPackageJson.version = version;
|
||||
await writeFile(
|
||||
stagedWasmPackagePath,
|
||||
`${JSON.stringify(stagedWasmPackageJson, null, 2)}\n`,
|
||||
);
|
||||
|
||||
await mkdir(outputDir, { recursive: true });
|
||||
const archiveName = `${packageJson.name}-${version}.tgz`;
|
||||
const archivePath = path.join(outputDir, archiveName);
|
||||
await rm(archivePath, { force: true });
|
||||
|
||||
await run("npm", ["pack", "--pack-destination", outputDir], { cwd: stageRoot });
|
||||
|
||||
try {
|
||||
await access(archivePath);
|
||||
} catch (error) {
|
||||
throw new Error(`npm pack did not create ${archiveName}`, { cause: error });
|
||||
}
|
||||
|
||||
return archivePath;
|
||||
} finally {
|
||||
await rm(stageRoot, { recursive: true, force: true });
|
||||
}
|
||||
}
|
||||
|
||||
async function main() {
|
||||
const options = parseArguments(process.argv.slice(2));
|
||||
if (options.help) {
|
||||
console.log(usage());
|
||||
return;
|
||||
}
|
||||
|
||||
const packageJson = await readJson(packageJsonPath);
|
||||
if (packageJson.name !== "mtp") {
|
||||
throw new Error("package.json must describe the 'mtp' Web package");
|
||||
}
|
||||
|
||||
const version = await readCargoVersion();
|
||||
console.log(`Using Cargo package version ${version}`);
|
||||
|
||||
if (!options.skipBuild) {
|
||||
await run("pnpm", ["run", "clean"]);
|
||||
await run("pnpm", ["run", "build"]);
|
||||
}
|
||||
|
||||
const archivePath = await createRelease({
|
||||
outputDir: options.outputDir,
|
||||
packageJson,
|
||||
version,
|
||||
});
|
||||
console.log(`\nCreated ${path.relative(repositoryRoot, archivePath) || archivePath}`);
|
||||
}
|
||||
|
||||
main().catch((error) => {
|
||||
console.error(`\n${error.message}`);
|
||||
process.exitCode = 1;
|
||||
});
|
||||
6
crypto/Cargo.lock
generated
6
crypto/Cargo.lock
generated
|
|
@ -181,9 +181,9 @@ dependencies = [
|
|||
|
||||
[[package]]
|
||||
name = "chacha20"
|
||||
version = "0.10.2"
|
||||
version = "0.10.1"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "65c35e4b699c7e15ccbe7ee35c005e4fc0a278d22238a2857e6ce2dadeda1b06"
|
||||
checksum = "d524456ba66e72eb8b115ff89e01e497f8e6d11d78b70b1aa13c0fbd97540a81"
|
||||
dependencies = [
|
||||
"cfg-if",
|
||||
"cpufeatures 0.3.0",
|
||||
|
|
@ -853,7 +853,7 @@ version = "0.10.2"
|
|||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "c7f5fa3a058cd35567ef9bfa5e75732bee0f9e4c55fa90477bef2dfcdbc4be80"
|
||||
dependencies = [
|
||||
"chacha20 0.10.2",
|
||||
"chacha20 0.10.1",
|
||||
"getrandom 0.4.3",
|
||||
"rand_core 0.10.1",
|
||||
]
|
||||
|
|
|
|||
|
|
@ -23,7 +23,6 @@ rand = "0.10.2"
|
|||
getrandom = "0.4.3"
|
||||
mlkem-tls = { version = "0.2", optional = true }
|
||||
ml-dsa = { version = "0.1.1", optional = true }
|
||||
argon2 = { version = "0.5", optional = true }
|
||||
serde = { version = "1", optional = true, features = ["derive"] }
|
||||
rcgen = { version = "0.14", optional = true }
|
||||
time = { version = "0.3", optional = true }
|
||||
|
|
@ -44,4 +43,3 @@ hkdf = ["dep:hkdf", "dep:sha2"]
|
|||
sha2 = ["dep:sha2"]
|
||||
tls = ["dep:rcgen", "dep:time"]
|
||||
parallel = ["dep:tokio"]
|
||||
password-kdf = ["dep:argon2"]
|
||||
|
|
|
|||
|
|
@ -6,8 +6,6 @@ pub enum CryptoError {
|
|||
EncryptionFailed,
|
||||
#[error("decryption failed")]
|
||||
DecryptionFailed,
|
||||
#[error("decryption output exceeds the caller's allocation limit")]
|
||||
AllocationLimit,
|
||||
#[error("malformed encryption envelope")]
|
||||
MalformedEnvelope,
|
||||
#[error("no encryption recipients")]
|
||||
|
|
|
|||
|
|
@ -40,114 +40,6 @@ pub struct MultiEncryptedMessage {
|
|||
pub ciphertext: Vec<u8>,
|
||||
}
|
||||
|
||||
/// Borrowed view of a canonical encrypted envelope.
|
||||
///
|
||||
/// The codec uses this view while validating an attacker-controlled envelope
|
||||
/// so parsing it does not first create a complete temporary copy of every
|
||||
/// recipient entry and the ciphertext.
|
||||
#[derive(Debug, Clone, Copy)]
|
||||
pub struct MultiEncryptedMessageRef<'a> {
|
||||
encryption_type: EncryptionType,
|
||||
purpose: u8,
|
||||
bytes: &'a [u8],
|
||||
entries_start: usize,
|
||||
entry_len: usize,
|
||||
count: usize,
|
||||
ciphertext_start: usize,
|
||||
}
|
||||
|
||||
impl<'a> MultiEncryptedMessageRef<'a> {
|
||||
pub fn from_bytes(bytes: &'a [u8]) -> Result<Self, CryptoError> {
|
||||
if bytes.len() < 4 {
|
||||
return Err(CryptoError::MalformedEnvelope);
|
||||
}
|
||||
let encryption_type =
|
||||
EncryptionType::from_byte(bytes[0]).ok_or(CryptoError::UnknownAlgorithm)?;
|
||||
let purpose = bytes[1];
|
||||
let count = u16::from_be_bytes([bytes[2], bytes[3]]) as usize;
|
||||
if count == 0 || count > MAX_RECIPIENTS {
|
||||
return Err(CryptoError::MalformedEnvelope);
|
||||
}
|
||||
let entry_len = encryption_type
|
||||
.kem_ciphertext_len()
|
||||
.checked_add(encryption_type.wrapped_key_len())
|
||||
.ok_or(CryptoError::MalformedEnvelope)?;
|
||||
let entries_len = count
|
||||
.checked_mul(entry_len)
|
||||
.ok_or(CryptoError::MalformedEnvelope)?;
|
||||
let entries_start = 4usize;
|
||||
let ciphertext_start = entries_start
|
||||
.checked_add(entries_len)
|
||||
.ok_or(CryptoError::MalformedEnvelope)?;
|
||||
let ciphertext_len = bytes
|
||||
.len()
|
||||
.checked_sub(ciphertext_start)
|
||||
.ok_or(CryptoError::MalformedEnvelope)?;
|
||||
if ciphertext_len < encryption_type.minimum_ciphertext_len() {
|
||||
return Err(CryptoError::MalformedEnvelope);
|
||||
}
|
||||
Ok(Self {
|
||||
encryption_type,
|
||||
purpose,
|
||||
bytes,
|
||||
entries_start,
|
||||
entry_len,
|
||||
count,
|
||||
ciphertext_start,
|
||||
})
|
||||
}
|
||||
|
||||
pub const fn encryption_type(&self) -> EncryptionType {
|
||||
self.encryption_type
|
||||
}
|
||||
|
||||
pub const fn purpose(&self) -> u8 {
|
||||
self.purpose
|
||||
}
|
||||
|
||||
pub const fn recipient_count(&self) -> usize {
|
||||
self.count
|
||||
}
|
||||
|
||||
pub fn recipient(&self, index: usize) -> Option<(&'a [u8], &'a [u8])> {
|
||||
if index >= self.count {
|
||||
return None;
|
||||
}
|
||||
let offset = self
|
||||
.entries_start
|
||||
.checked_add(index.checked_mul(self.entry_len)?)?;
|
||||
let kem_len = self.encryption_type.kem_ciphertext_len();
|
||||
let kem_end = offset.checked_add(kem_len)?;
|
||||
let end = offset.checked_add(self.entry_len)?;
|
||||
Some((
|
||||
self.bytes.get(offset..kem_end)?,
|
||||
self.bytes.get(kem_end..end)?,
|
||||
))
|
||||
}
|
||||
|
||||
pub fn ciphertext(&self) -> &'a [u8] {
|
||||
&self.bytes[self.ciphertext_start..]
|
||||
}
|
||||
|
||||
pub fn to_owned(self) -> MultiEncryptedMessage {
|
||||
let recipients = (0..self.count)
|
||||
.filter_map(|index| {
|
||||
let (kem_ciphertext, encrypted_key) = self.recipient(index)?;
|
||||
Some(RecipientEntry {
|
||||
kem_ciphertext: kem_ciphertext.to_vec(),
|
||||
encrypted_key: encrypted_key.to_vec(),
|
||||
})
|
||||
})
|
||||
.collect();
|
||||
MultiEncryptedMessage {
|
||||
encryption_type: self.encryption_type,
|
||||
purpose: self.purpose,
|
||||
recipients,
|
||||
ciphertext: self.ciphertext().to_vec(),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl MultiEncryptedMessage {
|
||||
/// Serialize the envelope body without redundant per-recipient lengths.
|
||||
pub fn to_bytes(&self) -> Result<Vec<u8>, CryptoError> {
|
||||
|
|
@ -180,7 +72,50 @@ impl MultiEncryptedMessage {
|
|||
|
||||
/// Parse the canonical envelope body.
|
||||
pub fn from_bytes(bytes: &[u8]) -> Result<Self, CryptoError> {
|
||||
Ok(MultiEncryptedMessageRef::from_bytes(bytes)?.to_owned())
|
||||
if bytes.len() < 4 {
|
||||
return Err(CryptoError::MalformedEnvelope);
|
||||
}
|
||||
let encryption_type =
|
||||
EncryptionType::from_byte(bytes[0]).ok_or(CryptoError::UnknownAlgorithm)?;
|
||||
let purpose = bytes[1];
|
||||
let count = u16::from_be_bytes([bytes[2], bytes[3]]) as usize;
|
||||
if count == 0 || count > MAX_RECIPIENTS {
|
||||
return Err(CryptoError::MalformedEnvelope);
|
||||
}
|
||||
let entry_len = encryption_type.kem_ciphertext_len() + encryption_type.wrapped_key_len();
|
||||
let entries_len = count
|
||||
.checked_mul(entry_len)
|
||||
.ok_or(CryptoError::MalformedEnvelope)?;
|
||||
let start = 4usize;
|
||||
let end = start
|
||||
.checked_add(entries_len)
|
||||
.ok_or(CryptoError::MalformedEnvelope)?;
|
||||
let ciphertext_len = bytes
|
||||
.len()
|
||||
.checked_sub(end)
|
||||
.ok_or(CryptoError::MalformedEnvelope)?;
|
||||
if ciphertext_len < encryption_type.minimum_ciphertext_len() {
|
||||
return Err(CryptoError::MalformedEnvelope);
|
||||
}
|
||||
|
||||
let mut offset = start;
|
||||
let mut recipients = Vec::with_capacity(count);
|
||||
for _ in 0..count {
|
||||
let kem_end = offset + encryption_type.kem_ciphertext_len();
|
||||
let wrapped_end = kem_end + encryption_type.wrapped_key_len();
|
||||
recipients.push(RecipientEntry {
|
||||
kem_ciphertext: bytes[offset..kem_end].to_vec(),
|
||||
encrypted_key: bytes[kem_end..wrapped_end].to_vec(),
|
||||
});
|
||||
offset = wrapped_end;
|
||||
}
|
||||
|
||||
Ok(Self {
|
||||
encryption_type,
|
||||
purpose,
|
||||
recipients,
|
||||
ciphertext: bytes[offset..].to_vec(),
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
|
|
@ -262,51 +197,20 @@ pub fn decrypt_multi_for(
|
|||
purpose: u8,
|
||||
keyring: &Keyring,
|
||||
) -> Result<Vec<u8>, CryptoError> {
|
||||
decrypt_multi_for_parts(
|
||||
message.encryption_type,
|
||||
message.purpose,
|
||||
&message.recipients,
|
||||
&message.ciphertext,
|
||||
purpose,
|
||||
keyring,
|
||||
)
|
||||
}
|
||||
|
||||
/// Decrypt an envelope represented by borrowed recipient and ciphertext
|
||||
/// slices. This keeps protected-value opening from cloning an already-owned
|
||||
/// envelope solely to call the cryptographic primitive.
|
||||
#[cfg(all(feature = "mlkem-tls", feature = "hkdf"))]
|
||||
pub fn decrypt_multi_for_parts(
|
||||
encryption_type: EncryptionType,
|
||||
envelope_purpose: u8,
|
||||
recipients: &[RecipientEntry],
|
||||
ciphertext: &[u8],
|
||||
purpose: u8,
|
||||
keyring: &Keyring,
|
||||
) -> Result<Vec<u8>, CryptoError> {
|
||||
if recipients.is_empty()
|
||||
|| recipients.len() > MAX_RECIPIENTS
|
||||
|| envelope_purpose != purpose
|
||||
|| ciphertext.len() < encryption_type.minimum_ciphertext_len()
|
||||
|| recipients.iter().any(|recipient| {
|
||||
recipient.kem_ciphertext.len() != encryption_type.kem_ciphertext_len()
|
||||
|| recipient.encrypted_key.len() != encryption_type.wrapped_key_len()
|
||||
if message.recipients.is_empty()
|
||||
|| message.recipients.len() > MAX_RECIPIENTS
|
||||
|| message.purpose != purpose
|
||||
|| message.ciphertext.len() < message.encryption_type.minimum_ciphertext_len()
|
||||
|| message.recipients.iter().any(|recipient| {
|
||||
recipient.kem_ciphertext.len() != message.encryption_type.kem_ciphertext_len()
|
||||
|| recipient.encrypted_key.len() != message.encryption_type.wrapped_key_len()
|
||||
})
|
||||
{
|
||||
return Err(CryptoError::MalformedEnvelope);
|
||||
}
|
||||
|
||||
let count = u16::try_from(recipients.len()).map_err(|_| CryptoError::MalformedEnvelope)?;
|
||||
let mut payload_aad = Vec::new();
|
||||
payload_aad.extend_from_slice(ENCRYPT_DOMAIN);
|
||||
payload_aad.push(encryption_type.to_byte());
|
||||
payload_aad.push(envelope_purpose);
|
||||
payload_aad.extend_from_slice(&count.to_be_bytes());
|
||||
for entry in recipients {
|
||||
payload_aad.extend_from_slice(&entry.kem_ciphertext);
|
||||
payload_aad.extend_from_slice(&entry.encrypted_key);
|
||||
}
|
||||
for entry in recipients {
|
||||
let payload_aad = payload_aad(message)?;
|
||||
for entry in &message.recipients {
|
||||
let shared_secret =
|
||||
match HybridKem::decapsulate(&keyring.kem_secret_key, &entry.kem_ciphertext) {
|
||||
Ok(secret) => secret,
|
||||
|
|
@ -315,51 +219,30 @@ pub fn decrypt_multi_for_parts(
|
|||
let wrap_key = Zeroizing::new(derive_encryption_key(
|
||||
&shared_secret,
|
||||
KEY_WRAP_DOMAIN,
|
||||
&[encryption_type.to_byte(), purpose],
|
||||
&[message.encryption_type.to_byte(), purpose],
|
||||
)?);
|
||||
let aad = wrap_aad(encryption_type, purpose, &entry.kem_ciphertext);
|
||||
let cek = match open_with_key(encryption_type, *wrap_key, &entry.encrypted_key, &aad) {
|
||||
let aad = wrap_aad(message.encryption_type, purpose, &entry.kem_ciphertext);
|
||||
let cek = match open_with_key(
|
||||
message.encryption_type,
|
||||
*wrap_key,
|
||||
&entry.encrypted_key,
|
||||
&aad,
|
||||
) {
|
||||
Ok(key) => key,
|
||||
Err(_) => continue,
|
||||
};
|
||||
let cek: [u8; 32] = cek.try_into().map_err(|_| CryptoError::DecryptionFailed)?;
|
||||
return open_with_key(encryption_type, cek, ciphertext, &payload_aad);
|
||||
return open_with_key(
|
||||
message.encryption_type,
|
||||
cek,
|
||||
&message.ciphertext,
|
||||
&payload_aad,
|
||||
);
|
||||
}
|
||||
|
||||
Err(CryptoError::NoMatchingRecipient)
|
||||
}
|
||||
|
||||
/// Decrypt a canonical envelope only when its plaintext can fit inside the
|
||||
/// caller's allocation budget.
|
||||
///
|
||||
/// The AEAD implementation allocates its output buffer internally. Checking
|
||||
/// the ciphertext upper bound before entering that implementation makes the
|
||||
/// codec's reservation meaningful instead of merely checking the result
|
||||
/// after the allocation has already happened.
|
||||
#[cfg(all(feature = "mlkem-tls", feature = "hkdf"))]
|
||||
pub fn decrypt_multi_for_parts_with_limit(
|
||||
encryption_type: EncryptionType,
|
||||
envelope_purpose: u8,
|
||||
recipients: &[RecipientEntry],
|
||||
ciphertext: &[u8],
|
||||
purpose: u8,
|
||||
keyring: &Keyring,
|
||||
max_plaintext_len: usize,
|
||||
) -> Result<Vec<u8>, CryptoError> {
|
||||
if ciphertext.len() > max_plaintext_len {
|
||||
return Err(CryptoError::AllocationLimit);
|
||||
}
|
||||
|
||||
decrypt_multi_for_parts(
|
||||
encryption_type,
|
||||
envelope_purpose,
|
||||
recipients,
|
||||
ciphertext,
|
||||
purpose,
|
||||
keyring,
|
||||
)
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
|
@ -442,30 +325,4 @@ mod tests {
|
|||
Err(CryptoError::MalformedEnvelope)
|
||||
));
|
||||
}
|
||||
|
||||
#[cfg(all(feature = "mlkem-tls", feature = "hkdf", feature = "chacha20poly1305"))]
|
||||
#[test]
|
||||
fn bounded_decryption_rejects_before_plaintext_allocation() -> Result<(), CryptoError> {
|
||||
let recipient = Keyring::generate();
|
||||
let message = encrypt_multi_for(
|
||||
EncryptionType::MlKemChaCha20Poly1305,
|
||||
1,
|
||||
b"bounded plaintext",
|
||||
&[recipient.public_key_bundle()],
|
||||
)?;
|
||||
|
||||
assert!(matches!(
|
||||
decrypt_multi_for_parts_with_limit(
|
||||
message.encryption_type,
|
||||
message.purpose,
|
||||
&message.recipients,
|
||||
&message.ciphertext,
|
||||
message.purpose,
|
||||
&recipient,
|
||||
message.ciphertext.len() - 1,
|
||||
),
|
||||
Err(CryptoError::AllocationLimit)
|
||||
));
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -36,29 +36,3 @@ pub fn derive_encryption_key(
|
|||
out.copy_from_slice(&key);
|
||||
Ok(out)
|
||||
}
|
||||
|
||||
#[cfg(feature = "password-kdf")]
|
||||
pub fn derive_password_key(
|
||||
passphrase: &[u8],
|
||||
salt: &[u8],
|
||||
memory_kib: u32,
|
||||
iterations: u32,
|
||||
lanes: u32,
|
||||
) -> Result<[u8; 32], CryptoError> {
|
||||
if passphrase.is_empty()
|
||||
|| salt.len() < 16
|
||||
|| !(8 * 1024..=256 * 1024).contains(&memory_kib)
|
||||
|| !(1..=10).contains(&iterations)
|
||||
|| !(1..=8).contains(&lanes)
|
||||
{
|
||||
return Err(CryptoError::KdfError);
|
||||
}
|
||||
let params = argon2::Params::new(memory_kib, iterations, lanes, Some(32))
|
||||
.map_err(|_| CryptoError::KdfError)?;
|
||||
let argon = argon2::Argon2::new(argon2::Algorithm::Argon2id, argon2::Version::V0x13, params);
|
||||
let mut key = [0u8; 32];
|
||||
argon
|
||||
.hash_password_into(passphrase, salt, &mut key)
|
||||
.map_err(|_| CryptoError::KdfError)?;
|
||||
Ok(key)
|
||||
}
|
||||
|
|
|
|||
|
|
@ -340,9 +340,9 @@ impl Keyring {
|
|||
Ok(out)
|
||||
}
|
||||
|
||||
#[deprecated(note = "use try_to_bytes for the primary fallible serializer")]
|
||||
pub fn to_bytes(&self) -> Result<Zeroizing<Vec<u8>>, crate::error::CryptoError> {
|
||||
pub fn to_bytes(&self) -> Zeroizing<Vec<u8>> {
|
||||
self.try_to_bytes()
|
||||
.expect("key material length exceeds wire limit")
|
||||
}
|
||||
|
||||
pub fn from_bytes(bytes: &[u8]) -> Result<Self, crate::error::CryptoError> {
|
||||
|
|
@ -383,26 +383,16 @@ impl Keyring {
|
|||
})
|
||||
}
|
||||
|
||||
#[deprecated(note = "use try_to_hex for the primary fallible serializer")]
|
||||
pub fn to_hex(&self) -> Result<String, crate::error::CryptoError> {
|
||||
self.try_to_hex()
|
||||
}
|
||||
|
||||
pub fn try_to_hex(&self) -> Result<String, crate::error::CryptoError> {
|
||||
Ok(bytes_to_hex(&self.try_to_bytes()?))
|
||||
pub fn to_hex(&self) -> String {
|
||||
bytes_to_hex(&self.to_bytes())
|
||||
}
|
||||
|
||||
pub fn from_hex(s: &str) -> Result<Self, crate::error::CryptoError> {
|
||||
Self::from_bytes(&hex_to_bytes(s)?)
|
||||
}
|
||||
|
||||
#[deprecated(note = "use try_to_base64 for the primary fallible serializer")]
|
||||
pub fn to_base64(&self) -> Result<String, crate::error::CryptoError> {
|
||||
self.try_to_base64()
|
||||
}
|
||||
|
||||
pub fn try_to_base64(&self) -> Result<String, crate::error::CryptoError> {
|
||||
Ok(bytes_to_base64(&self.try_to_bytes()?))
|
||||
pub fn to_base64(&self) -> String {
|
||||
bytes_to_base64(&self.to_bytes())
|
||||
}
|
||||
|
||||
pub fn from_base64(s: &str) -> Result<Self, crate::error::CryptoError> {
|
||||
|
|
@ -517,9 +507,9 @@ impl PublicKeyBundle {
|
|||
Ok(out)
|
||||
}
|
||||
|
||||
#[deprecated(note = "use try_as_bytes for the primary fallible serializer")]
|
||||
pub fn as_bytes(&self) -> Result<Vec<u8>, crate::error::CryptoError> {
|
||||
pub fn as_bytes(&self) -> Vec<u8> {
|
||||
self.try_as_bytes()
|
||||
.expect("public key bundle field exceeds wire limit")
|
||||
}
|
||||
|
||||
/// Parse a complete suite-compatible public bundle.
|
||||
|
|
@ -592,13 +582,8 @@ impl PublicKeyBundle {
|
|||
Self::from_bytes(bytes)
|
||||
}
|
||||
|
||||
#[deprecated(note = "use try_to_base64 for the primary fallible serializer")]
|
||||
pub fn to_base64(&self) -> Result<String, crate::error::CryptoError> {
|
||||
self.try_to_base64()
|
||||
}
|
||||
|
||||
pub fn try_to_base64(&self) -> Result<String, crate::error::CryptoError> {
|
||||
Ok(bytes_to_base64(&self.try_as_bytes()?))
|
||||
pub fn to_base64(&self) -> String {
|
||||
bytes_to_base64(&self.as_bytes())
|
||||
}
|
||||
|
||||
pub fn from_base64(s: &str) -> Result<Self, crate::error::CryptoError> {
|
||||
|
|
@ -617,6 +602,12 @@ impl TryFrom<&[u8]> for PublicKeyBundle {
|
|||
}
|
||||
}
|
||||
|
||||
impl From<&PublicKeyBundle> for Vec<u8> {
|
||||
fn from(bundle: &PublicKeyBundle) -> Vec<u8> {
|
||||
bundle.as_bytes()
|
||||
}
|
||||
}
|
||||
|
||||
impl fmt::Debug for PublicKeyBundle {
|
||||
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
|
||||
f.debug_struct("PublicKeyBundle")
|
||||
|
|
@ -638,7 +629,7 @@ mod tests {
|
|||
let cl = SignaturePublicKey::new(vec![3u8; 32]);
|
||||
|
||||
let bundle = PublicKeyBundle::new(kem, pq, cl);
|
||||
let bytes = bundle.try_as_bytes()?;
|
||||
let bytes = bundle.as_bytes();
|
||||
let recovered = PublicKeyBundle::from_bytes_unvalidated(&bytes)?;
|
||||
|
||||
assert_eq!(
|
||||
|
|
@ -681,9 +672,9 @@ mod tests {
|
|||
SignaturePqPublicKey::new(vec![0xCDu8; 96]),
|
||||
SignaturePublicKey::new(vec![0xEFu8; 32]),
|
||||
);
|
||||
let bytes = bundle.try_as_bytes()?;
|
||||
let bytes: Vec<u8> = Vec::from(&bundle);
|
||||
let recovered = PublicKeyBundle::from_bytes_unvalidated(bytes.as_slice())?;
|
||||
assert_eq!(bundle.try_as_bytes()?, recovered.try_as_bytes()?);
|
||||
assert_eq!(bundle.as_bytes(), recovered.as_bytes());
|
||||
Ok(())
|
||||
}
|
||||
|
||||
|
|
@ -697,8 +688,8 @@ mod tests {
|
|||
SignaturePublicKey::new(vec![5u8; 32]),
|
||||
SignaturePrivateKey::new(vec![6u8; 32]),
|
||||
);
|
||||
let bytes = keyring.try_to_bytes()?;
|
||||
let recovered = Keyring::from_bytes(bytes.as_slice())?;
|
||||
let bytes = keyring.to_bytes();
|
||||
let recovered = Keyring::from_bytes(&bytes)?;
|
||||
assert_eq!(
|
||||
keyring.kem_public_key.as_bytes(),
|
||||
recovered.kem_public_key.as_bytes()
|
||||
|
|
@ -731,7 +722,7 @@ mod tests {
|
|||
}
|
||||
|
||||
#[test]
|
||||
fn canonical_key_parsers_reject_trailing_bytes() -> Result<(), Box<dyn std::error::Error>> {
|
||||
fn canonical_key_parsers_reject_trailing_bytes() {
|
||||
let keyring = Keyring::new(
|
||||
KemPublicKey::new(vec![1u8; 16]),
|
||||
KemPrivateKey::new(vec![2u8; 16]),
|
||||
|
|
@ -740,40 +731,25 @@ mod tests {
|
|||
SignaturePublicKey::new(vec![5u8; 16]),
|
||||
SignaturePrivateKey::new(vec![6u8; 16]),
|
||||
);
|
||||
let mut keyring_bytes = keyring.try_to_bytes()?.to_vec();
|
||||
let mut keyring_bytes = keyring.to_bytes().to_vec();
|
||||
keyring_bytes.push(0xAA);
|
||||
assert!(Keyring::from_bytes(&keyring_bytes).is_err());
|
||||
|
||||
let bundle = keyring.public_key_bundle();
|
||||
let mut bundle_bytes = bundle.try_as_bytes()?;
|
||||
let mut bundle_bytes = bundle.as_bytes();
|
||||
bundle_bytes.push(0xBB);
|
||||
assert!(PublicKeyBundle::from_bytes(&bundle_bytes).is_err());
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn public_key_bundle_try_as_bytes_rejects_fields_larger_than_wire_length() {
|
||||
let bundle = PublicKeyBundle::new(
|
||||
KemPublicKey::new(vec![0u8; 65_536]),
|
||||
SignaturePqPublicKey::new(Vec::new()),
|
||||
SignaturePublicKey::new(Vec::new()),
|
||||
);
|
||||
assert!(matches!(
|
||||
bundle.try_as_bytes(),
|
||||
Err(crate::error::CryptoError::InvalidKeyLength)
|
||||
));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn validated_bundle_rejects_partial_suite_keys() -> Result<(), Box<dyn std::error::Error>> {
|
||||
fn validated_bundle_rejects_partial_suite_keys() {
|
||||
let bundle = PublicKeyBundle::new(
|
||||
KemPublicKey::new(vec![1u8; 32]),
|
||||
SignaturePqPublicKey::new(vec![2u8; 64]),
|
||||
SignaturePublicKey::new(vec![3u8; 32]),
|
||||
);
|
||||
assert!(bundle.validate().is_err());
|
||||
assert!(PublicKeyBundle::from_bytes_validated(&bundle.try_as_bytes()?).is_err());
|
||||
Ok(())
|
||||
assert!(PublicKeyBundle::from_bytes_validated(&bundle.as_bytes()).is_err());
|
||||
}
|
||||
|
||||
#[test]
|
||||
|
|
@ -786,9 +762,9 @@ mod tests {
|
|||
SignaturePublicKey::new(vec![4u8; 16]),
|
||||
SignaturePrivateKey::new(vec![5u8; 16]),
|
||||
);
|
||||
let bytes = keyring.try_to_bytes()?;
|
||||
let bytes = keyring.to_bytes();
|
||||
let recovered = Keyring::try_from(bytes.as_slice())?;
|
||||
assert_eq!(keyring.try_to_bytes()?, recovered.try_to_bytes()?);
|
||||
assert_eq!(keyring.to_bytes(), recovered.to_bytes());
|
||||
Ok(())
|
||||
}
|
||||
|
||||
|
|
@ -812,9 +788,9 @@ mod tests {
|
|||
SignaturePublicKey::new(vec![5u8; 16]),
|
||||
SignaturePrivateKey::new(vec![6u8; 16]),
|
||||
);
|
||||
let hex = keyring.try_to_hex()?;
|
||||
let hex = keyring.to_hex();
|
||||
let recovered = Keyring::from_hex(&hex)?;
|
||||
assert_eq!(keyring.try_to_bytes()?, recovered.try_to_bytes()?);
|
||||
assert_eq!(keyring.to_bytes(), recovered.to_bytes());
|
||||
Ok(())
|
||||
}
|
||||
|
||||
|
|
@ -828,9 +804,9 @@ mod tests {
|
|||
SignaturePublicKey::new(vec![5u8; 16]),
|
||||
SignaturePrivateKey::new(vec![6u8; 16]),
|
||||
);
|
||||
let b64 = keyring.try_to_base64()?;
|
||||
let b64 = keyring.to_base64();
|
||||
let recovered = Keyring::from_base64(&b64)?;
|
||||
assert_eq!(keyring.try_to_bytes()?, recovered.try_to_bytes()?);
|
||||
assert_eq!(keyring.to_bytes(), recovered.to_bytes());
|
||||
Ok(())
|
||||
}
|
||||
|
||||
|
|
@ -841,9 +817,9 @@ mod tests {
|
|||
SignaturePqPublicKey::new(vec![2u8; 64]),
|
||||
SignaturePublicKey::new(vec![3u8; 32]),
|
||||
);
|
||||
let b64 = bundle.try_to_base64()?;
|
||||
let b64 = bundle.to_base64();
|
||||
let recovered = PublicKeyBundle::from_base64_unvalidated(&b64)?;
|
||||
assert_eq!(bundle.try_as_bytes()?, recovered.try_as_bytes()?);
|
||||
assert_eq!(bundle.as_bytes(), recovered.as_bytes());
|
||||
Ok(())
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -60,8 +60,6 @@ pub use sign::{DualSignature, DualSigner, sign_dual};
|
|||
#[cfg(feature = "sha2")]
|
||||
pub use hash::{Sha256Hasher, sha256, sha256_double};
|
||||
|
||||
#[cfg(feature = "password-kdf")]
|
||||
pub use kdf::derive_password_key;
|
||||
#[cfg(feature = "hkdf")]
|
||||
pub use kdf::{derive_encryption_key, hkdf_expand, hkdf_extract};
|
||||
|
||||
|
|
@ -85,9 +83,7 @@ pub use helper::{ENCRYPT_DOMAIN, KEY_WRAP_DOMAIN};
|
|||
|
||||
#[cfg(all(feature = "mlkem-tls", feature = "hkdf"))]
|
||||
pub use helper::{
|
||||
MAX_RECIPIENTS, MultiEncryptedMessage, MultiEncryptedMessageRef, RecipientEntry,
|
||||
decrypt_multi_for, decrypt_multi_for_parts, decrypt_multi_for_parts_with_limit,
|
||||
encrypt_multi_for,
|
||||
MAX_RECIPIENTS, MultiEncryptedMessage, RecipientEntry, decrypt_multi_for, encrypt_multi_for,
|
||||
};
|
||||
|
||||
/* ================================ TESTS ================================ */
|
||||
|
|
@ -315,9 +311,7 @@ mod tests {
|
|||
#[test]
|
||||
fn keyring_serialize_roundtrip() {
|
||||
let kr = Keyring::generate();
|
||||
let bytes = kr
|
||||
.try_to_bytes()
|
||||
.expect("keyring serialization should succeed");
|
||||
let bytes = kr.to_bytes();
|
||||
let loaded = Keyring::from_bytes(&bytes).expect("keyring roundtrip should succeed");
|
||||
assert_eq!(
|
||||
kr.kem_public_key.as_bytes(),
|
||||
|
|
@ -338,9 +332,7 @@ mod tests {
|
|||
fn public_key_bundle_serialize_roundtrip() {
|
||||
let kr = Keyring::generate();
|
||||
let bundle = kr.public_key_bundle();
|
||||
let bytes = bundle
|
||||
.try_as_bytes()
|
||||
.expect("bundle serialization should succeed");
|
||||
let bytes = bundle.as_bytes();
|
||||
let loaded = PublicKeyBundle::from_bytes(&bytes).expect("bundle roundtrip should succeed");
|
||||
assert_eq!(
|
||||
bundle.kem_public_key.as_bytes(),
|
||||
|
|
|
|||
16
deny.toml
16
deny.toml
|
|
@ -8,23 +8,9 @@ ignore = []
|
|||
|
||||
[bans]
|
||||
# Flag multiple versions of the same crate so duplicate trees are visible.
|
||||
multiple-versions = "deny"
|
||||
multiple-versions = "warn"
|
||||
wildcards = "deny"
|
||||
|
||||
# These versions are required by incompatible upstream dependency lines:
|
||||
# - pem/rcgen/wtransport still use base64 0.22.
|
||||
# - ring and wasm-bindgen still use getrandom 0.2.
|
||||
# - current displaydoc/serde/thiserror/tokio and wasm-bindgen trees span syn 2
|
||||
# and syn 3.
|
||||
# - ring still uses windows-sys 0.52 while the Tokio/QUIC tree uses 0.61.
|
||||
# Keep the duplicate-version policy strict for every other crate/version.
|
||||
skip = [
|
||||
{ name = "base64", version = "0.22.1" },
|
||||
{ name = "getrandom", version = "0.2.17" },
|
||||
{ name = "syn", version = "2.0.119" },
|
||||
{ name = "windows-sys", version = "0.52.0" },
|
||||
]
|
||||
|
||||
[licenses]
|
||||
# Allowlist of licenses acceptable for this project's dependencies.
|
||||
allow = [
|
||||
|
|
|
|||
|
|
@ -1,23 +1,19 @@
|
|||
# MTP Connections
|
||||
|
||||
Native clients and server-side hosts expose parallel connection handles after the
|
||||
opening handshake. The client creates its handle; the host receives one from
|
||||
`accept()`.
|
||||
Native clients and hosts share the same connection shape after the opening handshake. The client creates the connection; the host receives it from `accept()`.
|
||||
|
||||
| Member | Native client | Native host | Web host (`WebMTPConnection`) |
|
||||
| --- | --- | --- | --- |
|
||||
| `version` | Compiled client version accepted by the host | Version selected by the registry | Version selected by the registry |
|
||||
| `sender` | Sends `CommunicationValue` frames | Sends `CommunicationValue` frames | Sends `CommunicationValue` frames |
|
||||
| `receiver` | Underlying receiver; use `receive()` for application frames | Underlying receiver; use `receive()` for application frames | Underlying receiver; use `receive()` for application frames |
|
||||
| `receiver` | Receives application frames | Receives application frames | Receives application frames |
|
||||
| `description` | Optional label sent during setup | Optional label received from the client | Optional label received from the client |
|
||||
| `client_id` | Confirmed or assigned ID with `crypto` | Authenticated or guest client ID with `crypto` | Authenticated or guest client ID with `crypto` |
|
||||
| `auth_state` | Authentication result with `crypto` | Authentication result with `crypto` | Authentication result with `crypto` |
|
||||
| `path` | — | Native hosts use `/` | WebTransport CONNECT path (e.g. `/mtp`) |
|
||||
| `request_path` | / | / | WebTransport CONNECT path (e.g. `/mtp`) |
|
||||
| `remote_addr` | Server `SocketAddr` when available | Peer `SocketAddr` | Peer `SocketAddr` |
|
||||
|
||||
`WebMTPConnection`, returned by `MTPWebServer::accept()`, exposes the same
|
||||
server-side members as the native host connection. Its `path` contains the
|
||||
HTTP/3 path used for the WebTransport extended CONNECT request.
|
||||
`WebMTPConnection`, returned by `MTPWebServer::accept()`, exposes the same members as the native host connection plus `request_path`, which contains the HTTP/3 path used for the WebTransport extended CONNECT request.
|
||||
|
||||
Server-side MTP connections expose `remote_addr`, the peer address observed by
|
||||
QUIC. HTTP route handlers receive the peer address as `HttpRequest::remote_addr`.
|
||||
|
|
|
|||
|
|
@ -4,28 +4,23 @@ This file documents the connection and version negotiation logic.
|
|||
|
||||
## Registry
|
||||
|
||||
The `registry` module provides a multi-version `Registry` used by the host for
|
||||
version negotiation. Accessed through the `mtp` facade (requires the `host`
|
||||
feature). In this repository, `Registry::builtin()` is generated from
|
||||
[`example/type-maps.yaml`](../example/type-maps.yaml), which currently contains
|
||||
protocol version 3.0 only. Downstream projects can register additional versions
|
||||
in their own YAML configuration.
|
||||
The `registry` module provides a multi-version `Registry` used by the host for version negotiation. Accessed through the `mtp` facade (requires the `host` feature):
|
||||
|
||||
```rust
|
||||
use mtp::codec::{Version, registry::Registry};
|
||||
use mtp::codec::registry::Registry;
|
||||
|
||||
let registry = Registry::builtin(); // loads all TypeMaps from the build config
|
||||
let registry = Registry::builtin(); // loads all TypeMaps from config
|
||||
|
||||
// Check if a version is supported
|
||||
assert!(registry.supports(&Version(3, 0)));
|
||||
assert!(registry.supports(&Version(1, 0)));
|
||||
|
||||
// Find highest mutual version for a client
|
||||
let client_versions = &[Version(2, 0), Version(3, 0)];
|
||||
let client_versions = &[Version(0, 0), Version(1, 0)];
|
||||
let negotiated = registry.negotiate(client_versions);
|
||||
assert_eq!(negotiated, Some(Version(3, 0)));
|
||||
assert_eq!(negotiated, Some(Version(1, 0)));
|
||||
|
||||
// Look up a version's TypeMap
|
||||
let tm = registry.get(&Version(3, 0)).unwrap();
|
||||
let tm = registry.get(&Version(2, 0)).unwrap();
|
||||
```
|
||||
|
||||
The `Registry::builtin()` constructor uses the `TypeMap::vX_Y()` methods generated from the config.
|
||||
|
|
@ -59,9 +54,9 @@ let mut host = MTPHost::new(config).await?;
|
|||
while let Some(conn) = host.accept().await? {
|
||||
// conn.version is the negotiated version
|
||||
// conn.codec is a VersionedCodec scoped to that version
|
||||
// conn.sender / conn.receive() for application CommunicationValue I/O
|
||||
// conn.sender / conn.receiver for raw CommunicationValue I/O
|
||||
|
||||
let msg = conn.receive().await?;
|
||||
let msg = conn.receiver.receive().await?;
|
||||
}
|
||||
```
|
||||
|
||||
|
|
@ -99,31 +94,29 @@ The client's `PROTOCOL_VERSION` constant is set by `protocol_version` in `type-m
|
|||
## Version Negotiation Flow
|
||||
|
||||
```
|
||||
Client (v3.0) Host (v3.0)
|
||||
Client (v2.0) Host (v0.0, v1.0, v2.0)
|
||||
| |
|
||||
| QUIC connect |
|
||||
|----------------------->|
|
||||
| |
|
||||
| CommValue{ Ident. } |
|
||||
| Version -> "3.0" |
|
||||
| Version -> "2.0" |
|
||||
| Id -> 8765 |
|
||||
| (unsigned hello; auth |
|
||||
| challenge follows) |
|
||||
|----------------------->|
|
||||
| | registry.negotiate(&[Version(3,0)])
|
||||
| | -> Some(Version(3,0))
|
||||
| | registry.negotiate(&[Version(2,0)])
|
||||
| | -> Some(Version(2,0))
|
||||
| |
|
||||
| Response | selected v3.0 TypeMap
|
||||
| Response | selected v2.0 TypeMap
|
||||
|<-----------------------|
|
||||
| Status, version |
|
||||
| |
|
||||
| subsequent messages |
|
||||
| use v3.0 TypeMap |
|
||||
| use v2.0 TypeMap |
|
||||
```
|
||||
|
||||
If the client sends an unsupported version (for example, v2.0 to the current
|
||||
repository builtin host), `negotiate` returns `None` and the connection is
|
||||
closed.
|
||||
If the client sends an unsupported version (e.g. v3.0 when the host only knows up to v2.0), `negotiate` returns `None` and the connection is closed.
|
||||
|
||||
## Protocol Ping and Pong
|
||||
|
||||
|
|
|
|||
|
|
@ -12,8 +12,6 @@ MTP reports codec failures separately from connection and transport failures.
|
|||
| `ReservedCommunicationType` | An application attempted to use a reserved communication type ID. |
|
||||
| `InvalidEncoding` | Bytes do not match the MTP value or frame format. |
|
||||
| `TooManyEntries` | A serialized value or frame exceeds its representable size. |
|
||||
| `MissingTypeMap` | A versioned codec was asked to encode a value without a retained negotiated type map. |
|
||||
| `TypeMapMismatch` | A value was created with a different protocol type map from the codec or peer operation. |
|
||||
| `CryptoFailed` | Signing, verification, encryption, or decryption failed while encoding or decoding. |
|
||||
| `MissingField` | A required typed field is absent. |
|
||||
|
||||
|
|
@ -45,4 +43,4 @@ Native builds may expose additional variants wrapping QUIC and WebTransport erro
|
|||
|
||||
## Authentication Rejections
|
||||
|
||||
The host reports unsupported or missing protocol versions through `AcceptError`. Authentication failures return `AcceptError::AuthenticationFailed` after the host sends a rejected handshake response; a handshake that exceeds the configured limit returns `AcceptError::AuthenticationTimedOut`. The authentication flow and its signed fields are defined in [Security](SECURITY.md).
|
||||
The host reports unsupported or missing protocol versions through `AcceptError`. Authentication failures return `AcceptError::AuthenticationFailed` after the host sends a rejected handshake response. The authentication flow and its signed fields are defined in [Security](SECURITY.md).
|
||||
|
|
|
|||
|
|
@ -19,7 +19,7 @@ let request = CommunicationValue::new(CommunicationType::Ping).with_id(1);
|
|||
conn.sender.send(&request).await?;
|
||||
let response = conn.receive().await?;
|
||||
println!("received {:?}", response.id());
|
||||
conn.sender.close().await;
|
||||
conn.sender.close();
|
||||
```
|
||||
|
||||
## Configuration
|
||||
|
|
@ -142,7 +142,7 @@ let conn = MTPClient::auth_register(config, &keyring, &host_pk).await?;
|
|||
|
||||
// Save for next session
|
||||
let id = conn.client_id;
|
||||
let keyring_bytes = keyring.try_to_bytes()?;
|
||||
let keyring_bytes = keyring.to_bytes();
|
||||
```
|
||||
|
||||
When callers already know whether a saved client ID exists, the convenience helper uses `Some(id)` for login and `None` for registration:
|
||||
|
|
@ -175,7 +175,7 @@ pub struct Keyring {
|
|||
}
|
||||
```
|
||||
|
||||
- Serialise: `keyring.try_to_bytes()` -> `Result<Zeroizing<Vec<u8>>, CryptoError>`
|
||||
- Serialise: `keyring.to_bytes()` -> `Vec<u8>`
|
||||
- Deserialise: `Keyring::from_bytes(&bytes)` -> `Result<Keyring, CryptoError>`
|
||||
- Get public half: `keyring.public_key_bundle()` -> `PublicKeyBundle`
|
||||
|
||||
|
|
@ -230,7 +230,7 @@ let response = conn
|
|||
|
||||
Requests are routed by id through the connection's receive dispatcher. Frames with other ids remain available through `conn.receive()`.
|
||||
|
||||
Two send modes (configured via `mtp::client::Policy`):
|
||||
Two send modes (configured via `mtp::transport::Policy`):
|
||||
- `PersistentStream` (default): reuses one QUIC unidirectional stream
|
||||
- `SingleStreamPerMessage`: opens a new stream per message
|
||||
|
||||
|
|
@ -248,16 +248,12 @@ Inbound frames are queued internally. The `receive()` method returns the next av
|
|||
### Close
|
||||
|
||||
```rust
|
||||
conn.sender.close().await;
|
||||
conn.sender.close();
|
||||
// or
|
||||
conn.receiver.close();
|
||||
```
|
||||
|
||||
`Sender::close().await` gracefully finishes the active send stream, sends the
|
||||
MTP close frame, and waits for `force_close_delay` (default 300ms) before
|
||||
force-closing the QUIC connection if necessary. `Sender::close_immediate()` is
|
||||
the fire-and-forget variant. `Receiver::close()` closes the local receive
|
||||
handle without performing the sender's graceful close sequence.
|
||||
Sends a close frame and signals the peer. The `Sender::close()` spawns an async task that sends the frame, waits for `force_close_delay` (default 300ms), then force-closes the QUIC connection if the peer has not already done so.
|
||||
|
||||
### Pipes
|
||||
|
||||
|
|
@ -313,7 +309,7 @@ For `public_signer`, call `verify` and `into_verified` before calling `decrypt`;
|
|||
The `Policy` struct controls transport behaviour:
|
||||
|
||||
```rust
|
||||
use mtp::client::{Policy, SendMode};
|
||||
use mtp::transport::{Policy, SendMode};
|
||||
|
||||
let policy = Policy {
|
||||
send_mode: SendMode::PersistentStream,
|
||||
|
|
|
|||
|
|
@ -130,7 +130,7 @@ while let Some(connection) = server.accept().await? {
|
|||
```
|
||||
> `MTPWebServer::new` consumes a `HostConfig` (not an `MTPHost` instance). It creates its own QUIC endpoint and does not share a port with a running `MTPHost`.
|
||||
|
||||
`server.accept()` returns `Option<WebMTPConnection>` for each WebTransport session. Ordinary HTTP routes do not surface through `accept()` because the server dispatches them internally. `WebMTPConnection` retains the negotiated version, codec, `path`, remote address, description, sender, and receiver used by native MTP connections.
|
||||
`server.accept()` returns `Option<WebMTPConnection>` for each WebTransport session. Ordinary HTTP routes do not surface through `accept()` because the server dispatches them internally. `WebMTPConnection` retains the negotiated version, codec, request path, remote address, description, sender, and receiver used by native MTP connections.
|
||||
|
||||
## Deployment
|
||||
|
||||
|
|
@ -138,7 +138,7 @@ For direct browser access, leave `serve_tcp_https(true)` enabled. The server adv
|
|||
|
||||
When a reverse proxy or another process owns TCP, use `WebServerConfig::new().serve_tcp_https(false)`. This retains the UDP HTTP/3/WebTransport endpoint and its shared router without claiming the TCP port.
|
||||
|
||||
With port `0` and TCP enabled, construction binds TCP first and binds UDP to the selected TCP port, so `local_addr()` reports the common address. With TCP disabled, Quinn selects the UDP port as before. `shutdown().await` stops both accept loops, gracefully finishes active HTTP requests until `drain_timeout`, closes Quinn, and then aborts remaining work. `close().await` and dropping the server stop both listeners immediately.
|
||||
With port `0` and TCP enabled, construction binds TCP first and binds UDP to the selected TCP port, so `local_addr()` reports the common address. With TCP disabled, Quinn selects the UDP port as before. `shutdown()` stops both accept loops, gracefully finishes active HTTP requests until `drain_timeout`, closes Quinn, and then aborts remaining work. `close()` and dropping the server stop both listeners immediately.
|
||||
|
||||
### Authentication
|
||||
|
||||
|
|
@ -164,18 +164,12 @@ On success, the connection has `AuthState::Authenticated`, the assigned `client_
|
|||
|
||||
## Errors
|
||||
|
||||
`MTPWebServer::new` returns `CommunicationError` for certificate parsing,
|
||||
certificate loading, and bind failures. Authentication policy is evaluated when
|
||||
WebTransport sessions are accepted, not rejected during construction.
|
||||
`MTPWebServer::new` returns `CommunicationError` for certificate parsing, certificate loading, bind failures, and rejected authentication policy.
|
||||
`accept()` returns `AcceptError` for a missing or unsupported version, a receive failure, or a send failure during the WebTransport opening handshake. HTTP route failures are reported through `WebServerMetrics::error_occurred` when metrics are configured. See [Errors](ERRORS.md) for shared error variants.
|
||||
|
||||
`WebServerMetrics` has these callbacks:
|
||||
|
||||
```rust
|
||||
use std::time::Duration;
|
||||
|
||||
fn connection_accepted(&self)
|
||||
fn connection_closed(&self, duration: Duration, reason: &str)
|
||||
fn request_started(&self, path: &str)
|
||||
fn request_completed(&self, path: &str, status: u16, duration: Duration)
|
||||
fn error_occurred(&self, error: &WebServerError)
|
||||
|
|
|
|||
|
|
@ -83,18 +83,18 @@ network metadata, not an authenticated client identity.
|
|||
|
||||
## Version Negotiation
|
||||
|
||||
`accept()` uses the version-bearing opening frame and registry flow in [Connector](CONNECTOR.md). The host registry is built from the type maps in [`example/type-maps.yaml`](../example/type-maps.yaml) by `Registry::builtin()` in this repository; downstream builds can provide their own `MTP_TYPE_MAPS` configuration.
|
||||
`accept()` uses the version-bearing opening frame and registry flow in [Connector](CONNECTOR.md). The host registry is built from the type maps in `type-maps.yaml` by `Registry::builtin()`.
|
||||
|
||||
### Registry
|
||||
|
||||
```rust
|
||||
use mtp::codec::Version;
|
||||
use mtp::codec::registry::Registry;
|
||||
|
||||
let registry = host.registry();
|
||||
assert!(registry.supports(&Version(3, 0)));
|
||||
assert!(registry.supports(&Version(2, 0)));
|
||||
|
||||
let negotiated = registry.negotiate(&[Version(2, 0), Version(3, 0)]);
|
||||
// -> Some(Version(3, 0)) for this repository's builtin map
|
||||
let negotiated = registry.negotiate(&[Version(1, 0), Version(2, 0)]);
|
||||
// -> Some(Version(2, 0)) if both versions are registered
|
||||
```
|
||||
|
||||
## Authentication Flow
|
||||
|
|
@ -105,15 +105,13 @@ After a successful handshake, `MTPConnection` exposes `AuthState::Authenticated`
|
|||
|
||||
## Handling Messages
|
||||
|
||||
Use `conn.sender` and `conn.receive()` for bidirectional message exchange. The
|
||||
connection dispatcher owns the underlying receiver, especially when `pipes` is
|
||||
enabled:
|
||||
Use `conn.sender` and `conn.receiver` for bidirectional message exchange:
|
||||
|
||||
```rust
|
||||
while let Some(conn) = host.accept().await? {
|
||||
tokio::spawn(async move {
|
||||
loop {
|
||||
match conn.receive().await {
|
||||
match conn.receiver.receive().await {
|
||||
Ok(msg) => {
|
||||
let response = process_message(&msg, &conn);
|
||||
conn.sender.send(&response).await.ok();
|
||||
|
|
@ -214,7 +212,7 @@ let (kem_sk, kem_pk) = HybridKem::generate_keypair();
|
|||
let host_keyring = Keyring::new(kem_pk, kem_sk, sig_pq_pk, sig_pq_sk, sig_pk, sig_sk);
|
||||
|
||||
// Save to disk
|
||||
let bytes = host_keyring.try_to_bytes()?;
|
||||
let bytes = host_keyring.to_bytes();
|
||||
std::fs::write("host_keys.bin", bytes)?;
|
||||
```
|
||||
|
||||
|
|
|
|||
|
|
@ -39,4 +39,4 @@ Back up host keyrings and client keyrings as protected secrets. Test restoring a
|
|||
|
||||
### Graceful Shutdown
|
||||
|
||||
Stop accepting new connections, reject new work at the application layer, and allow active requests and pipe writers to finish. For `MTPWebServer`, call `shutdown().await`; its `drain_timeout` controls graceful TCP HTTP completion and the QUIC drain period before remaining connection tasks are terminated.
|
||||
Stop accepting new connections, reject new work at the application layer, and allow active requests and pipe writers to finish. For `MTPWebServer`, call `shutdown()`; its `drain_timeout` controls graceful TCP HTTP completion and the QUIC drain period before remaining connection tasks are terminated.
|
||||
|
|
|
|||
|
|
@ -167,9 +167,6 @@ if let Some(writer) = handle.wait().await? {
|
|||
```rust
|
||||
// Host
|
||||
use mtp_transport::{PipeSessionParameters, accept_pipe_session};
|
||||
use sha2::{Digest, Sha256};
|
||||
|
||||
// The streaming digest below requires `sha2` as a direct application dependency.
|
||||
|
||||
while let Ok(request) = conn.receive_pipe().await {
|
||||
if request.description() != "file-upload" {
|
||||
|
|
@ -185,7 +182,7 @@ while let Ok(request) = conn.receive_pipe().await {
|
|||
let mut reader = accept_pipe_session(
|
||||
reader.into_inner(), ¶ms, &own_keyring, &client_public_bundle,
|
||||
).await?;
|
||||
let mut hasher = Sha256::new();
|
||||
let mut hasher = sha2::Sha256::new();
|
||||
while let Some(chunk) = reader.read_record().await? {
|
||||
hasher.update(&chunk);
|
||||
process_chunk(&chunk).await?;
|
||||
|
|
|
|||
|
|
@ -55,28 +55,8 @@ Verified SDK results expose the authenticated `protectedVersion` and
|
|||
`finalRecipientId` alongside the application content.
|
||||
|
||||
Native applications use the same schema through `ProtectedMessageBuilder` and
|
||||
the replay-explicit `open_protected_checked` or `open_protected_without_replay`
|
||||
APIs; language bindings delegate envelope construction and opening to this
|
||||
codec boundary.
|
||||
|
||||
Message processing uses the replay-required native APIs
|
||||
`open_protected_checked` and `open_relay_metadata_checked` (or the equivalent
|
||||
browser client path). Stored-message or forensic tooling must opt into the
|
||||
explicit `*_without_replay` APIs. Native in-memory guards are bounded and
|
||||
configurable; durable guards must perform an atomic insert-if-absent on
|
||||
`(signer ID, MessageId)`.
|
||||
|
||||
Protected identifiers have semantic limits separate from the generic codec
|
||||
blob limit. The default maximum `MessageId` is 256 UTF-8 bytes and relay
|
||||
metadata is limited to 1 MiB of encoded metadata. Deployments can provide
|
||||
stricter limits through the receive policy. Limits are checked after
|
||||
authentication and before retained values enter replay or application state.
|
||||
|
||||
Transport-derived resource policies use a conservative decoder allocation
|
||||
factor of `4 * max_message_size`, in addition to the frame-size output limit.
|
||||
This factor accounts for owned wrapper, recipient, ciphertext, and decoded
|
||||
value copies; it is an implementation admission policy rather than a wire
|
||||
field.
|
||||
`open_protected`; language bindings delegate envelope construction and opening
|
||||
to this codec boundary.
|
||||
|
||||
## Authentication Flow
|
||||
|
||||
|
|
@ -97,24 +77,10 @@ Client Host
|
|||
|
||||
Login proof binds the protocol version, client ID, host challenge, and client nonce. Registration proof binds the protocol version, public key bundle, host challenge, and client nonce. The host challenge is generated per connection.
|
||||
|
||||
Authentication attempts pass through a deployment-configurable limiter before
|
||||
client lookup, key validation, challenge signing, or registration callbacks.
|
||||
The default host configuration uses a bounded in-memory window. Hosts may key
|
||||
limits by connection, peer identity, claimed client ID, or registration flow.
|
||||
When identity concealment is enabled, an unknown client ID follows a dummy
|
||||
challenge/proof path and receives the same generic authentication failure as a
|
||||
known client with an invalid proof; disabling concealment restores the legacy
|
||||
identity-specific response for deployments where IDs are public.
|
||||
|
||||
`ForceAuthentication` requires login or registration. `AllowAuthentication` accepts authenticated and unauthenticated clients. `Unauthenticated` rejects authentication attempts. The connection states are `Pending`, `Authenticated`, `Unauthenticated`, and `Failed`.
|
||||
|
||||
## Version Negotiation
|
||||
|
||||
The client sends one compiled-in protocol version. The host compares it with the versions in its registry and returns the selected version in the opening response. Subsequent frames use that version's type map. An unsupported version closes the connection with `AcceptError::UnsupportedVersion`.
|
||||
|
||||
The current self-delimiting `DataValue` codec and three-bit communication header
|
||||
are used by the repository's protocol 3.0 map. The checked-in builtin registry
|
||||
contains only 3.0, so its native clients and hosts do not provide legacy map
|
||||
fallbacks. Type-map versions are configuration-driven; a custom registry may
|
||||
register another version number, but its map must use the current codec format
|
||||
and is not a fallback for a different legacy wire format.
|
||||
The self-delimiting `DataValue` codec and the three-bit communication header begin at protocol version `3.0`. A peer offering an older codec version is rejected during version negotiation; the new decoder does not attempt legacy flag, ID, or crypto-container fallbacks.
|
||||
|
|
|
|||
|
|
@ -27,7 +27,7 @@ For rotation, publish the replacement certificate or key before changing the ser
|
|||
|
||||
### Development Certificates
|
||||
|
||||
The `tls` feature exposes `mtp_crypto::tls::generate_self_signed_cert`. It creates an ECDSA P-256 server certificate for the requested domain, `127.0.0.1`, and `::1`; the certificate is valid for 13 days. The lower-level `mtp_transport::HostConfig::self_signed` provides a transport-level self-signed setup without the crypto certificate helper.
|
||||
The `tls` feature exposes `mtp_crypto::tls::generate_self_signed_cert`. It creates an ECDSA P-256 server certificate for the requested domain, `127.0.0.1`, and `::1`; the certificate is valid for 13 days. `HostConfig::self_signed` provides a transport-level self-signed setup without the crypto certificate helper.
|
||||
|
||||
Self-signed certificates are for development. Production deployments should use a certificate trusted by the client or an explicitly pinned certificate.
|
||||
|
||||
|
|
@ -160,14 +160,6 @@ process boundaries. A guard should atomically record a new ID before
|
|||
dispatching application content. Transport frame IDs must not be used for
|
||||
this purpose.
|
||||
|
||||
Native message-processing boundaries require a replay guard through the
|
||||
checked opening APIs. Reopening stored or forensic frames without a guard is
|
||||
available only through an explicitly named `without_replay` API. The reference
|
||||
in-memory guard is bounded and FIFO-evicts old entries, so it is a duplicate
|
||||
suppression cache rather than durable replay protection. A durable deployment
|
||||
must use an atomic insert-if-absent operation keyed by `(signer ID, MessageId)`;
|
||||
a separate read followed by insert is race-prone.
|
||||
|
||||
`VerifiedRelayMetadata` is an authenticated capability rather than a caller
|
||||
constructed data transfer object. Rust fields are private and the browser
|
||||
implementation keeps authenticated state behind a branded class. Content
|
||||
|
|
@ -182,11 +174,9 @@ fallback.
|
|||
Verification takes a receiver-side `SignaturePolicy`/`ProtectionPolicy`.
|
||||
`AnySupported` is useful for compatibility at the low-level codec boundary,
|
||||
but protocol receivers should select `Ed25519` or `Dual`. The browser SDK uses
|
||||
an explicit `ed25519` default and permits an operation or client override. Its
|
||||
`MTPSecurityProfile` resolves protected-message sender/receiver suites,
|
||||
encrypted-pipe suites, and the authentication PQ requirement together;
|
||||
`any-supported` remains an explicit compatibility value. It never derives
|
||||
receive policy from the recipient keyring. Signature policy must be applied
|
||||
an explicit `ed25519` default and permits an operation or client override. It
|
||||
never derives receive policy from the recipient keyring. The sender's
|
||||
signature suite remains a separate choice. Signature policy must be applied
|
||||
independently to relay metadata, relay content, and pipe session establishment.
|
||||
|
||||
### Key history and rotation
|
||||
|
|
@ -212,7 +202,6 @@ The crate's feature groups are:
|
|||
| `serde` | Serialization support for key types |
|
||||
| `wasm` | `getrandom` support for WebAssembly |
|
||||
| `tls` | Development certificate generation |
|
||||
| `password-kdf` | Argon2id password derivation for protected keyring files |
|
||||
|
||||
The main types are `Keyring`, `PublicKeyBundle`, `EncryptionType`, `HybridKem`, `XChaCha20Poly1305` (with the legacy `ChaCha20Poly1305` alias), `Aes256Gcm`, `Ed25519Signer`, and `MlDsaSigner`. Hashing and KDF helpers include `sha256`, `sha256_double`, `hkdf_extract`, `hkdf_expand`, and `derive_encryption_key`. Handshake payload builders are in `mtp_crypto::auth`.
|
||||
|
||||
|
|
@ -272,43 +261,16 @@ to proceed with an invalid local decryption key.
|
|||
Applications remain responsible for storage at rest. The `files` feature writes passphrase-protected keyrings to `.mk` files and public bundles to `.mpkb` files. Protected `.mk` files store the Argon2id identifier, parameters, salt, and AEAD ciphertext; they do not derive their key with HKDF. On Unix, keyring files are created with owner-only `0600` permissions.
|
||||
Restrict those files to the owning account and protect backups. Browser applications should treat the configured credential storage as sensitive application data.
|
||||
|
||||
Key-material parsing is explicit in the SDK: use the hex, Base64, or byte
|
||||
helpers for encoded key material. Arbitrary strings are no longer treated as
|
||||
passphrases by the compatibility `secretKeyFromString` helper. Applications
|
||||
migrating data written by the old implicit-HKDF behavior can use the explicitly
|
||||
named, deprecated `legacySecretKeyFromStringV1` helper only for that migration;
|
||||
new data must not use it. Passwords must use the explicit Argon2id passphrase
|
||||
API with a stored per-record salt and versioned parameters. The SDK's
|
||||
`deriveKeyFromPassphrase` uses a worker when browser workers are available;
|
||||
the explicitly named `deriveKeyFromPassphraseSync` form is for workers and
|
||||
command-line migrations. HKDF helpers are for high-entropy key material and
|
||||
are not password-hardening functions.
|
||||
|
||||
## Resource Limits and Operational Controls
|
||||
|
||||
`Policy::default()` sets a 16 MiB application message limit and a 64 KiB handshake message limit. It also sets a 30 second read timeout, a 30 second maximum idle timeout, a receiver queue capacity of 1000, and a maximum of 128 concurrent stream tasks. Tune these values for the deployment and peer trust level.
|
||||
|
||||
The recursive codec applies additional defaults while parsing untrusted values:
|
||||
maximum nesting depth 64, 65,536 value nodes, 16 MiB per blob or envelope,
|
||||
64 encrypted recipients, and a 64 MiB cumulative decoder allocation budget.
|
||||
Decrypted values are parsed with the same limits. Transport derives the blob,
|
||||
allocation, and encoder output budgets from its admitted frame size rather than
|
||||
serializing an unrestricted recursive value first. The default transport
|
||||
allocation budget is four times the admitted frame size to cover conservative
|
||||
owned-copy and crypto-buffer accounting; deployments may choose another
|
||||
factor with `DecodeLimits::for_transport_message_size_with_allocation_factor`.
|
||||
and 64 encrypted recipients. Decrypted values are parsed with the same limits.
|
||||
|
||||
The host applies an authentication-attempt limiter before storage lookups,
|
||||
public-key validation, challenge signing, and registration callbacks. The
|
||||
default limiter is a bounded in-memory sliding window; configure a durable or
|
||||
distributed limiter when limits must coordinate across host instances. Unknown
|
||||
client IDs are sent through a fixed dummy challenge/proof path by default, so
|
||||
they receive a generic authentication failure instead of an enumeration hint.
|
||||
Deployments that intentionally publish client IDs can disable this concealment.
|
||||
|
||||
Keepalive Pong observation is bounded and accepts only the currently pending
|
||||
ping ID. Unsolicited Pongs are dropped before they can consume application
|
||||
receiver capacity.
|
||||
The host does not provide a general authentication-attempt rate limiter.
|
||||
Deploy authentication endpoints behind a rate-limiting proxy or add admission control through the host callbacks, including `GuestIdGenerator` where guest connections are permitted.
|
||||
|
||||
## Security Limitations
|
||||
|
||||
|
|
|
|||
|
|
@ -59,7 +59,7 @@ When `require_pq` is true, both Ed25519 and ML-DSA-65 keys and signatures must b
|
|||
|
||||
**Prevention:** Treat generated type maps as versioned build artifacts.
|
||||
|
||||
`CodecError::UnknownVersion` means the codec was created for a version absent from its registry. `UnknownCommunicationType` and `UnknownDataType` mean the selected `TypeMap` has no mapping for the value being encoded. `MissingTypeMap` means a versioned value lost its retained negotiated map; `TypeMapMismatch` means it was combined with a value or codec for another version. Select the negotiated type map and do not send an unmapped variant.
|
||||
`CodecError::UnknownVersion` means the codec was created for a version absent from its registry. `UnknownCommunicationType` and `UnknownDataType` mean the selected `TypeMap` has no mapping for the value being encoded. Select the negotiated type map and do not send an unmapped variant.
|
||||
|
||||
`ReservedCommunicationType` means application code attempted to use a reserved wire ID. Use generated communication types instead of assigning protocol IDs manually. `MissingField` means a required typed field was not present.
|
||||
|
||||
|
|
|
|||
|
|
@ -1,14 +1,6 @@
|
|||
# Type Map
|
||||
|
||||
This file documents the type-map and registry configuration used by MTP. The
|
||||
repository workspace uses [`example/type-maps.yaml`](../example/type-maps.yaml)
|
||||
through [`.cargo/config.toml`](../.cargo/config.toml); that map currently
|
||||
selects protocol version 3.0. The root [`example-type-maps.yaml`](../example-type-maps.yaml)
|
||||
is a separate illustrative multi-version configuration used by the manual WASM
|
||||
build script. Downstream applications should provide their own map.
|
||||
|
||||
The protocol version selects the generated codec/type-map build, while the
|
||||
type-map entries define the available application types and their IDs.
|
||||
This file documents the Type Map & Registry configuration used by the MTP protocol. It will assume you are working with the [example-type-maps.yaml](./../example-type-maps.yaml).
|
||||
|
||||
## Binary Frame Format
|
||||
|
||||
|
|
@ -93,15 +85,6 @@ The envelope length counts the bytes after the length field. A recipient entry i
|
|||
|
||||
Protection nesting directly represents both signer-visibility choices: `Encrypted(Signed(Container))` keeps signer metadata private, while `Signed(Encrypted(Container))` exposes it. A frame with no outer sender and an `Encrypted(Signed(Container))` payload uses sealed sender. Sealed sender adds no flag or distinct wire type.
|
||||
|
||||
### Container ordering and signatures
|
||||
|
||||
Container entries are ordered sequences in the current format. Insertion order
|
||||
is therefore semantic: two containers with the same field/value pairs in a
|
||||
different order have different serialized bytes and different signatures. The
|
||||
decoder rejects duplicate field IDs. Applications that need map semantics must
|
||||
canonicalize their own input before signing; a future canonical map encoding
|
||||
requires a protocol-format version and cannot be inferred by a receiver.
|
||||
|
||||
## TypeMap & Compile-Time Type Safety
|
||||
|
||||
A `TypeMap` maps Communication-Types and Data-Types to their wire IDs. Each protocol version has its own `TypeMap` because the same type name may use different wire IDs in different versions.
|
||||
|
|
@ -137,7 +120,7 @@ After editing the config and rebuilding, `CommunicationType` and `DataType` enum
|
|||
use mtp::type_map::{CommunicationType, DataType, TypeMap};
|
||||
|
||||
let tm = TypeMap::v3_0();
|
||||
let id = tm.data_id_enum(DataType::ExampleText).unwrap();
|
||||
let id = tm.data_id_enum(DataType::SomeType).unwrap();
|
||||
```
|
||||
|
||||
For native builds with the `registry` feature, the enums are a **union across
|
||||
|
|
@ -151,31 +134,30 @@ compiled by the Vite plugin.
|
|||
Encoding/decoding uses a `TypeMap` to resolve type names to wire IDs:
|
||||
|
||||
```rust
|
||||
use mtp::codec::{CommunicationType, CommunicationValue, DataType, DataValue};
|
||||
use mtp::codec::{encode, decode, DataValue};
|
||||
use mtp::type_map::TypeMap;
|
||||
|
||||
let tm = TypeMap::v3_0();
|
||||
let value = CommunicationValue::new_with_type_map(CommunicationType::Ping, &tm)
|
||||
.add_typed(DataType::Description, &tm, DataValue::Str("hello".into()));
|
||||
let tm = TypeMap::v2_0();
|
||||
let value = DataValue::Str("hello".into());
|
||||
|
||||
let bytes = value.to_bytes().unwrap();
|
||||
let decoded = CommunicationValue::from_bytes_with(&bytes, &tm).unwrap();
|
||||
let bytes = encode(&value, &tm).unwrap();
|
||||
let decoded = decode(&bytes, &tm).unwrap();
|
||||
```
|
||||
|
||||
```rust
|
||||
let tm_v3 = TypeMap::v3_0();
|
||||
assert!(tm_v3.data_id_enum(DataType::ExampleText).is_some());
|
||||
assert!(tm_v3.data_id_enum(DataType::SomeType).is_some());
|
||||
```
|
||||
|
||||
When communicating with a peer on another version, encode only variants that map in the negotiated version. If an incoming frame names a type absent from the selected map, reject it as a protocol or type-map compatibility error; do not reinterpret its wire ID using another version's map. The current repository map uses the self-delimiting codec format for protocol version `3.0`; a custom registry may register other version numbers, but those maps are not legacy wire-format fallbacks.
|
||||
When communicating with a peer on another version, encode only variants that map in the negotiated version. If an incoming frame names a type absent from the selected map, reject it as a protocol or type-map compatibility error; do not reinterpret its wire ID using another version's map. The self-delimiting codec begins at protocol version `3.0`; older versions are not codec fallbacks.
|
||||
|
||||
### Forward/Backward Compatibility Between Versions
|
||||
|
||||
Because enums are a union of all types across versions, a variant might exist that has no wire mapping in the *negotiated* version:
|
||||
|
||||
```
|
||||
v3.0 client sends DataType::ExampleText → host encodes with v3.0 TypeMap → wire ID 43
|
||||
v3.0 host receives a version absent from the registry → version negotiation error
|
||||
v3.0 client sends DataType::SomeType → host encodes with v3.0 TypeMap → wire ID 32
|
||||
v3.0 host receives an unsupported pre-v3.0 peer → version negotiation error
|
||||
```
|
||||
|
||||
Encoding a frame with an unmapped communication or data type returns `CodecError::UnknownCommunicationType` or `CodecError::UnknownDataType`. Select a mapped variant from the compiled-in version before sending it.
|
||||
|
|
@ -193,33 +175,17 @@ mtp = { path = "..", features = ["host"] }
|
|||
|
||||
```rust
|
||||
use mtp::codec::registry::{Registry, VersionedCodec};
|
||||
use mtp::codec::{CommunicationType, CommunicationValue, DataValue};
|
||||
use mtp_type_map::Version;
|
||||
|
||||
let registry = Registry::builtin();
|
||||
let codec = VersionedCodec::for_version(registry, Version(3, 0)).unwrap();
|
||||
let value = CommunicationValue::new_with_type_map(
|
||||
CommunicationType::Ping,
|
||||
codec.type_map(),
|
||||
).with_payload(DataValue::Null);
|
||||
let codec = VersionedCodec::new(registry);
|
||||
|
||||
// The value must retain the negotiated map used to construct it.
|
||||
let bytes = codec.encode(&value).unwrap();
|
||||
// Encode with a specific version
|
||||
let bytes = codec.encode(&value, Version(3, 0)).unwrap();
|
||||
|
||||
let decoded = codec.decode(&bytes).unwrap();
|
||||
|
||||
// A clear value can be migrated explicitly when the application has chosen
|
||||
// that behavior. Protected values are not silently remapped.
|
||||
let migrated = codec.encode_migrating(&value).unwrap();
|
||||
// Decode with a specific version
|
||||
let decoded = codec.decode(&bytes, Version(3, 0)).unwrap();
|
||||
```
|
||||
|
||||
`VersionedCodec::encode` compares the retained map identity (its protocol
|
||||
version) and returns `CodecError::MissingTypeMap` or
|
||||
`CodecError::TypeMapMismatch` on failure. `reply_to` retains the request's
|
||||
map, while `try_merge` rejects frames from different maps before copying any
|
||||
fields. The deprecated `merge` method records the error for compatibility; new
|
||||
code should migrate to `try_merge` and handle the result.
|
||||
|
||||
## Customizing Type Maps in Downstream Projects
|
||||
|
||||
External projects must provide their own type map configuration. Browser projects use the Vite plugin from [Defining Type Maps](#defining-type-maps) and do not need to publish, fork, or copy a generated WASM package.
|
||||
|
|
|
|||
|
|
@ -104,9 +104,6 @@ if (!MTPClient.isSupported()) {
|
|||
| `requestTimeoutMs` | 30 seconds | Default `request()` timeout. |
|
||||
| `pings` | `false` | Protocol pings, or an object with `intervalMs`. |
|
||||
| `logger` | No-op | Receives SDK state and error events. |
|
||||
| `schemas` | None | Client-wide request and response schema registry. |
|
||||
| `throwProtocolErrors` | `false` | Reject requests whose correlated response is an `Error*` frame. |
|
||||
| `onValidationError` | No-op | Receives subscription validation failures. |
|
||||
| `sessionStorage` | In-memory | E2EE session state storage. |
|
||||
| `encryptedSecretProvider` | In-memory | Independent caller-managed encrypted secret storage. |
|
||||
| `defaultSignatureVerificationPolicy` | `"ed25519"` | Receiver policy for protected signatures. |
|
||||
|
|
@ -474,60 +471,6 @@ const unsubscribe = client.subscribe("SomeType", (message) => {
|
|||
unsubscribe();
|
||||
```
|
||||
|
||||
### Zod request and response schemas
|
||||
|
||||
Applications can provide their request and response schemas once when creating
|
||||
the client. MTP uses `parseAsync`, so synchronous schemas, async refinements,
|
||||
defaults, coercions, and transforms all work. MTP has no runtime dependency on
|
||||
Zod; the application supplies its preferred Zod version.
|
||||
|
||||
```typescript
|
||||
import { z } from "zod";
|
||||
import { MTPClient, MTPValidationError } from "mtp";
|
||||
|
||||
const schemas = {
|
||||
GetUser: {
|
||||
request: z.object({ UserId: z.number().int().positive() }),
|
||||
response: z.object({
|
||||
UserId: z.number().int().positive(),
|
||||
Display: z.string(),
|
||||
}),
|
||||
},
|
||||
};
|
||||
|
||||
const client = await MTPClient.create({
|
||||
url,
|
||||
schemas,
|
||||
throwProtocolErrors: true,
|
||||
onValidationError(error) {
|
||||
console.error(error.messageType, error.cause);
|
||||
},
|
||||
});
|
||||
|
||||
const response = await client.request("GetUser", { UserId: 42 });
|
||||
console.log(response.data.Display);
|
||||
```
|
||||
|
||||
Request schemas run before frame encoding and transmission. Their transformed
|
||||
output is sent. Response schemas run after request correlation, and their
|
||||
transformed output replaces `frame.data`; `frame.raw`, when present, remains the
|
||||
original wire frame. Invalid requests and responses reject with
|
||||
`MTPValidationError`. Invalid subscription messages do not reach the handler
|
||||
and are reported through `onValidationError`.
|
||||
|
||||
`throwProtocolErrors: true` converts correlated `Error*` frames into
|
||||
`MTPProtocolError`. It defaults to `false` for compatibility.
|
||||
|
||||
`MTPProxyConnection` applies the same schema registry to another TypeScript
|
||||
request/subscription transport, such as a Tauri command and event proxy:
|
||||
|
||||
```typescript
|
||||
const connection = new MTPProxyConnection(adapter, {
|
||||
schemas,
|
||||
throwProtocolErrors: true,
|
||||
});
|
||||
```
|
||||
|
||||
Protocol ping behavior is defined in [Protocol Reference](PROTOCOL-REFERENCE.md#protocol-keepalive). The SDK configuration is:
|
||||
|
||||
```typescript
|
||||
|
|
@ -659,19 +602,8 @@ The SDK logger receives parsed events:
|
|||
|
||||
```typescript
|
||||
type MTPLogEvent =
|
||||
| {
|
||||
hint: "info" | "warning";
|
||||
type: string;
|
||||
data: unknown;
|
||||
direction?: "send" | "recv";
|
||||
}
|
||||
| {
|
||||
hint: "error";
|
||||
type: string | "error";
|
||||
error: string;
|
||||
data?: unknown;
|
||||
direction?: "send" | "recv";
|
||||
};
|
||||
| { hint: "info" | "warning"; type: string; data: unknown }
|
||||
| { hint: "error"; type: string | "error"; error: string };
|
||||
```
|
||||
|
||||
Incoming non-error frames and sent frames are logged as `info`. Error frames and transport errors are logged as `error`.
|
||||
|
|
|
|||
10
example/Cargo.lock
generated
10
example/Cargo.lock
generated
|
|
@ -225,9 +225,9 @@ dependencies = [
|
|||
|
||||
[[package]]
|
||||
name = "chacha20"
|
||||
version = "0.10.2"
|
||||
version = "0.10.1"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "65c35e4b699c7e15ccbe7ee35c005e4fc0a278d22238a2857e6ce2dadeda1b06"
|
||||
checksum = "d524456ba66e72eb8b115ff89e01e497f8e6d11d78b70b1aa13c0fbd97540a81"
|
||||
dependencies = [
|
||||
"cfg-if",
|
||||
"cpufeatures 0.3.0",
|
||||
|
|
@ -1305,6 +1305,7 @@ name = "mtp-common"
|
|||
version = "0.3.0"
|
||||
dependencies = [
|
||||
"quinn",
|
||||
"rustls",
|
||||
"thiserror 2.0.20",
|
||||
"wtransport",
|
||||
]
|
||||
|
|
@ -1313,7 +1314,6 @@ dependencies = [
|
|||
name = "mtp-crypto"
|
||||
version = "0.3.0"
|
||||
dependencies = [
|
||||
"argon2",
|
||||
"base64 0.22.1",
|
||||
"chacha20poly1305",
|
||||
"ed25519-dalek",
|
||||
|
|
@ -1337,6 +1337,7 @@ dependencies = [
|
|||
name = "mtp-files"
|
||||
version = "0.3.0"
|
||||
dependencies = [
|
||||
"argon2",
|
||||
"mtp-crypto",
|
||||
"rand",
|
||||
"thiserror 2.0.20",
|
||||
|
|
@ -1352,7 +1353,6 @@ dependencies = [
|
|||
"mtp-crypto",
|
||||
"mtp-transport",
|
||||
"rand",
|
||||
"thiserror 2.0.20",
|
||||
"tokio",
|
||||
"tracing",
|
||||
"wtransport",
|
||||
|
|
@ -1701,7 +1701,7 @@ version = "0.10.2"
|
|||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "c7f5fa3a058cd35567ef9bfa5e75732bee0f9e4c55fa90477bef2dfcdbc4be80"
|
||||
dependencies = [
|
||||
"chacha20 0.10.2",
|
||||
"chacha20 0.10.1",
|
||||
"getrandom 0.4.3",
|
||||
"rand_core 0.10.1",
|
||||
]
|
||||
|
|
|
|||
|
|
@ -62,7 +62,7 @@ pub fn build_demo_message(
|
|||
)
|
||||
.add_typed_default(
|
||||
DataType::Timestamp,
|
||||
DataValue::UnsignedNumber(timestamp),
|
||||
DataValue::UnsignedNumber(timestamp as u128),
|
||||
)
|
||||
.add_typed_default(DataType::Data, DataValue::Str("Hello, MTP!".into()))
|
||||
.add_typed_default(DataType::Flags, DataValue::BoolTrue)
|
||||
|
|
|
|||
|
|
@ -79,7 +79,6 @@ pub struct ClientMetrics {
|
|||
}
|
||||
|
||||
impl ClientMetrics {
|
||||
#[cfg(test)]
|
||||
pub fn new() -> Self {
|
||||
Self {
|
||||
sessions: Vec::new(),
|
||||
|
|
|
|||
|
|
@ -3,9 +3,8 @@ use std::time::{Duration, Instant};
|
|||
use mtp::client::MTPConnection;
|
||||
use mtp::codec::{
|
||||
CommunicationType, DataType, DataValue, ProtectionPolicy, ProtectedMessageBuilder,
|
||||
ProtectionPurpose, RelayOpenOptions, SealedRelayBuilder, SignaturePolicy, TypeMap,
|
||||
open_relay_content_with_limits_without_replay,
|
||||
open_relay_metadata_without_replay,
|
||||
ProtectionPurpose, SealedRelayBuilder, SignaturePolicy, TypeMap, open_relay_content,
|
||||
open_relay_metadata,
|
||||
};
|
||||
use mtp::common::unix_time_millis;
|
||||
use mtp::crypto::{Ed25519Signer, Keyring, PublicKeyBundle};
|
||||
|
|
@ -160,22 +159,22 @@ pub async fn send_sealed_relay(
|
|||
return Err("relay forwarding changed the sealed-sender boundary".into());
|
||||
}
|
||||
|
||||
let metadata = open_relay_metadata_without_replay(
|
||||
let metadata = open_relay_metadata(
|
||||
&forwarded,
|
||||
&final_recipient_keyring,
|
||||
signer_id,
|
||||
&signer_keyring.public_key_bundle(),
|
||||
RelayOpenOptions::new(RELAY_SIGNATURE_POLICY),
|
||||
RELAY_SIGNATURE_POLICY,
|
||||
)?;
|
||||
let application_metadata = metadata
|
||||
.metadata()
|
||||
.ok_or("forwarded relay metadata was missing")?;
|
||||
let content = open_relay_content_with_limits_without_replay(
|
||||
let content = open_relay_content(
|
||||
&metadata,
|
||||
&[&final_recipient_keyring],
|
||||
&[signer_keyring.public_key_bundle()],
|
||||
Some(FINAL_RECIPIENT_ID),
|
||||
RelayOpenOptions::new(RELAY_SIGNATURE_POLICY),
|
||||
&final_recipient_keyring,
|
||||
&signer_keyring.public_key_bundle(),
|
||||
FINAL_RECIPIENT_ID,
|
||||
RELAY_SIGNATURE_POLICY,
|
||||
)?;
|
||||
if content.message_type != "ProtectedMessage" {
|
||||
return Err(format!("unexpected relay message type: {}", content.message_type).into());
|
||||
|
|
|
|||
|
|
@ -17,22 +17,14 @@ fn main() -> Result<(), files::FileError> {
|
|||
/* Read both back to confirm the files round-trip through the on-disk format. */
|
||||
let loaded_keyring = load_keyring_raw(&keyring_path)?;
|
||||
let loaded_bundle = load_public_key_bundle(&bundle_path)?;
|
||||
assert_eq!(keyring.try_to_bytes()?, loaded_keyring.try_to_bytes()?);
|
||||
let bundle_bytes = keyring.public_key_bundle().try_as_bytes()?;
|
||||
let loaded_bundle_bytes = loaded_bundle.try_as_bytes()?;
|
||||
assert_eq!(keyring.to_bytes(), loaded_keyring.to_bytes());
|
||||
assert_eq!(
|
||||
bundle_bytes,
|
||||
loaded_bundle_bytes
|
||||
);
|
||||
println!(
|
||||
"\nPrivateKeyRing (base64):\n{}",
|
||||
keyring.try_to_base64()?
|
||||
keyring.public_key_bundle().as_bytes(),
|
||||
loaded_bundle.as_bytes()
|
||||
);
|
||||
println!("\nPrivateKeyRing (base64):\n{}", keyring.to_base64());
|
||||
|
||||
println!(
|
||||
"\nPublicKeyBundle (base64):\n{}",
|
||||
loaded_bundle.try_to_base64()?
|
||||
);
|
||||
println!("\nPublicKeyBundle (base64):\n{}", loaded_bundle.to_base64());
|
||||
|
||||
println!("Wrote keyring -> {}", keyring_path.display());
|
||||
println!("Wrote bundle -> {}", bundle_path.display());
|
||||
|
|
|
|||
|
|
@ -2,11 +2,8 @@ use std::collections::HashMap;
|
|||
|
||||
use mtp::codec::{
|
||||
CommunicationType, CommunicationValue, DataType, DataTypeId, DataValue, InMemoryReplayGuard,
|
||||
ProtectedOpenOptions, ProtectionPolicy, ProtectionPurpose, RelayOpenOptions, SignaturePolicy,
|
||||
TypeMap,
|
||||
forward_relay_frame, open_protected_with_checked,
|
||||
open_relay_content_with_limits_without_replay,
|
||||
open_relay_metadata_with_checked,
|
||||
ProtectedOpenOptions, ProtectionPolicy, ProtectionPurpose, SignaturePolicy, TypeMap,
|
||||
forward_relay_frame, open_protected_with, open_relay_content, open_relay_metadata_with,
|
||||
};
|
||||
use mtp::crypto::{Keyring, PublicKeyBundle};
|
||||
|
||||
|
|
@ -47,7 +44,7 @@ fn pong(tm: &TypeMap, data: impl Into<String>) -> Result<CommunicationValue, Str
|
|||
CommunicationValue::from_comm(CommunicationType::Pong, tm)
|
||||
.add_data(desc_id, DataValue::Str("MTP example response".into()))
|
||||
.map_err(|e| e.to_string())?
|
||||
.add_data(ts_id, DataValue::UnsignedNumber(now))
|
||||
.add_data(ts_id, DataValue::UnsignedNumber(now as u128))
|
||||
.map_err(|e| e.to_string())?
|
||||
.add_data(data_id, DataValue::Str(data.into()))
|
||||
.map_err(|e| e.to_string())
|
||||
|
|
@ -68,7 +65,7 @@ fn process_direct_protected(
|
|||
));
|
||||
}
|
||||
|
||||
let opened = open_protected_with_checked(
|
||||
let opened = open_protected_with(
|
||||
msg,
|
||||
std::slice::from_ref(&host_keyring),
|
||||
None,
|
||||
|
|
@ -79,7 +76,7 @@ fn process_direct_protected(
|
|||
ProtectionPurpose::from(DIRECT_ENCRYPTION_PURPOSE),
|
||||
SIGNATURE_POLICY,
|
||||
),
|
||||
accepted_messages,
|
||||
Some(accepted_messages),
|
||||
)
|
||||
.map_err(|e| format!("direct protected message could not be authenticated: {e}"))?;
|
||||
let signer_id = opened.signer_id;
|
||||
|
|
@ -136,15 +133,15 @@ fn process_sealed_relay(
|
|||
));
|
||||
}
|
||||
|
||||
let metadata = open_relay_metadata_with_checked(
|
||||
let metadata = open_relay_metadata_with(
|
||||
msg,
|
||||
std::slice::from_ref(&host_keyring),
|
||||
None,
|
||||
|signer_id| {
|
||||
resolve_signer_key(signer_id, registered_clients).map(|key| vec![key])
|
||||
},
|
||||
RelayOpenOptions::new(SIGNATURE_POLICY),
|
||||
accepted_messages,
|
||||
SIGNATURE_POLICY,
|
||||
Some(accepted_messages),
|
||||
)
|
||||
.map_err(|e| format!("metadata relay could not authenticate metadata: {e}"))?;
|
||||
println!(
|
||||
|
|
@ -163,13 +160,13 @@ fn process_sealed_relay(
|
|||
.len()
|
||||
);
|
||||
|
||||
let content_result = open_relay_content_with_limits_without_replay(
|
||||
let content_result = open_relay_content(
|
||||
&metadata,
|
||||
&[host_keyring],
|
||||
&[resolve_signer_key(metadata.signer_id(), registered_clients)
|
||||
.ok_or("metadata signer key disappeared")?],
|
||||
Some(FINAL_RECIPIENT_ID),
|
||||
RelayOpenOptions::new(SIGNATURE_POLICY),
|
||||
host_keyring,
|
||||
&resolve_signer_key(metadata.signer_id(), registered_clients)
|
||||
.ok_or("metadata signer key disappeared")?,
|
||||
FINAL_RECIPIENT_ID,
|
||||
SIGNATURE_POLICY,
|
||||
);
|
||||
if content_result.is_ok() {
|
||||
return Err("metadata relay unexpectedly decrypted final-recipient content".into());
|
||||
|
|
@ -310,22 +307,12 @@ pub fn process_and_respond(
|
|||
let signer_id = sig.as_signed().map(|signed| signed.signer_id);
|
||||
if let Some(signer_id) = signer_id
|
||||
&& sig
|
||||
.verify_with_policy(
|
||||
signer_id,
|
||||
pk_bundle,
|
||||
mtp::codec::ProtectionPurpose::from(2),
|
||||
SIGNATURE_POLICY,
|
||||
)
|
||||
.verify(signer_id, pk_bundle, mtp::codec::ProtectionPurpose::from(2))
|
||||
.is_ok()
|
||||
{
|
||||
let dv = sig
|
||||
.clone()
|
||||
.into_verified_with_policy(
|
||||
signer_id,
|
||||
pk_bundle,
|
||||
mtp::codec::ProtectionPurpose::from(2),
|
||||
SIGNATURE_POLICY,
|
||||
)
|
||||
.into_verified(signer_id, pk_bundle, mtp::codec::ProtectionPurpose::from(2))
|
||||
.ok();
|
||||
if let Some(entries) = dv.and_then(|value| value.as_container()) {
|
||||
println!(" Verified SignedPayload: {:?}", entries);
|
||||
|
|
@ -346,22 +333,16 @@ pub fn process_and_respond(
|
|||
if let Ok(opened) = secure.decrypt(host_keyring, mtp::codec::ProtectionPurpose::from(4))
|
||||
&& let Some(signed) = opened.as_signed()
|
||||
&& opened
|
||||
.verify_with_policy(
|
||||
.verify(
|
||||
signed.signer_id,
|
||||
pk_bundle,
|
||||
mtp::codec::ProtectionPurpose::from(3),
|
||||
SIGNATURE_POLICY,
|
||||
)
|
||||
.is_ok()
|
||||
{
|
||||
let signer_id = signed.signer_id;
|
||||
let dv = opened
|
||||
.into_verified_with_policy(
|
||||
signer_id,
|
||||
pk_bundle,
|
||||
mtp::codec::ProtectionPurpose::from(3),
|
||||
SIGNATURE_POLICY,
|
||||
)
|
||||
.into_verified(signer_id, pk_bundle, mtp::codec::ProtectionPurpose::from(3))
|
||||
.ok();
|
||||
if let Some(entries) = dv.and_then(|value| value.as_container()) {
|
||||
println!(" Verified SecurePayload: {:?}", entries);
|
||||
|
|
@ -388,7 +369,7 @@ pub fn process_and_respond(
|
|||
let response = CommunicationValue::from_comm(CommunicationType::Pong, tm)
|
||||
.add_data(desc_id, description)
|
||||
.map_err(|e| e.to_string())?
|
||||
.add_data(ts_id, DataValue::UnsignedNumber(now))
|
||||
.add_data(ts_id, DataValue::UnsignedNumber(now as u128))
|
||||
.map_err(|e| e.to_string())?
|
||||
.add_data(
|
||||
data_id,
|
||||
|
|
|
|||
|
|
@ -27,7 +27,7 @@ pub async fn export_host_public_keys(
|
|||
save_public_key_bundle(&bundle, "host.mpkb")?;
|
||||
|
||||
/* The web client fetches the bundle as hex over HTTP. */
|
||||
let bundle_hex = hex::encode(bundle.try_as_bytes()?);
|
||||
let bundle_hex = hex::encode(bundle.as_bytes());
|
||||
fs::write("host_public_key_bundle.hex", &bundle_hex).await?;
|
||||
fs::create_dir_all("web-client/public").await?;
|
||||
fs::write("web-client/public/host_public_key_bundle.hex", &bundle_hex).await?;
|
||||
|
|
|
|||
|
|
@ -38,7 +38,6 @@ async fn handle_pipe_loopback(
|
|||
conn: &mtp::webserver::WebMTPConnection,
|
||||
request: mtp::host::PipeRequest<
|
||||
mtp::webserver::WebMtpSender,
|
||||
mtp::webserver::WebMtpReceiver,
|
||||
mtp::webserver::H3TransportReceiver,
|
||||
>,
|
||||
) -> Result<u64, Box<dyn std::error::Error>> {
|
||||
|
|
@ -101,20 +100,19 @@ async fn main() -> Result<(), Box<dyn std::error::Error>> {
|
|||
serde_json::to_string_pretty(&*db).ok()
|
||||
};
|
||||
|
||||
if let Some(json) = json
|
||||
&& let Err(error) = tokio::fs::write("clients.json", json).await
|
||||
{
|
||||
if let Some(json) = json {
|
||||
if let Err(error) = tokio::fs::write("clients.json", json).await {
|
||||
eprintln!("Failed to persist clients.json: {error}");
|
||||
}
|
||||
}
|
||||
|
||||
println!("Registered new client with ID: {id}");
|
||||
id
|
||||
}) as Pin<Box<dyn Future<Output = u64> + Send>>
|
||||
};
|
||||
|
||||
let decrypt_keyring_bytes = host_keyring.try_to_bytes()?;
|
||||
let decrypt_keyring = Arc::new(
|
||||
match mtp::crypto::Keyring::from_bytes(&decrypt_keyring_bytes) {
|
||||
match mtp::crypto::Keyring::from_bytes(&host_keyring.to_bytes()) {
|
||||
Ok(keyring) => keyring,
|
||||
Err(e) => {
|
||||
return Err(format!("failed to re-load host keyring for decryption: {e}").into());
|
||||
|
|
|
|||
|
|
@ -108,7 +108,6 @@ pub struct ServerMetrics {
|
|||
}
|
||||
|
||||
impl ServerMetrics {
|
||||
#[cfg(test)]
|
||||
pub fn new() -> Self {
|
||||
Self {
|
||||
inner: Mutex::new(Inner {
|
||||
|
|
@ -219,7 +218,6 @@ impl ServerMetrics {
|
|||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
pub fn snapshot(&self) -> ServerMetricsFile {
|
||||
let inner = self.inner.lock().unwrap();
|
||||
self.to_file(&inner)
|
||||
|
|
@ -555,11 +553,11 @@ mod tests {
|
|||
let metrics = ServerMetrics::new();
|
||||
|
||||
for i in 0..3 {
|
||||
let mut session = metrics.start_session(1000 + i, format!("session {i}"));
|
||||
let mut session = metrics.start_session(1000 + i as u64, format!("session {i}"));
|
||||
for _ in 0..(i + 1) * 2 {
|
||||
session.record_message(Duration::from_millis(1 + i), true);
|
||||
}
|
||||
session.record_pipe((i + 1) * 1000);
|
||||
session.record_pipe((i as u64 + 1) * 1000);
|
||||
session.finish(format!("exit {i}"));
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -6,7 +6,8 @@ edition = "2024"
|
|||
[dependencies]
|
||||
# Only the plain key types (`Keyring`, `PublicKeyBundle`, `CryptoError`) are
|
||||
# needed here; those are always compiled, so no crypto features are required.
|
||||
mtp-crypto = { version = "0.3.0", path = "../crypto", default-features = false, features = ["chacha20poly1305", "hkdf", "password-kdf"] }
|
||||
mtp-crypto = { version = "0.3.0", path = "../crypto", default-features = false, features = ["chacha20poly1305", "hkdf"] }
|
||||
argon2 = "0.5"
|
||||
rand = "0.10.2"
|
||||
|
||||
thiserror = "2"
|
||||
|
|
|
|||
|
|
@ -119,20 +119,6 @@ fn temporary_path(path: &Path, attempt: u64) -> io::Result<PathBuf> {
|
|||
|
||||
static TEMP_COUNTER: AtomicU64 = AtomicU64::new(0);
|
||||
|
||||
#[cfg(unix)]
|
||||
fn sync_parent_directory(path: &Path) -> io::Result<()> {
|
||||
let parent = path
|
||||
.parent()
|
||||
.filter(|parent| !parent.as_os_str().is_empty())
|
||||
.unwrap_or_else(|| Path::new("."));
|
||||
fs::File::open(parent)?.sync_all()
|
||||
}
|
||||
|
||||
#[cfg(not(unix))]
|
||||
fn sync_parent_directory(_path: &Path) -> io::Result<()> {
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn write_secret_atomic(path: &Path, bytes: &[u8]) -> io::Result<()> {
|
||||
use std::io::Write;
|
||||
|
||||
|
|
@ -160,7 +146,7 @@ fn write_secret_atomic(path: &Path, bytes: &[u8]) -> io::Result<()> {
|
|||
let _ = fs::remove_file(&temporary);
|
||||
return Err(error);
|
||||
}
|
||||
sync_parent_directory(path)
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn derive_key(
|
||||
|
|
@ -170,12 +156,21 @@ fn derive_key(
|
|||
iterations: u32,
|
||||
lanes: u32,
|
||||
) -> Result<Zeroizing<[u8; 32]>, FileError> {
|
||||
if salt.len() != SALT_LEN {
|
||||
if salt.len() != SALT_LEN
|
||||
|| !(8 * 1024..=256 * 1024).contains(&memory_kib)
|
||||
|| !(1..=10).contains(&iterations)
|
||||
|| !(1..=8).contains(&lanes)
|
||||
{
|
||||
return Err(FileError::Crypto(CryptoError::KdfError));
|
||||
}
|
||||
Ok(Zeroizing::new(mtp_crypto::derive_password_key(
|
||||
passphrase, salt, memory_kib, iterations, lanes,
|
||||
)?))
|
||||
let params = argon2::Params::new(memory_kib, iterations, lanes, Some(32))
|
||||
.map_err(|_| FileError::Crypto(CryptoError::KdfError))?;
|
||||
let argon = argon2::Argon2::new(argon2::Algorithm::Argon2id, argon2::Version::V0x13, params);
|
||||
let mut key = Zeroizing::new([0u8; 32]);
|
||||
argon
|
||||
.hash_password_into(passphrase, salt, key.as_mut())
|
||||
.map_err(|_| FileError::Crypto(CryptoError::KdfError))?;
|
||||
Ok(key)
|
||||
}
|
||||
|
||||
fn protected_header_aad(parameters: &[u8]) -> Vec<u8> {
|
||||
|
|
@ -211,7 +206,7 @@ pub fn save_keyring(
|
|||
parameters.extend_from_slice(&ARGON2_LANES.to_be_bytes());
|
||||
parameters.extend_from_slice(&salt);
|
||||
let cipher = ChaCha20Poly1305::new(*key);
|
||||
let plaintext = keyring.try_to_bytes()?;
|
||||
let plaintext = keyring.to_bytes();
|
||||
let encrypted = cipher.encrypt(&plaintext, &protected_header_aad(¶meters))?;
|
||||
let mut payload = Vec::with_capacity(PROTECTED_PARAMS_LEN + encrypted.len());
|
||||
payload.extend_from_slice(¶meters);
|
||||
|
|
@ -264,7 +259,7 @@ pub fn load_keyring(path: impl AsRef<Path>, passphrase: &[u8]) -> Result<Keyring
|
|||
/// Explicitly save the legacy plaintext format for tests and development.
|
||||
#[cfg(any(test, feature = "raw"))]
|
||||
pub fn save_keyring_raw(keyring: &Keyring, path: impl AsRef<Path>) -> Result<(), FileError> {
|
||||
let payload = keyring.try_to_bytes()?;
|
||||
let payload = keyring.to_bytes();
|
||||
let bytes = Zeroizing::new(encode(KEYRING_MAGIC, RAW_FORMAT_VERSION, &payload));
|
||||
write_secret_atomic(path.as_ref(), &bytes)?;
|
||||
Ok(())
|
||||
|
|
@ -291,8 +286,7 @@ pub fn save_public_key_bundle(
|
|||
bundle: &PublicKeyBundle,
|
||||
path: impl AsRef<Path>,
|
||||
) -> Result<(), FileError> {
|
||||
let bundle_bytes = bundle.try_as_bytes()?;
|
||||
let bytes = encode(BUNDLE_MAGIC, BUNDLE_FORMAT_VERSION, &bundle_bytes);
|
||||
let bytes = encode(BUNDLE_MAGIC, BUNDLE_FORMAT_VERSION, &bundle.as_bytes());
|
||||
fs::write(path, bytes)?;
|
||||
Ok(())
|
||||
}
|
||||
|
|
@ -345,7 +339,7 @@ mod tests {
|
|||
let keyring = sample_keyring();
|
||||
save_keyring(&keyring, &path, b"correct horse battery staple")?;
|
||||
let loaded = load_keyring(&path, b"correct horse battery staple")?;
|
||||
assert_eq!(keyring.try_to_bytes()?, loaded.try_to_bytes()?);
|
||||
assert_eq!(keyring.to_bytes(), loaded.to_bytes());
|
||||
let _ = fs::remove_file(&path);
|
||||
Ok(())
|
||||
}
|
||||
|
|
@ -356,7 +350,7 @@ mod tests {
|
|||
let bundle = Keyring::generate().public_key_bundle();
|
||||
save_public_key_bundle(&bundle, &path)?;
|
||||
let loaded = load_public_key_bundle(&path)?;
|
||||
assert_eq!(bundle.try_as_bytes()?, loaded.try_as_bytes()?);
|
||||
assert_eq!(bundle.as_bytes(), loaded.as_bytes());
|
||||
let _ = fs::remove_file(&path);
|
||||
Ok(())
|
||||
}
|
||||
|
|
@ -420,7 +414,7 @@ mod tests {
|
|||
Err(FileError::UnprotectedKeyring)
|
||||
));
|
||||
let loaded = load_keyring_raw(&path)?;
|
||||
assert_eq!(keyring.try_to_bytes()?, loaded.try_to_bytes()?);
|
||||
assert_eq!(keyring.to_bytes(), loaded.to_bytes());
|
||||
let _ = fs::remove_file(&path);
|
||||
Ok(())
|
||||
}
|
||||
|
|
@ -429,7 +423,7 @@ mod tests {
|
|||
fn protected_keyring_is_not_plaintext() -> Result<(), Box<dyn std::error::Error>> {
|
||||
let path = temp_path(KEYRING_EXTENSION);
|
||||
let keyring = sample_keyring();
|
||||
let serialized = keyring.try_to_bytes()?;
|
||||
let serialized = keyring.to_bytes();
|
||||
save_keyring(&keyring, &path, b"passphrase")?;
|
||||
let stored = fs::read(&path)?;
|
||||
assert!(
|
||||
|
|
|
|||
32
flake.nix
32
flake.nix
|
|
@ -1,32 +1,29 @@
|
|||
{
|
||||
description = "MTP - Methanium Transport Protocol";
|
||||
|
||||
inputs = {
|
||||
nixpkgs.url = "github:NixOS/nixpkgs/nixos-unstable";
|
||||
rust-overlay.url = "github:oxalica/rust-overlay";
|
||||
};
|
||||
|
||||
outputs =
|
||||
{
|
||||
outputs = {
|
||||
self,
|
||||
nixpkgs,
|
||||
rust-overlay,
|
||||
}:
|
||||
let
|
||||
}: let
|
||||
systems = [
|
||||
"aarch64-darwin"
|
||||
"aarch64-linux"
|
||||
"x86_64-darwin"
|
||||
"x86_64-linux"
|
||||
];
|
||||
eachSystem =
|
||||
f:
|
||||
eachSystem = f:
|
||||
nixpkgs.lib.foldl' nixpkgs.lib.recursiveUpdate {} (
|
||||
map (system: nixpkgs.lib.mapAttrs (_: value: {${system} = value;}) (f system)) systems
|
||||
);
|
||||
in
|
||||
eachSystem (
|
||||
system:
|
||||
let
|
||||
system: let
|
||||
overlays = [rust-overlay.overlays.default];
|
||||
pkgs = import nixpkgs {inherit system overlays;};
|
||||
|
||||
|
|
@ -58,15 +55,7 @@
|
|||
|
||||
buildAll = pkgs.writeShellApplication {
|
||||
name = "mtp-build-all";
|
||||
runtimeInputs = [
|
||||
rustToolchain
|
||||
pkgs.cargo-deny
|
||||
pkgs.wasm-pack
|
||||
pkgs.pnpm
|
||||
pkgs.coreutils
|
||||
clippyCheck
|
||||
macheteCheck
|
||||
];
|
||||
runtimeInputs = [rustToolchain pkgs.cargo-deny pkgs.wasm-pack pkgs.pnpm pkgs.coreutils clippyCheck macheteCheck];
|
||||
text = ''
|
||||
export MTP_TYPE_MAPS="''${MTP_TYPE_MAPS:-$PWD/example/type-maps.yaml}"
|
||||
|
||||
|
|
@ -78,6 +67,7 @@
|
|||
cargo check --manifest-path example/Cargo.toml --workspace --all-targets --all-features
|
||||
mtp-clippy
|
||||
mtp-machete
|
||||
pnpm run dup
|
||||
pnpm run build
|
||||
RUSTFLAGS="--cfg web_sys_unstable_apis" wasm-pack test --node wasm
|
||||
pnpm run test:e2e
|
||||
|
|
@ -90,17 +80,13 @@
|
|||
|
||||
healthCheck = pkgs.writeShellApplication {
|
||||
name = "mtp-health";
|
||||
runtimeInputs = [
|
||||
clippyCheck
|
||||
macheteCheck
|
||||
];
|
||||
runtimeInputs = [clippyCheck macheteCheck];
|
||||
text = ''
|
||||
mtp-clippy
|
||||
mtp-machete
|
||||
'';
|
||||
};
|
||||
in
|
||||
{
|
||||
in {
|
||||
devShells = {
|
||||
default = pkgs.mkShell {
|
||||
name = "mtp-dev";
|
||||
|
|
|
|||
|
|
@ -9,7 +9,6 @@ mtp-codec = { version = "0.3.0", path = "../codec", features = ["registry"] }
|
|||
mtp-transport = { version = "0.3.0", path = "../transport", features = ["host"] }
|
||||
mtp-crypto = { version = "0.3.0", path = "../crypto", optional = true }
|
||||
rand = "0.10"
|
||||
thiserror = "2"
|
||||
tokio = { version = "1", features = ["macros", "rt", "time", "sync"] }
|
||||
tracing = "0.1"
|
||||
wtransport = "0.7"
|
||||
|
|
|
|||
|
|
@ -5,14 +5,10 @@ use std::collections::HashMap;
|
|||
#[cfg(feature = "crypto")]
|
||||
use std::collections::HashSet;
|
||||
#[cfg(feature = "crypto")]
|
||||
use std::collections::VecDeque;
|
||||
#[cfg(feature = "crypto")]
|
||||
use std::pin::Pin;
|
||||
#[cfg(feature = "crypto")]
|
||||
use std::sync::{Arc, Mutex};
|
||||
#[cfg(feature = "crypto")]
|
||||
use std::time::{Duration as StdDuration, Instant};
|
||||
#[cfg(feature = "crypto")]
|
||||
use tokio::time::Duration;
|
||||
|
||||
pub use mtp_transport::Policy;
|
||||
|
|
@ -86,159 +82,6 @@ pub enum AuthenticationPolicy {
|
|||
Unauthenticated,
|
||||
}
|
||||
|
||||
/// Transport-supplied identity used to scope authentication attempt limits.
|
||||
/// Concrete hosts should populate these fields from the accepted connection;
|
||||
/// the zero/empty defaults exist only for transport-neutral callers.
|
||||
#[derive(Debug, Clone, Default, PartialEq, Eq)]
|
||||
pub struct AuthenticationContext {
|
||||
pub peer_network_identity: Option<String>,
|
||||
pub connection_id: u64,
|
||||
}
|
||||
|
||||
#[cfg(feature = "crypto")]
|
||||
#[derive(Debug, Clone, PartialEq, Eq)]
|
||||
pub struct AuthenticationAttempt {
|
||||
pub peer_network_identity: Option<String>,
|
||||
pub connection_id: u64,
|
||||
pub claimed_client_id: Option<u64>,
|
||||
pub registration: bool,
|
||||
}
|
||||
|
||||
#[cfg(feature = "crypto")]
|
||||
#[derive(Debug, thiserror::Error)]
|
||||
pub enum AuthenticationLimitError {
|
||||
#[error("authentication limiter storage is unavailable")]
|
||||
Store,
|
||||
}
|
||||
|
||||
#[cfg(feature = "crypto")]
|
||||
pub trait AuthenticationAttemptLimiter: Send + Sync {
|
||||
fn allow(&self, context: &AuthenticationAttempt) -> Result<bool, AuthenticationLimitError>;
|
||||
}
|
||||
|
||||
#[cfg(feature = "crypto")]
|
||||
#[derive(Debug, Clone, Hash, PartialEq, Eq)]
|
||||
enum AuthenticationLimitKey {
|
||||
Peer(String),
|
||||
Connection(u64),
|
||||
Client(u64),
|
||||
Registration,
|
||||
Global,
|
||||
}
|
||||
|
||||
#[cfg(feature = "crypto")]
|
||||
#[derive(Debug)]
|
||||
pub struct InMemoryAuthenticationAttemptLimiter {
|
||||
max_attempts: usize,
|
||||
window: StdDuration,
|
||||
max_keys: usize,
|
||||
by_peer: bool,
|
||||
by_connection: bool,
|
||||
by_client: bool,
|
||||
by_registration: bool,
|
||||
attempts: Mutex<HashMap<AuthenticationLimitKey, VecDeque<Instant>>>,
|
||||
}
|
||||
|
||||
#[cfg(feature = "crypto")]
|
||||
impl InMemoryAuthenticationAttemptLimiter {
|
||||
pub fn new(max_attempts: usize, window: StdDuration) -> Self {
|
||||
Self {
|
||||
max_attempts,
|
||||
window,
|
||||
max_keys: 100_000,
|
||||
by_peer: true,
|
||||
by_connection: true,
|
||||
by_client: true,
|
||||
by_registration: true,
|
||||
attempts: Mutex::new(HashMap::new()),
|
||||
}
|
||||
}
|
||||
|
||||
pub fn with_keys(
|
||||
mut self,
|
||||
by_peer: bool,
|
||||
by_connection: bool,
|
||||
by_client: bool,
|
||||
by_registration: bool,
|
||||
) -> Self {
|
||||
self.by_peer = by_peer;
|
||||
self.by_connection = by_connection;
|
||||
self.by_client = by_client;
|
||||
self.by_registration = by_registration;
|
||||
self
|
||||
}
|
||||
|
||||
pub fn with_max_keys(mut self, max_keys: usize) -> Self {
|
||||
self.max_keys = max_keys.max(1);
|
||||
self
|
||||
}
|
||||
|
||||
fn keys(&self, context: &AuthenticationAttempt) -> Vec<AuthenticationLimitKey> {
|
||||
let mut keys = Vec::with_capacity(5);
|
||||
if self.by_peer
|
||||
&& let Some(peer) = context.peer_network_identity.as_ref()
|
||||
{
|
||||
keys.push(AuthenticationLimitKey::Peer(peer.clone()));
|
||||
}
|
||||
if self.by_connection && context.connection_id != 0 {
|
||||
keys.push(AuthenticationLimitKey::Connection(context.connection_id));
|
||||
}
|
||||
if self.by_client
|
||||
&& let Some(client_id) = context.claimed_client_id
|
||||
{
|
||||
keys.push(AuthenticationLimitKey::Client(client_id));
|
||||
}
|
||||
if self.by_registration && context.registration {
|
||||
keys.push(AuthenticationLimitKey::Registration);
|
||||
}
|
||||
// Keep one global bucket as a backstop when an attacker varies the
|
||||
// claimed client ID or presents no peer/connection identity.
|
||||
keys.push(AuthenticationLimitKey::Global);
|
||||
keys
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(feature = "crypto")]
|
||||
impl AuthenticationAttemptLimiter for InMemoryAuthenticationAttemptLimiter {
|
||||
fn allow(&self, context: &AuthenticationAttempt) -> Result<bool, AuthenticationLimitError> {
|
||||
if self.max_attempts == 0 {
|
||||
return Ok(false);
|
||||
}
|
||||
let now = Instant::now();
|
||||
let cutoff = now.checked_sub(self.window);
|
||||
let keys = self.keys(context);
|
||||
let mut attempts = self
|
||||
.attempts
|
||||
.lock()
|
||||
.map_err(|_| AuthenticationLimitError::Store)?;
|
||||
|
||||
for key in &keys {
|
||||
if let Some(history) = attempts.get_mut(key) {
|
||||
while history
|
||||
.front()
|
||||
.is_some_and(|timestamp| cutoff.is_some_and(|cutoff| *timestamp <= cutoff))
|
||||
{
|
||||
history.pop_front();
|
||||
}
|
||||
if history.len() >= self.max_attempts {
|
||||
return Ok(false);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
for key in keys {
|
||||
if !attempts.contains_key(&key)
|
||||
&& attempts.len() >= self.max_keys
|
||||
&& let Some(oldest) = attempts.keys().next().cloned()
|
||||
{
|
||||
attempts.remove(&oldest);
|
||||
}
|
||||
attempts.entry(key).or_default().push_back(now);
|
||||
}
|
||||
Ok(true)
|
||||
}
|
||||
}
|
||||
|
||||
pub struct HostConfig {
|
||||
pub ip: IpAddr,
|
||||
pub port: u16,
|
||||
|
|
@ -251,8 +94,6 @@ pub struct HostConfig {
|
|||
#[cfg(feature = "crypto")]
|
||||
pub authentication_policy: AuthenticationPolicy,
|
||||
#[cfg(feature = "crypto")]
|
||||
authentication_policy_explicit: bool,
|
||||
#[cfg(feature = "crypto")]
|
||||
pub auth_timeout: Duration,
|
||||
#[cfg(feature = "crypto")]
|
||||
pub require_pq: bool,
|
||||
|
|
@ -272,10 +113,6 @@ pub struct HostConfig {
|
|||
pub complete_register: CompleteRegister,
|
||||
#[cfg(feature = "crypto")]
|
||||
pub find_registered_client: Option<FindRegisteredClient>,
|
||||
#[cfg(feature = "crypto")]
|
||||
pub auth_limiter: Arc<dyn AuthenticationAttemptLimiter>,
|
||||
#[cfg(feature = "crypto")]
|
||||
pub conceal_authentication_identities: bool,
|
||||
}
|
||||
|
||||
impl HostConfig {
|
||||
|
|
@ -290,8 +127,6 @@ impl HostConfig {
|
|||
#[cfg(feature = "crypto")]
|
||||
authentication_policy: AuthenticationPolicy::Unauthenticated,
|
||||
#[cfg(feature = "crypto")]
|
||||
authentication_policy_explicit: false,
|
||||
#[cfg(feature = "crypto")]
|
||||
auth_timeout: Duration::from_secs(30),
|
||||
#[cfg(feature = "crypto")]
|
||||
require_pq: true,
|
||||
|
|
@ -318,13 +153,6 @@ impl HostConfig {
|
|||
complete_register: Box::new(|_, _| Box::pin(async { 0 })),
|
||||
#[cfg(feature = "crypto")]
|
||||
find_registered_client: None,
|
||||
#[cfg(feature = "crypto")]
|
||||
auth_limiter: Arc::new(InMemoryAuthenticationAttemptLimiter::new(
|
||||
32,
|
||||
StdDuration::from_secs(60),
|
||||
)),
|
||||
#[cfg(feature = "crypto")]
|
||||
conceal_authentication_identities: true,
|
||||
}
|
||||
}
|
||||
|
||||
|
|
@ -345,9 +173,7 @@ impl HostConfig {
|
|||
get_existing_client: GetExistingClient,
|
||||
complete_register: CompleteRegister,
|
||||
) -> Self {
|
||||
if !self.authentication_policy_explicit {
|
||||
self.authentication_policy = AuthenticationPolicy::ForceAuthentication;
|
||||
}
|
||||
self.host_keyring = host_keyring;
|
||||
self.get_existing_client = Box::new(get_existing_client);
|
||||
self.complete_register = Box::new(complete_register);
|
||||
|
|
@ -357,7 +183,6 @@ impl HostConfig {
|
|||
#[cfg(feature = "crypto")]
|
||||
pub fn with_authentication_policy(mut self, policy: AuthenticationPolicy) -> Self {
|
||||
self.authentication_policy = policy;
|
||||
self.authentication_policy_explicit = true;
|
||||
self
|
||||
}
|
||||
|
||||
|
|
@ -385,141 +210,4 @@ impl HostConfig {
|
|||
self.find_registered_client = Some(lookup);
|
||||
self
|
||||
}
|
||||
|
||||
#[cfg(feature = "crypto")]
|
||||
pub fn with_authentication_limiter(
|
||||
mut self,
|
||||
limiter: Arc<dyn AuthenticationAttemptLimiter>,
|
||||
) -> Self {
|
||||
self.auth_limiter = limiter;
|
||||
self
|
||||
}
|
||||
|
||||
#[cfg(feature = "crypto")]
|
||||
pub fn with_authentication_identity_concealment(mut self, conceal: bool) -> Self {
|
||||
self.conceal_authentication_identities = conceal;
|
||||
self
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(all(test, feature = "crypto"))]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn authentication_attempt_limiter_rejects_repeated_attempts() {
|
||||
let limiter = InMemoryAuthenticationAttemptLimiter::new(1, StdDuration::from_secs(60))
|
||||
.with_keys(false, true, false, false);
|
||||
let attempt = AuthenticationAttempt {
|
||||
peer_network_identity: None,
|
||||
connection_id: 9,
|
||||
claimed_client_id: Some(42),
|
||||
registration: false,
|
||||
};
|
||||
|
||||
assert!(limiter.allow(&attempt).expect("first attempt decision"));
|
||||
assert!(!limiter.allow(&attempt).expect("second attempt decision"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn authentication_attempt_limiter_can_scope_registration_separately() {
|
||||
let limiter = InMemoryAuthenticationAttemptLimiter::new(2, StdDuration::from_secs(60))
|
||||
.with_keys(false, false, false, true);
|
||||
let login = AuthenticationAttempt {
|
||||
peer_network_identity: None,
|
||||
connection_id: 1,
|
||||
claimed_client_id: None,
|
||||
registration: false,
|
||||
};
|
||||
let registration = AuthenticationAttempt {
|
||||
registration: true,
|
||||
..login.clone()
|
||||
};
|
||||
|
||||
assert!(limiter.allow(&login).expect("login attempt decision"));
|
||||
assert!(
|
||||
limiter
|
||||
.allow(®istration)
|
||||
.expect("registration attempt decision")
|
||||
);
|
||||
assert!(
|
||||
!limiter
|
||||
.allow(®istration)
|
||||
.expect("repeated registration decision")
|
||||
);
|
||||
}
|
||||
|
||||
fn test_keyring() -> mtp_crypto::Keyring {
|
||||
mtp_crypto::Keyring::new(
|
||||
mtp_crypto::KemPublicKey::new(Vec::new()),
|
||||
mtp_crypto::KemPrivateKey::new(Vec::new()),
|
||||
mtp_crypto::SignaturePqPublicKey::new(Vec::new()),
|
||||
mtp_crypto::SignaturePqPrivateKey::new(Vec::new()),
|
||||
mtp_crypto::SignaturePublicKey::new(Vec::new()),
|
||||
mtp_crypto::SignaturePrivateKey::new(Vec::new()),
|
||||
)
|
||||
}
|
||||
|
||||
fn test_get_existing_client() -> GetExistingClient {
|
||||
Box::new(|_, _| Box::pin(async { None }))
|
||||
}
|
||||
|
||||
fn test_complete_register() -> CompleteRegister {
|
||||
Box::new(|_, _| Box::pin(async { 1 }))
|
||||
}
|
||||
|
||||
fn test_config() -> HostConfig {
|
||||
HostConfig::new(
|
||||
IpAddr::V4(std::net::Ipv4Addr::LOCALHOST),
|
||||
4433,
|
||||
Vec::new(),
|
||||
Vec::new(),
|
||||
)
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn with_authentication_defaults_to_force_authentication() {
|
||||
let config = test_config().with_authentication(
|
||||
test_keyring(),
|
||||
test_get_existing_client(),
|
||||
test_complete_register(),
|
||||
);
|
||||
|
||||
assert_eq!(
|
||||
config.authentication_policy,
|
||||
AuthenticationPolicy::ForceAuthentication
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn explicit_authentication_policy_before_with_authentication_is_preserved() {
|
||||
let config = test_config()
|
||||
.with_authentication_policy(AuthenticationPolicy::AllowAuthentication)
|
||||
.with_authentication(
|
||||
test_keyring(),
|
||||
test_get_existing_client(),
|
||||
test_complete_register(),
|
||||
);
|
||||
|
||||
assert_eq!(
|
||||
config.authentication_policy,
|
||||
AuthenticationPolicy::AllowAuthentication
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn explicit_authentication_policy_after_with_authentication_is_preserved() {
|
||||
let config = test_config()
|
||||
.with_authentication(
|
||||
test_keyring(),
|
||||
test_get_existing_client(),
|
||||
test_complete_register(),
|
||||
)
|
||||
.with_authentication_policy(AuthenticationPolicy::AllowAuthentication);
|
||||
|
||||
assert_eq!(
|
||||
config.authentication_policy,
|
||||
AuthenticationPolicy::AllowAuthentication
|
||||
);
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -11,10 +11,7 @@ use tokio::sync::{Mutex, mpsc};
|
|||
#[cfg(feature = "crypto")]
|
||||
use crate::error::random_client_id;
|
||||
#[cfg(feature = "pipes")]
|
||||
use crate::pipe::{
|
||||
PendingCreationGuard, PipeDispatcher, PipeReceiver, PipeRequest, PipeSender,
|
||||
is_expired_creation, run_dispatcher,
|
||||
};
|
||||
use crate::pipe::{PipeDispatcher, PipeReceiver, PipeRequest, PipeSender, run_dispatcher};
|
||||
#[cfg(feature = "pipes")]
|
||||
use mtp_transport::Policy;
|
||||
|
||||
|
|
@ -73,7 +70,7 @@ pub struct MTPConnection<
|
|||
#[cfg(feature = "pipes")]
|
||||
pub(crate) app_rx: Mutex<mpsc::Receiver<Result<CommunicationValue, CommunicationError>>>,
|
||||
#[cfg(feature = "pipes")]
|
||||
pub(crate) pipe_req_rx: Mutex<mpsc::Receiver<PipeRequest<S, R, P>>>,
|
||||
pub(crate) pipe_req_rx: Mutex<mpsc::Receiver<PipeRequest<S, P>>>,
|
||||
#[cfg(feature = "pipes")]
|
||||
pub(crate) pipe_dispatcher: Arc<PipeDispatcher<P>>,
|
||||
#[cfg(not(feature = "pipes"))]
|
||||
|
|
@ -148,12 +145,10 @@ where
|
|||
remote_addr: Option<SocketAddr>,
|
||||
) -> Self {
|
||||
let policy = Arc::new(Policy::default());
|
||||
let receiver_queue_capacity = policy.receiver_queue_capacity.max(1);
|
||||
let (app_tx, app_rx) = mpsc::channel(receiver_queue_capacity);
|
||||
let (pipe_req_tx, pipe_req_rx) = mpsc::channel(receiver_queue_capacity);
|
||||
let (app_tx, app_rx) = mpsc::channel(policy.receiver_queue_capacity);
|
||||
let (pipe_req_tx, pipe_req_rx) = mpsc::channel(policy.receiver_queue_capacity);
|
||||
let dispatcher = Arc::new(PipeDispatcher {
|
||||
pending_creations: std::sync::Mutex::new(std::collections::HashMap::new()),
|
||||
expired_creations: std::sync::Mutex::new(std::collections::HashMap::new()),
|
||||
pending_creations: Mutex::new(std::collections::HashMap::new()),
|
||||
pending_pipes: Mutex::new(std::collections::HashMap::new()),
|
||||
policy,
|
||||
type_map: codec.type_map().clone(),
|
||||
|
|
@ -203,12 +198,10 @@ where
|
|||
remote_addr: Option<SocketAddr>,
|
||||
policy: Arc<Policy>,
|
||||
) -> Self {
|
||||
let receiver_queue_capacity = policy.receiver_queue_capacity.max(1);
|
||||
let (app_tx, app_rx) = mpsc::channel(receiver_queue_capacity);
|
||||
let (pipe_req_tx, pipe_req_rx) = mpsc::channel(receiver_queue_capacity);
|
||||
let (app_tx, app_rx) = mpsc::channel(policy.receiver_queue_capacity);
|
||||
let (pipe_req_tx, pipe_req_rx) = mpsc::channel(policy.receiver_queue_capacity);
|
||||
let dispatcher = Arc::new(PipeDispatcher {
|
||||
pending_creations: std::sync::Mutex::new(std::collections::HashMap::new()),
|
||||
expired_creations: std::sync::Mutex::new(std::collections::HashMap::new()),
|
||||
pending_creations: Mutex::new(std::collections::HashMap::new()),
|
||||
pending_pipes: Mutex::new(std::collections::HashMap::new()),
|
||||
policy,
|
||||
type_map: codec.type_map().clone(),
|
||||
|
|
@ -328,37 +321,19 @@ where
|
|||
pub async fn create_pipe(
|
||||
&self,
|
||||
description: &str,
|
||||
) -> Result<crate::pipe::PipeHandle<S, P>, mtp_common::PipeError> {
|
||||
) -> Result<crate::pipe::PipeHandle<S>, mtp_common::PipeError> {
|
||||
let (response_tx, response_rx) = tokio::sync::oneshot::channel();
|
||||
let pipe_id = {
|
||||
let mut pending = self
|
||||
.pipe_dispatcher
|
||||
.pending_creations
|
||||
.lock()
|
||||
.map_err(|_| mtp_common::PipeError::ConnectionClosed)?;
|
||||
let mut pending = self.pipe_dispatcher.pending_creations.lock().await;
|
||||
let pipe_id = loop {
|
||||
let candidate = rand::random::<u32>();
|
||||
if candidate != 0
|
||||
&& !pending.contains_key(&candidate)
|
||||
&& !is_expired_creation(&self.pipe_dispatcher, candidate)
|
||||
{
|
||||
if candidate != 0 && !pending.contains_key(&candidate) {
|
||||
break candidate;
|
||||
}
|
||||
};
|
||||
let token = Arc::new(());
|
||||
pending.insert(
|
||||
pipe_id,
|
||||
crate::pipe::PendingCreation {
|
||||
token: token.clone(),
|
||||
sender: response_tx,
|
||||
},
|
||||
);
|
||||
drop(pending);
|
||||
(pipe_id, token)
|
||||
pending.insert(pipe_id, response_tx);
|
||||
pipe_id
|
||||
};
|
||||
let (pipe_id, token) = pipe_id;
|
||||
let mut creation_guard =
|
||||
PendingCreationGuard::new(self.pipe_dispatcher.clone(), pipe_id, token.clone());
|
||||
|
||||
let request = CommunicationValue::new_with_type_map(
|
||||
CommunicationType::PipeRequest,
|
||||
|
|
@ -367,21 +342,23 @@ where
|
|||
.with_id(pipe_id)
|
||||
.add_typed_default(DataType::Description, DataValue::Str(description.into()));
|
||||
if let Err(error) = self.sender.send_pipe_message(&request).await {
|
||||
self.pipe_dispatcher
|
||||
.pending_creations
|
||||
.lock()
|
||||
.await
|
||||
.remove(&pipe_id);
|
||||
return Err(mtp_common::PipeError::from(error));
|
||||
}
|
||||
|
||||
creation_guard.disarm();
|
||||
Ok(crate::pipe::PipeHandle {
|
||||
pipe_id,
|
||||
description: description.to_owned(),
|
||||
sender: self.sender.clone(),
|
||||
response_rx,
|
||||
dispatcher: self.pipe_dispatcher.clone(),
|
||||
token,
|
||||
})
|
||||
}
|
||||
|
||||
pub async fn receive_pipe(&self) -> Result<PipeRequest<S, R, P>, CommunicationError> {
|
||||
pub async fn receive_pipe(&self) -> Result<PipeRequest<S, P>, CommunicationError> {
|
||||
self.pipe_req_rx
|
||||
.lock()
|
||||
.await
|
||||
|
|
|
|||
161
host/src/engine.rs
Executable file → Normal file
161
host/src/engine.rs
Executable file → Normal file
|
|
@ -4,7 +4,7 @@
|
|||
//! and the web server's `MTPWebServer` to perform the MTP opening handshake,
|
||||
//! version negotiation, authentication, and guest assignment.
|
||||
|
||||
use crate::config::{AuthenticationContext, HostConfig};
|
||||
use crate::config::HostConfig;
|
||||
use crate::error::AcceptError;
|
||||
use mtp_codec::{
|
||||
CommunicationType, CommunicationValue, DataType, DataValue, TypeMap, Version,
|
||||
|
|
@ -127,32 +127,19 @@ impl HandshakeEngine {
|
|||
&self,
|
||||
sender: &S,
|
||||
receiver: &R,
|
||||
) -> Result<HandshakeResult, AcceptError> {
|
||||
self.accept_with_context(sender, receiver, AuthenticationContext::default())
|
||||
.await
|
||||
}
|
||||
|
||||
/// Run the opening handshake with transport-provided authentication
|
||||
/// scoping information.
|
||||
pub async fn accept_with_context<S: HandshakeSender, R: HandshakeReceiver>(
|
||||
&self,
|
||||
sender: &S,
|
||||
receiver: &R,
|
||||
context: AuthenticationContext,
|
||||
) -> Result<HandshakeResult, AcceptError> {
|
||||
#[cfg(feature = "crypto")]
|
||||
{
|
||||
self.accept_until_with_context(
|
||||
self.accept_until(
|
||||
sender,
|
||||
receiver,
|
||||
tokio::time::Instant::now() + self.config.auth_timeout,
|
||||
context,
|
||||
)
|
||||
.await
|
||||
}
|
||||
#[cfg(not(feature = "crypto"))]
|
||||
{
|
||||
let result = self.accept_inner(sender, receiver, &context).await;
|
||||
let result = self.accept_inner(sender, receiver).await;
|
||||
if result.is_err() {
|
||||
sender.close();
|
||||
}
|
||||
|
|
@ -172,22 +159,7 @@ impl HandshakeEngine {
|
|||
receiver: &R,
|
||||
deadline: tokio::time::Instant,
|
||||
) -> Result<HandshakeResult, AcceptError> {
|
||||
self.accept_until_with_context(sender, receiver, deadline, AuthenticationContext::default())
|
||||
.await
|
||||
}
|
||||
|
||||
/// Run the crypto handshake until a deadline with transport-provided
|
||||
/// authentication scoping information.
|
||||
#[cfg(feature = "crypto")]
|
||||
pub async fn accept_until_with_context<S: HandshakeSender, R: HandshakeReceiver>(
|
||||
&self,
|
||||
sender: &S,
|
||||
receiver: &R,
|
||||
deadline: tokio::time::Instant,
|
||||
context: AuthenticationContext,
|
||||
) -> Result<HandshakeResult, AcceptError> {
|
||||
match tokio::time::timeout_at(deadline, self.accept_inner(sender, receiver, &context)).await
|
||||
{
|
||||
match tokio::time::timeout_at(deadline, self.accept_inner(sender, receiver)).await {
|
||||
Ok(result) => {
|
||||
if result.is_err() {
|
||||
sender.close();
|
||||
|
|
@ -214,15 +186,8 @@ impl HandshakeEngine {
|
|||
&self,
|
||||
sender: &S,
|
||||
receiver: &R,
|
||||
_authentication_context: &AuthenticationContext,
|
||||
) -> Result<HandshakeResult, AcceptError> {
|
||||
let mut first_msg = receiver.receive().await.map_err(AcceptError::Receive)?;
|
||||
tracing::debug!(
|
||||
message_type = ?first_msg.get_type(),
|
||||
version = ?first_msg.get_str(DataType::Version),
|
||||
client_id = ?first_msg.get_data(DataType::Id),
|
||||
"received MTP opening message"
|
||||
);
|
||||
|
||||
let version_str = match first_msg.get_data(DataType::Version) {
|
||||
Some(DataValue::Str(s)) => s.clone(),
|
||||
|
|
@ -277,11 +242,6 @@ impl HandshakeEngine {
|
|||
return Err(AcceptError::UnsupportedVersion(client_version));
|
||||
}
|
||||
};
|
||||
tracing::debug!(
|
||||
client_version = %client_version,
|
||||
negotiated_version = %negotiated,
|
||||
"MTP protocol version negotiated"
|
||||
);
|
||||
|
||||
let codec = VersionedCodec::for_version(self.registry.clone(), negotiated.clone())
|
||||
.ok_or_else(|| AcceptError::UnsupportedVersion(negotiated.clone()))?;
|
||||
|
|
@ -297,48 +257,6 @@ impl HandshakeEngine {
|
|||
|
||||
#[cfg(feature = "crypto")]
|
||||
{
|
||||
let claimed_client_id = match first_msg.get_data(DataType::Id) {
|
||||
Some(DataValue::UnsignedNumber(value)) => u64::try_from(*value).ok(),
|
||||
_ => None,
|
||||
};
|
||||
let registration = Some(first_msg.get_type())
|
||||
== CommunicationType::Register.try_to_id(codec.type_map());
|
||||
let authentication_requested = matches!(
|
||||
self.config.authentication_policy,
|
||||
crate::config::AuthenticationPolicy::ForceAuthentication
|
||||
) || registration
|
||||
|| first_msg.get_data(DataType::PublicKeys).is_some()
|
||||
|| claimed_client_id.is_some_and(|client_id| client_id != 0);
|
||||
tracing::info!(
|
||||
claimed_client_id = ?claimed_client_id,
|
||||
registration,
|
||||
authentication_requested,
|
||||
"classified MTP opening authentication mode"
|
||||
);
|
||||
if authentication_requested {
|
||||
let attempt = crate::config::AuthenticationAttempt {
|
||||
peer_network_identity: _authentication_context.peer_network_identity.clone(),
|
||||
connection_id: _authentication_context.connection_id,
|
||||
claimed_client_id,
|
||||
registration,
|
||||
};
|
||||
match self.config.auth_limiter.allow(&attempt) {
|
||||
Ok(true) => {}
|
||||
Ok(false) | Err(_) => {
|
||||
let error =
|
||||
AcceptError::AuthenticationFailed("authentication rejected".into());
|
||||
send_rejection_generic(
|
||||
sender,
|
||||
RejectionReason::RateLimited,
|
||||
Some(codec.type_map()),
|
||||
)
|
||||
.await;
|
||||
sender.close();
|
||||
return Err(error);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
match self.config.authentication_policy {
|
||||
crate::config::AuthenticationPolicy::ForceAuthentication => {
|
||||
self.force_auth_handshake(
|
||||
|
|
@ -490,14 +408,7 @@ impl HandshakeEngine {
|
|||
return Err(error);
|
||||
}
|
||||
};
|
||||
let pk_bytes = match bundle.try_as_bytes() {
|
||||
Ok(bytes) => bytes,
|
||||
Err(error) => {
|
||||
let error = AcceptError::AuthenticationFailed(error.to_string());
|
||||
reject_error_generic(sender, &error, tm).await;
|
||||
return Err(error);
|
||||
}
|
||||
};
|
||||
let pk_bytes = bundle.as_bytes();
|
||||
return self
|
||||
.complete_auth_handshake(
|
||||
sender,
|
||||
|
|
@ -542,28 +453,6 @@ impl HandshakeEngine {
|
|||
// Unknown or zero ID: an Identification carrying PublicKeys is an
|
||||
// explicit authentication attempt, not a guest connection.
|
||||
if first_msg.get_data(DataType::PublicKeys).is_some() {
|
||||
if self.config.conceal_authentication_identities && cid > 0 {
|
||||
/* Keep an unknown authenticated ID on the same
|
||||
challenge/proof path as a known ID. The fixed host
|
||||
identity makes the eventual proof fail without
|
||||
disclosing whether the lookup succeeded. */
|
||||
return self
|
||||
.complete_auth_handshake(
|
||||
sender,
|
||||
receiver,
|
||||
Flow::Login {
|
||||
id: cid,
|
||||
bundle: self.config.host_keyring.public_key_bundle(),
|
||||
},
|
||||
CommunicationType::IdentificationResponse,
|
||||
&negotiated,
|
||||
&codec,
|
||||
description,
|
||||
version_str,
|
||||
client_version,
|
||||
)
|
||||
.await;
|
||||
}
|
||||
let error = AcceptError::AuthenticationFailed(
|
||||
"unknown authenticated client identity".into(),
|
||||
);
|
||||
|
|
@ -572,7 +461,6 @@ impl HandshakeEngine {
|
|||
}
|
||||
|
||||
// Unknown or zero ID: fall back to guest
|
||||
tracing::info!("allocating MTP guest identity");
|
||||
let guest_id_lease = match self.assign_guest_id().await {
|
||||
Ok(lease) => lease,
|
||||
Err(error) => {
|
||||
|
|
@ -581,7 +469,6 @@ impl HandshakeEngine {
|
|||
}
|
||||
};
|
||||
let guest_id = guest_id_lease.id;
|
||||
tracing::info!(guest_id, "allocated MTP guest identity");
|
||||
send_accepted_generic(sender, &negotiated, tm, Some(guest_id))
|
||||
.await
|
||||
.map_err(AcceptError::Send)?;
|
||||
|
|
@ -638,14 +525,6 @@ impl HandshakeEngine {
|
|||
let bundle = match (self.config.get_existing_client)(cid, description.clone()).await {
|
||||
Some(b) => b,
|
||||
None => {
|
||||
if self.config.conceal_authentication_identities {
|
||||
// Use a valid fixed-cost dummy identity so an unknown
|
||||
// client follows the same challenge/proof sequence as
|
||||
// a registered client. The host public bundle is
|
||||
// already public and the peer cannot produce its
|
||||
// private-key proof.
|
||||
self.config.host_keyring.public_key_bundle()
|
||||
} else {
|
||||
let rejection = CommunicationValue::new_with_type_map(
|
||||
CommunicationType::IdentificationResponse,
|
||||
tm,
|
||||
|
|
@ -661,7 +540,6 @@ impl HandshakeEngine {
|
|||
"unknown client id".into(),
|
||||
));
|
||||
}
|
||||
}
|
||||
};
|
||||
(
|
||||
Flow::Login { id: cid, bundle },
|
||||
|
|
@ -675,14 +553,7 @@ impl HandshakeEngine {
|
|||
return Err(error);
|
||||
}
|
||||
};
|
||||
let pk_bytes = match bundle.try_as_bytes() {
|
||||
Ok(bytes) => bytes,
|
||||
Err(error) => {
|
||||
let error = AcceptError::AuthenticationFailed(error.to_string());
|
||||
reject_error_generic(sender, &error, tm).await;
|
||||
return Err(error);
|
||||
}
|
||||
};
|
||||
let pk_bytes = bundle.as_bytes();
|
||||
(
|
||||
Flow::Register { bundle, pk_bytes },
|
||||
CommunicationType::RegisterResponse,
|
||||
|
|
@ -916,14 +787,7 @@ impl HandshakeEngine {
|
|||
Flow::Login { id, bundle } => (id, bundle),
|
||||
Flow::Register { bundle, .. } => {
|
||||
let _registration_guard = self.config.registration_lock.lock().await;
|
||||
let identity = match bundle.try_as_bytes() {
|
||||
Ok(identity) => identity,
|
||||
Err(error) => {
|
||||
let error = AcceptError::AuthenticationFailed(error.to_string());
|
||||
reject_error_generic(sender, &error, tm).await;
|
||||
return Err(error);
|
||||
}
|
||||
};
|
||||
let identity = bundle.as_bytes();
|
||||
let cached_id = self
|
||||
.config
|
||||
.registration_ids
|
||||
|
|
@ -1124,12 +988,6 @@ async fn send_rejection_generic<S: HandshakeSender>(
|
|||
.add_typed_default(DataType::Connected, DataValue::BoolFalse)
|
||||
.add_typed_default(DataType::ErrorMessage, DataValue::Str(reason.to_string())),
|
||||
};
|
||||
tracing::debug!(
|
||||
reason = %reason,
|
||||
response_type = ?response.get_type(),
|
||||
has_version = response.get_data(DataType::Version).is_some(),
|
||||
"sending MTP handshake rejection"
|
||||
);
|
||||
let _ = sender.send(&response).await;
|
||||
}
|
||||
|
||||
|
|
@ -1163,11 +1021,6 @@ async fn send_accepted_generic<S: HandshakeSender>(
|
|||
if let Some(id) = assigned_id {
|
||||
response = response.add_typed_default(DataType::Id, DataValue::UnsignedNumber(id as u128));
|
||||
}
|
||||
tracing::debug!(
|
||||
version = %version,
|
||||
assigned_id = ?assigned_id,
|
||||
"sending accepted MTP handshake response"
|
||||
);
|
||||
sender.send(&response).await?;
|
||||
sender.finish_stream().await
|
||||
}
|
||||
|
|
|
|||
|
|
@ -11,7 +11,7 @@ use std::time::Instant;
|
|||
#[cfg(feature = "pipes")]
|
||||
use tokio::sync::mpsc;
|
||||
|
||||
use crate::config::{AuthenticationContext, HostConfig};
|
||||
use crate::config::HostConfig;
|
||||
use crate::connection::MTPConnection;
|
||||
use crate::engine::HandshakeEngine;
|
||||
use crate::error::AcceptError;
|
||||
|
|
@ -135,16 +135,7 @@ impl HandshakeContext {
|
|||
receiver: Receiver,
|
||||
) -> Result<Option<MTPConnection>, AcceptError> {
|
||||
let engine = HandshakeEngine::new(self.registry.clone(), self.config.clone());
|
||||
let authentication_context = AuthenticationContext {
|
||||
peer_network_identity: sender
|
||||
.handle()
|
||||
.remote_addr()
|
||||
.map(|address| address.to_string()),
|
||||
connection_id: sender.handle().connection_id(),
|
||||
};
|
||||
let result = engine
|
||||
.accept_with_context(&sender, &receiver, authentication_context)
|
||||
.await?;
|
||||
let result = engine.accept(&sender, &receiver).await?;
|
||||
#[cfg(feature = "crypto")]
|
||||
{
|
||||
Ok(Some(self.connection_from_handshake_result(
|
||||
|
|
@ -180,13 +171,12 @@ impl HandshakeContext {
|
|||
receiver.respond_to_pings(sender.clone());
|
||||
}
|
||||
|
||||
let receiver_queue_capacity = self.config.policy.receiver_queue_capacity.max(1);
|
||||
let (app_tx, app_rx) = mpsc::channel(receiver_queue_capacity);
|
||||
let (pipe_req_tx, pipe_req_rx) = mpsc::channel(receiver_queue_capacity);
|
||||
let (app_tx, app_rx) = mpsc::channel(self.config.policy.receiver_queue_capacity);
|
||||
let (pipe_req_tx, pipe_req_rx) =
|
||||
mpsc::channel(self.config.policy.receiver_queue_capacity);
|
||||
|
||||
let dispatcher = Arc::new(PipeDispatcher {
|
||||
pending_creations: std::sync::Mutex::new(std::collections::HashMap::new()),
|
||||
expired_creations: std::sync::Mutex::new(std::collections::HashMap::new()),
|
||||
pending_creations: tokio::sync::Mutex::new(std::collections::HashMap::new()),
|
||||
pending_pipes: tokio::sync::Mutex::new(std::collections::HashMap::new()),
|
||||
policy: Arc::new(self.config.policy),
|
||||
type_map: type_map.clone(),
|
||||
|
|
@ -270,13 +260,12 @@ impl HandshakeContext {
|
|||
receiver.respond_to_pings(sender.clone());
|
||||
}
|
||||
|
||||
let receiver_queue_capacity = self.config.policy.receiver_queue_capacity.max(1);
|
||||
let (app_tx, app_rx) = mpsc::channel(receiver_queue_capacity);
|
||||
let (pipe_req_tx, pipe_req_rx) = mpsc::channel(receiver_queue_capacity);
|
||||
let (app_tx, app_rx) = mpsc::channel(self.config.policy.receiver_queue_capacity);
|
||||
let (pipe_req_tx, pipe_req_rx) =
|
||||
mpsc::channel(self.config.policy.receiver_queue_capacity);
|
||||
|
||||
let dispatcher = Arc::new(PipeDispatcher {
|
||||
pending_creations: std::sync::Mutex::new(std::collections::HashMap::new()),
|
||||
expired_creations: std::sync::Mutex::new(std::collections::HashMap::new()),
|
||||
pending_creations: tokio::sync::Mutex::new(std::collections::HashMap::new()),
|
||||
pending_pipes: tokio::sync::Mutex::new(std::collections::HashMap::new()),
|
||||
policy: Arc::new(self.config.policy),
|
||||
type_map,
|
||||
|
|
|
|||
|
|
@ -29,9 +29,8 @@ pub use mtp_codec::registry::Registry;
|
|||
|
||||
#[cfg(feature = "crypto")]
|
||||
pub use config::{
|
||||
AuthenticationAttempt, AuthenticationAttemptLimiter, AuthenticationContext,
|
||||
AuthenticationLimitError, AuthenticationPolicy, CompleteRegister, FindRegisteredClient,
|
||||
GetExistingClient, GuestIdGenerator, InMemoryAuthenticationAttemptLimiter,
|
||||
AuthenticationPolicy, CompleteRegister, FindRegisteredClient, GetExistingClient,
|
||||
GuestIdGenerator,
|
||||
};
|
||||
#[cfg(feature = "crypto")]
|
||||
pub use error::AuthState;
|
||||
|
|
|
|||
266
host/src/pipe.rs
266
host/src/pipe.rs
|
|
@ -3,7 +3,6 @@ use mtp_common::{CommunicationError, PipeError};
|
|||
use mtp_transport::{PipeReader, PipeWriter, Policy, TransportEvent};
|
||||
use std::collections::HashMap;
|
||||
use std::sync::Arc;
|
||||
use std::sync::Mutex as StdMutex;
|
||||
use tokio::sync::{Mutex, mpsc};
|
||||
|
||||
/// The sender operations needed by the transport-independent pipe protocol.
|
||||
|
|
@ -27,10 +26,6 @@ pub trait PipeReceiver<P>: Clone + Send + Sync + 'static
|
|||
where
|
||||
P: tokio::io::AsyncRead + Send + Unpin + 'static,
|
||||
{
|
||||
fn expect_pipe(&self, pipe_id: u32) -> Result<(), CommunicationError>;
|
||||
|
||||
fn cancel_expected_pipe(&self, pipe_id: u32);
|
||||
|
||||
fn receive_pipe_event(
|
||||
&self,
|
||||
) -> impl std::future::Future<Output = Result<TransportEvent<P>, CommunicationError>> + Send;
|
||||
|
|
@ -56,14 +51,6 @@ impl PipeSender for mtp_transport::Sender {
|
|||
}
|
||||
|
||||
impl PipeReceiver<wtransport::RecvStream> for mtp_transport::Receiver {
|
||||
fn expect_pipe(&self, pipe_id: u32) -> Result<(), CommunicationError> {
|
||||
self.expect_pipe(pipe_id)
|
||||
}
|
||||
|
||||
fn cancel_expected_pipe(&self, pipe_id: u32) {
|
||||
self.cancel_expected_pipe(pipe_id);
|
||||
}
|
||||
|
||||
async fn receive_pipe_event(
|
||||
&self,
|
||||
) -> Result<TransportEvent<wtransport::RecvStream>, CommunicationError> {
|
||||
|
|
@ -99,14 +86,6 @@ where
|
|||
C: mtp_transport::TransportConnection,
|
||||
C::RecvStream: tokio::io::AsyncRead + Send + Unpin + 'static,
|
||||
{
|
||||
fn expect_pipe(&self, pipe_id: u32) -> Result<(), CommunicationError> {
|
||||
self.expect_pipe(pipe_id)
|
||||
}
|
||||
|
||||
fn cancel_expected_pipe(&self, pipe_id: u32) {
|
||||
self.cancel_expected_pipe(pipe_id);
|
||||
}
|
||||
|
||||
async fn receive_pipe_event(
|
||||
&self,
|
||||
) -> Result<TransportEvent<C::RecvStream>, CommunicationError> {
|
||||
|
|
@ -114,20 +93,14 @@ where
|
|||
}
|
||||
}
|
||||
|
||||
pub struct PipeHandle<S: PipeSender, P = wtransport::RecvStream> {
|
||||
pub struct PipeHandle<S: PipeSender> {
|
||||
pub(crate) pipe_id: u32,
|
||||
pub(crate) description: String,
|
||||
pub(crate) sender: S,
|
||||
pub(crate) response_rx: tokio::sync::oneshot::Receiver<Result<bool, PipeError>>,
|
||||
pub(crate) dispatcher: Arc<PipeDispatcher<P>>,
|
||||
pub(crate) token: Arc<()>,
|
||||
}
|
||||
|
||||
impl<S, P> PipeHandle<S, P>
|
||||
where
|
||||
S: PipeSender,
|
||||
P: tokio::io::AsyncRead + Send + Unpin + 'static,
|
||||
{
|
||||
impl<S: PipeSender> PipeHandle<S> {
|
||||
pub fn pipe_id(&self) -> u32 {
|
||||
self.pipe_id
|
||||
}
|
||||
|
|
@ -136,96 +109,31 @@ where
|
|||
&self.description
|
||||
}
|
||||
|
||||
pub async fn wait(mut self) -> Result<Option<PipeWriter<S::Writer>>, PipeError> {
|
||||
let response =
|
||||
tokio::time::timeout(self.dispatcher.policy.read_timeout, &mut self.response_rx).await;
|
||||
match response {
|
||||
Ok(Ok(Ok(true))) => self
|
||||
pub async fn wait(self) -> Result<Option<PipeWriter<S::Writer>>, PipeError> {
|
||||
match self.response_rx.await {
|
||||
Ok(Ok(true)) => self
|
||||
.sender
|
||||
.open_pipe_stream(self.pipe_id, &self.description)
|
||||
.await
|
||||
.map(Some)
|
||||
.map_err(PipeError::from),
|
||||
Ok(Ok(Ok(false))) => Ok(None),
|
||||
Ok(Ok(Err(error))) => {
|
||||
expire_pending_creation(&self.dispatcher, self.pipe_id, &self.token);
|
||||
Err(error)
|
||||
}
|
||||
Ok(Err(_)) => {
|
||||
expire_pending_creation(&self.dispatcher, self.pipe_id, &self.token);
|
||||
Err(PipeError::StreamClosed)
|
||||
}
|
||||
Err(_) => {
|
||||
expire_pending_creation(&self.dispatcher, self.pipe_id, &self.token);
|
||||
Err(PipeError::HandshakeTimeout)
|
||||
}
|
||||
Ok(Ok(false)) => Ok(None),
|
||||
Ok(Err(error)) => Err(error),
|
||||
Err(_) => Err(PipeError::StreamClosed),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl<S, P> Drop for PipeHandle<S, P>
|
||||
where
|
||||
S: PipeSender,
|
||||
{
|
||||
fn drop(&mut self) {
|
||||
expire_pending_creation(&self.dispatcher, self.pipe_id, &self.token);
|
||||
}
|
||||
}
|
||||
|
||||
pub struct PipeRequest<S, R, P> {
|
||||
pub struct PipeRequest<S, P> {
|
||||
pub(crate) pipe_id: u32,
|
||||
pub(crate) description: String,
|
||||
pub(crate) sender: S,
|
||||
pub(crate) receiver: R,
|
||||
pub(crate) dispatcher: Arc<PipeDispatcher<P>>,
|
||||
}
|
||||
|
||||
struct ExpectedPipeGuard<R, P>
|
||||
where
|
||||
R: PipeReceiver<P>,
|
||||
P: tokio::io::AsyncRead + Send + Unpin + 'static,
|
||||
{
|
||||
receiver: R,
|
||||
pipe_id: u32,
|
||||
armed: bool,
|
||||
_stream: std::marker::PhantomData<P>,
|
||||
}
|
||||
|
||||
impl<R, P> ExpectedPipeGuard<R, P>
|
||||
where
|
||||
R: PipeReceiver<P>,
|
||||
P: tokio::io::AsyncRead + Send + Unpin + 'static,
|
||||
{
|
||||
fn new(receiver: R, pipe_id: u32) -> Self {
|
||||
Self {
|
||||
receiver,
|
||||
pipe_id,
|
||||
armed: true,
|
||||
_stream: std::marker::PhantomData,
|
||||
}
|
||||
}
|
||||
|
||||
fn disarm(&mut self) {
|
||||
self.armed = false;
|
||||
}
|
||||
}
|
||||
|
||||
impl<R, P> Drop for ExpectedPipeGuard<R, P>
|
||||
where
|
||||
R: PipeReceiver<P>,
|
||||
P: tokio::io::AsyncRead + Send + Unpin + 'static,
|
||||
{
|
||||
fn drop(&mut self) {
|
||||
if self.armed {
|
||||
self.receiver.cancel_expected_pipe(self.pipe_id);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl<S, R, P> PipeRequest<S, R, P>
|
||||
impl<S, P> PipeRequest<S, P>
|
||||
where
|
||||
S: PipeSender,
|
||||
R: PipeReceiver<P>,
|
||||
P: tokio::io::AsyncRead + Send + Unpin + 'static,
|
||||
{
|
||||
pub fn id(&self) -> u32 {
|
||||
|
|
@ -237,10 +145,6 @@ where
|
|||
}
|
||||
|
||||
pub async fn accept(self) -> Result<PipeReader<P>, PipeError> {
|
||||
self.receiver
|
||||
.expect_pipe(self.pipe_id)
|
||||
.map_err(PipeError::from)?;
|
||||
let mut expected_pipe = ExpectedPipeGuard::<R, P>::new(self.receiver.clone(), self.pipe_id);
|
||||
let (pipe_tx, pipe_rx) = tokio::sync::oneshot::channel();
|
||||
self.dispatcher
|
||||
.pending_pipes
|
||||
|
|
@ -264,10 +168,7 @@ where
|
|||
}
|
||||
|
||||
match tokio::time::timeout(self.dispatcher.policy.read_timeout, pipe_rx).await {
|
||||
Ok(Ok(reader)) => {
|
||||
expected_pipe.disarm();
|
||||
Ok(reader)
|
||||
}
|
||||
Ok(Ok(reader)) => Ok(reader),
|
||||
Ok(Err(_)) => {
|
||||
self.dispatcher
|
||||
.pending_pipes
|
||||
|
|
@ -302,137 +203,18 @@ where
|
|||
}
|
||||
|
||||
pub(crate) struct PipeDispatcher<P> {
|
||||
pub(crate) pending_creations: StdMutex<HashMap<u32, PendingCreation>>,
|
||||
pub(crate) expired_creations: StdMutex<HashMap<u32, tokio::time::Instant>>,
|
||||
pub(crate) pending_creations:
|
||||
Mutex<HashMap<u32, tokio::sync::oneshot::Sender<Result<bool, PipeError>>>>,
|
||||
pub(crate) pending_pipes: Mutex<HashMap<u32, tokio::sync::oneshot::Sender<PipeReader<P>>>>,
|
||||
pub(crate) policy: Arc<Policy>,
|
||||
pub(crate) type_map: TypeMap,
|
||||
}
|
||||
|
||||
pub(crate) struct PendingCreation {
|
||||
pub(crate) token: Arc<()>,
|
||||
pub(crate) sender: tokio::sync::oneshot::Sender<Result<bool, PipeError>>,
|
||||
}
|
||||
|
||||
pub(crate) struct PendingCreationGuard<P> {
|
||||
dispatcher: Arc<PipeDispatcher<P>>,
|
||||
pipe_id: u32,
|
||||
token: Arc<()>,
|
||||
armed: bool,
|
||||
}
|
||||
|
||||
impl<P> PendingCreationGuard<P> {
|
||||
pub(crate) fn new(dispatcher: Arc<PipeDispatcher<P>>, pipe_id: u32, token: Arc<()>) -> Self {
|
||||
Self {
|
||||
dispatcher,
|
||||
pipe_id,
|
||||
token,
|
||||
armed: true,
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) fn disarm(&mut self) {
|
||||
self.armed = false;
|
||||
}
|
||||
}
|
||||
|
||||
impl<P> Drop for PendingCreationGuard<P> {
|
||||
fn drop(&mut self) {
|
||||
if self.armed {
|
||||
expire_pending_creation(&self.dispatcher, self.pipe_id, &self.token);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
const EXPIRED_CREATION_TOMBSTONE_TTL: tokio::time::Duration = tokio::time::Duration::from_secs(60);
|
||||
const MAX_EXPIRED_CREATION_TOMBSTONES: usize = 1024;
|
||||
|
||||
pub(crate) fn expire_pending_creation<P>(
|
||||
dispatcher: &PipeDispatcher<P>,
|
||||
pipe_id: u32,
|
||||
token: &Arc<()>,
|
||||
) {
|
||||
let removed = dispatcher
|
||||
.pending_creations
|
||||
.lock()
|
||||
.ok()
|
||||
.and_then(|mut pending| {
|
||||
if pending
|
||||
.get(&pipe_id)
|
||||
.is_some_and(|entry| Arc::ptr_eq(&entry.token, token))
|
||||
{
|
||||
pending.remove(&pipe_id);
|
||||
Some(())
|
||||
} else {
|
||||
None
|
||||
}
|
||||
});
|
||||
if removed.is_none() {
|
||||
return;
|
||||
}
|
||||
let Ok(mut expired) = dispatcher.expired_creations.lock() else {
|
||||
return;
|
||||
};
|
||||
let now = tokio::time::Instant::now();
|
||||
expired.retain(|_, expires_at| *expires_at > now);
|
||||
if expired.len() >= MAX_EXPIRED_CREATION_TOMBSTONES
|
||||
&& let Some(oldest) = expired
|
||||
.iter()
|
||||
.min_by_key(|(_, expires_at)| **expires_at)
|
||||
.map(|(id, _)| *id)
|
||||
{
|
||||
expired.remove(&oldest);
|
||||
}
|
||||
expired.insert(pipe_id, now + EXPIRED_CREATION_TOMBSTONE_TTL);
|
||||
}
|
||||
|
||||
fn consume_expired_creation<P>(dispatcher: &PipeDispatcher<P>, pipe_id: u32) -> bool {
|
||||
let Ok(mut expired) = dispatcher.expired_creations.lock() else {
|
||||
return false;
|
||||
};
|
||||
let now = tokio::time::Instant::now();
|
||||
expired.retain(|_, expires_at| *expires_at > now);
|
||||
expired.remove(&pipe_id).is_some()
|
||||
}
|
||||
|
||||
pub(crate) fn is_expired_creation<P>(dispatcher: &PipeDispatcher<P>, pipe_id: u32) -> bool {
|
||||
let Ok(mut expired) = dispatcher.expired_creations.lock() else {
|
||||
return true;
|
||||
};
|
||||
let now = tokio::time::Instant::now();
|
||||
expired.retain(|_, expires_at| *expires_at > now);
|
||||
expired.contains_key(&pipe_id)
|
||||
}
|
||||
|
||||
pub(crate) fn fail_pending_creations<P>(
|
||||
dispatcher: &PipeDispatcher<P>,
|
||||
error: &CommunicationError,
|
||||
) {
|
||||
let pending = dispatcher
|
||||
.pending_creations
|
||||
.lock()
|
||||
.ok()
|
||||
.map(|mut pending| std::mem::take(&mut *pending));
|
||||
if let Some(pending) = pending {
|
||||
let error = PipeError::from(error.clone());
|
||||
for (_, pending) in pending {
|
||||
let _ = pending.sender.send(Err(error.clone()));
|
||||
}
|
||||
}
|
||||
if let Ok(mut expired) = dispatcher.expired_creations.lock() {
|
||||
expired.clear();
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) async fn fail_pending_pipes<P>(dispatcher: &PipeDispatcher<P>) {
|
||||
dispatcher.pending_pipes.lock().await.clear();
|
||||
}
|
||||
|
||||
pub(crate) async fn run_dispatcher<S, R, P>(
|
||||
receiver: R,
|
||||
sender: S,
|
||||
app_tx: mpsc::Sender<Result<CommunicationValue, CommunicationError>>,
|
||||
pipe_req_tx: mpsc::Sender<PipeRequest<S, R, P>>,
|
||||
pipe_req_tx: mpsc::Sender<PipeRequest<S, P>>,
|
||||
dispatcher: Arc<PipeDispatcher<P>>,
|
||||
) where
|
||||
S: PipeSender,
|
||||
|
|
@ -459,7 +241,6 @@ pub(crate) async fn run_dispatcher<S, R, P>(
|
|||
.unwrap_or("")
|
||||
.to_owned(),
|
||||
sender: sender.clone(),
|
||||
receiver: receiver.clone(),
|
||||
dispatcher: dispatcher.clone(),
|
||||
};
|
||||
let _ = pipe_req_tx.send(request).await;
|
||||
|
|
@ -475,17 +256,10 @@ pub(crate) async fn run_dispatcher<S, R, P>(
|
|||
}
|
||||
continue;
|
||||
};
|
||||
let pending = dispatcher
|
||||
.pending_creations
|
||||
.lock()
|
||||
.ok()
|
||||
.and_then(|mut pending| pending.remove(&pipe_id));
|
||||
if let Some(entry) = pending {
|
||||
let _ = entry
|
||||
.sender
|
||||
.send(Ok(message.get_bool(DataType::Accepted).unwrap_or(false)));
|
||||
} else if consume_expired_creation(&dispatcher, pipe_id) {
|
||||
tracing::debug!(pipe_id, "ignored late pipe creation response");
|
||||
let mut pending = dispatcher.pending_creations.lock().await;
|
||||
if let Some(reply) = pending.remove(&pipe_id) {
|
||||
let _ =
|
||||
reply.send(Ok(message.get_bool(DataType::Accepted).unwrap_or(false)));
|
||||
}
|
||||
continue;
|
||||
}
|
||||
|
|
@ -518,18 +292,14 @@ pub(crate) async fn run_dispatcher<S, R, P>(
|
|||
pipe_id,
|
||||
description: reader.description().to_owned(),
|
||||
sender: sender.clone(),
|
||||
receiver: receiver.clone(),
|
||||
dispatcher: dispatcher.clone(),
|
||||
};
|
||||
let _ = pipe_req_tx.send(request).await;
|
||||
}
|
||||
Err(error) => {
|
||||
fail_pending_creations(&dispatcher, &error);
|
||||
fail_pending_pipes(&dispatcher).await;
|
||||
if app_tx.send(Err(error)).await.is_err() {
|
||||
break;
|
||||
}
|
||||
break;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -143,11 +143,6 @@ pub(crate) async fn run_driver(
|
|||
return;
|
||||
}
|
||||
};
|
||||
tracing::debug!(
|
||||
remote = %remote_addr,
|
||||
session_id = ?session.session_id(),
|
||||
"accepted WebTransport MTP session"
|
||||
);
|
||||
tokio::spawn(run_session_requests(
|
||||
session.clone(),
|
||||
router.clone(),
|
||||
|
|
|
|||
|
|
@ -32,8 +32,6 @@ pub struct H3TransportSender {
|
|||
|
||||
pub struct H3TransportReceiver {
|
||||
stream: H3RecvStream,
|
||||
quinn: quinn::Connection,
|
||||
read_exact_calls: u64,
|
||||
}
|
||||
|
||||
impl H3TransportConnection {
|
||||
|
|
@ -44,11 +42,6 @@ impl H3TransportConnection {
|
|||
pub(crate) fn remote_addr(&self) -> std::net::SocketAddr {
|
||||
self.quinn.remote_address()
|
||||
}
|
||||
|
||||
#[cfg(feature = "crypto")]
|
||||
pub(crate) fn connection_id(&self) -> u64 {
|
||||
self.quinn.stable_id() as u64
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait::async_trait]
|
||||
|
|
@ -57,14 +50,14 @@ impl TransportSendStream for H3TransportSender {
|
|||
self.stream
|
||||
.write_all(buf)
|
||||
.await
|
||||
.map_err(|_| CommunicationError::DeliveryUnknown)?;
|
||||
.map_err(|_| CommunicationError::StreamError)?;
|
||||
// Control/authentication frames use a persistent stream. h3 keeps
|
||||
// those writes buffered until flushed; without this the peer can wait
|
||||
// for the challenge while the server waits for its proof.
|
||||
self.stream
|
||||
.flush()
|
||||
.await
|
||||
.map_err(|_| CommunicationError::DeliveryUnknown)
|
||||
.map_err(|_| CommunicationError::StreamError)
|
||||
}
|
||||
|
||||
async fn finish(&mut self) -> Result<(), CommunicationError> {
|
||||
|
|
@ -73,53 +66,23 @@ impl TransportSendStream for H3TransportSender {
|
|||
.await
|
||||
.map_err(|_| CommunicationError::StreamError)
|
||||
}
|
||||
|
||||
fn reset(&mut self, code: u32) -> Result<(), CommunicationError> {
|
||||
h3::quic::SendStream::reset(&mut self.stream, code as u64);
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait::async_trait]
|
||||
impl TransportRecvStream for H3TransportReceiver {
|
||||
async fn read_exact(&mut self, buf: &mut [u8]) -> Result<(), CommunicationError> {
|
||||
let first_read = self.read_exact_calls == 0;
|
||||
self.read_exact_calls += 1;
|
||||
self.stream
|
||||
.read_exact(buf)
|
||||
.await
|
||||
.map(|_| {
|
||||
if first_read {
|
||||
tracing::debug!(
|
||||
remote = %self.quinn.remote_address(),
|
||||
bytes = buf.len(),
|
||||
header = ?buf,
|
||||
"received first bytes from WebTransport MTP stream"
|
||||
);
|
||||
}
|
||||
})
|
||||
.map(|_| ())
|
||||
.map_err(|error| {
|
||||
if error.kind() == std::io::ErrorKind::UnexpectedEof
|
||||
|| self.quinn.close_reason().is_some()
|
||||
{
|
||||
/*
|
||||
* Reaching FIN, or losing the enclosing QUIC connection,
|
||||
* is a normal stream-closure path. Do not turn it into a
|
||||
* frame-header failure and close the connection again.
|
||||
*/
|
||||
if error.kind() == std::io::ErrorKind::UnexpectedEof {
|
||||
// Browser control frames are sent on one-frame uni streams.
|
||||
// Reaching FIN while looking for another frame is normal.
|
||||
return CommunicationError::StreamClosed;
|
||||
}
|
||||
error!(
|
||||
"[mtp-webserver] receive stream read_exact failed ({} bytes): {error}",
|
||||
buf.len()
|
||||
);
|
||||
tracing::warn!(
|
||||
remote = %self.quinn.remote_address(),
|
||||
first_read,
|
||||
len = buf.len(),
|
||||
%error,
|
||||
"WebTransport receive stream read_exact failed"
|
||||
);
|
||||
error!("[mtp-webserver] receive stream read_exact failed ({} bytes): {error}", buf.len());
|
||||
tracing::warn!(len = buf.len(), %error, "WebTransport receive stream read_exact failed");
|
||||
CommunicationError::StreamError
|
||||
})
|
||||
}
|
||||
|
|
@ -133,9 +96,6 @@ impl TransportRecvStream for H3TransportReceiver {
|
|||
Ok(Some(buf))
|
||||
}
|
||||
Err(error) => {
|
||||
if self.quinn.close_reason().is_some() {
|
||||
return Err(CommunicationError::StreamClosed);
|
||||
}
|
||||
error!(
|
||||
"[mtp-webserver] receive stream read failed (max {} bytes): {error}",
|
||||
max
|
||||
|
|
@ -145,11 +105,6 @@ impl TransportRecvStream for H3TransportReceiver {
|
|||
}
|
||||
}
|
||||
}
|
||||
|
||||
fn stop(mut self, code: u32) -> Result<(), CommunicationError> {
|
||||
h3::quic::RecvStream::stop_sending(&mut self.stream, code as u64);
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
|
||||
impl tokio::io::AsyncWrite for H3TransportSender {
|
||||
|
|
@ -207,27 +162,10 @@ impl TransportConnection for H3TransportConnection {
|
|||
loop {
|
||||
match self.session.accept_uni().await {
|
||||
Ok(Some((id, stream))) if id == self.session.session_id() => {
|
||||
let stream_id = h3::quic::RecvStream::recv_id(&stream);
|
||||
tracing::debug!(
|
||||
remote = %self.quinn.remote_address(),
|
||||
session_id = ?self.session.session_id(),
|
||||
stream_id = ?stream_id,
|
||||
"accepted WebTransport MTP receive stream"
|
||||
);
|
||||
return Ok(H3TransportReceiver {
|
||||
stream,
|
||||
quinn: self.quinn.clone(),
|
||||
read_exact_calls: 0,
|
||||
});
|
||||
return Ok(H3TransportReceiver { stream });
|
||||
}
|
||||
Ok(Some((stream_session_id, _stream))) => {
|
||||
Ok(Some(_)) => {
|
||||
consecutive_errors = 0;
|
||||
tracing::debug!(
|
||||
remote = %self.quinn.remote_address(),
|
||||
session_id = ?self.session.session_id(),
|
||||
stream_session_id = ?stream_session_id,
|
||||
"ignored WebTransport receive stream belonging to another session"
|
||||
);
|
||||
continue;
|
||||
}
|
||||
Ok(None) => return Err(CommunicationError::StreamClosed),
|
||||
|
|
@ -345,8 +283,6 @@ async fn accept_web_connection_inner(
|
|||
let max_message_size = policy.max_message_size;
|
||||
let transport = H3TransportConnection::new(session, quinn);
|
||||
let remote_addr = transport.remote_addr();
|
||||
#[cfg(feature = "crypto")]
|
||||
let connection_id = transport.connection_id();
|
||||
let policy = Arc::new(policy);
|
||||
let sender = WebMtpSender::new(transport.clone(), policy.clone());
|
||||
let receiver = WebMtpReceiver::new(transport, policy.clone());
|
||||
|
|
@ -354,27 +290,12 @@ async fn accept_web_connection_inner(
|
|||
let engine = mtp_host::HandshakeEngine::new(Registry::builtin(), host_config);
|
||||
#[cfg(feature = "crypto")]
|
||||
let result = engine
|
||||
.accept_until_with_context(
|
||||
.accept_until(
|
||||
&sender,
|
||||
&receiver,
|
||||
deadline.expect("crypto WebTransport handshakes have a deadline"),
|
||||
mtp_host::AuthenticationContext {
|
||||
peer_network_identity: Some(remote_addr.to_string()),
|
||||
connection_id,
|
||||
},
|
||||
)
|
||||
.await;
|
||||
#[cfg(feature = "crypto")]
|
||||
if let Err(error) = &result {
|
||||
tracing::warn!(
|
||||
remote = %remote_addr,
|
||||
connection_id,
|
||||
%error,
|
||||
"WebTransport MTP handshake failed"
|
||||
);
|
||||
}
|
||||
#[cfg(feature = "crypto")]
|
||||
let result = result?;
|
||||
.await?;
|
||||
#[cfg(not(feature = "crypto"))]
|
||||
let result = engine.accept(&sender, &receiver).await?;
|
||||
|
||||
|
|
|
|||
|
|
@ -58,20 +58,19 @@
|
|||
"build:wasm": "MTP_TYPE_MAPS=$PWD/example-type-maps.yaml wasm-pack build wasm --target web --out-dir pkg --release && rm -f wasm/pkg/.gitignore",
|
||||
"build:ts": "rm -rf dist && tsc",
|
||||
"build": "pnpm run build:wasm && pnpm run build:ts",
|
||||
"pack": "pnpm run release:web",
|
||||
"release:web": "node create-web-release.mjs",
|
||||
"pack": "pnpm run clean && pnpm run build && pnpm pack",
|
||||
"build:all": "nix run .#build-all",
|
||||
"dup": "jscpd --pattern '**/*.{rs,ts}' --ignore 'target/**' --ignore 'wasm/pkg/**' --ignore '.git/**' --min-lines 8 --min-tokens 80 --threshold 4 --reporters console --noTips .",
|
||||
"test:e2e": "tsc && node test/e2ee.mjs",
|
||||
"test:secrets": "tsc && node --test --test-isolation=none test/encrypted-secret.mjs",
|
||||
"test:wasm-init": "tsc && node --test test/wasm-init.mjs",
|
||||
"test:types": "tsc -p tsconfig.type-tests.json --noEmit",
|
||||
"test:vite": "tsc && node test/vite-type-map.mjs",
|
||||
"test:boundary": "node --test test/package-boundary.mjs",
|
||||
"test": "pnpm run test:e2e && pnpm run test:secrets && pnpm run test:wasm-init && pnpm run test:types && pnpm run test:vite && pnpm run test:boundary"
|
||||
"test": "pnpm run test:e2e && pnpm run test:secrets && pnpm run test:types && pnpm run test:vite && pnpm run test:boundary"
|
||||
},
|
||||
"devDependencies": {
|
||||
"@types/node": "^26.0.1",
|
||||
"jscpd": "5.0.14",
|
||||
"jscpd": "4.2.5",
|
||||
"typescript": "^7.0.0"
|
||||
},
|
||||
"dependencies": {
|
||||
|
|
|
|||
935
pnpm-lock.yaml
generated
935
pnpm-lock.yaml
generated
File diff suppressed because it is too large
Load diff
2760
src/sdk/client.ts
2760
src/sdk/client.ts
File diff suppressed because it is too large
Load diff
959
src/sdk/codec.ts
959
src/sdk/codec.ts
|
|
@ -1,959 +0,0 @@
|
|||
import * as bindings from "mtp/raw";
|
||||
import type { MTPCommunicationType } from "../type-map/index";
|
||||
import { RESERVED_COMMUNICATION_TYPE_IDS } from "../type-map/reserved.js";
|
||||
import { utf8Encode } from "./utils.js";
|
||||
import { initWasmOnce } from "./wasm-init.js";
|
||||
import type {
|
||||
MTPBytesInput,
|
||||
MTPCodec,
|
||||
MTPCodecOptions,
|
||||
MTPDataValue,
|
||||
MTPDataValueInput,
|
||||
MTPEncodeLimits,
|
||||
MTPEncodedBytesInput,
|
||||
MTPKeyringKeys,
|
||||
MTPKeyMaterialInput,
|
||||
MTPReceiveLimits,
|
||||
MTPCrypto,
|
||||
MTPPublicKeyBundleKeys,
|
||||
MTPProtectedFrameInput,
|
||||
MTPProtectionSignatureSuite,
|
||||
ParsedFrame,
|
||||
} from "./client.js";
|
||||
|
||||
const checkedKeyringGenerator = (
|
||||
bindings as typeof bindings & {
|
||||
keyring_generate_checked?: () => Uint8Array;
|
||||
}
|
||||
).keyring_generate_checked;
|
||||
|
||||
export const crypto: MTPCrypto = {
|
||||
generateKeyring: () =>
|
||||
checkedKeyringGenerator?.() ?? bindings.keyring_generate(),
|
||||
generateEd25519: () => bindings.ed25519_generate(),
|
||||
keyringFromEd25519: (secretKey, publicKey) =>
|
||||
bindings.keyring_from_ed25519(secretKey, publicKey),
|
||||
verifyEd25519: (publicKey, message, signature) =>
|
||||
bindings.ed25519_verify(publicKey, message, signature),
|
||||
deriveEncryptionKey: (ikm, salt, context) =>
|
||||
bindings.wasm_derive_encryption_key(ikm, salt, context),
|
||||
hkdfExpand: (ikm, salt, info, len) =>
|
||||
bindings.wasm_hkdf_expand(ikm, salt, info, len),
|
||||
sha256: (data) => bindings.wasm_sha256(data),
|
||||
sha256Double: (data) => bindings.wasm_sha256_double(data),
|
||||
keyringToKeys: (keyring) => keyringToKeys(keyring),
|
||||
publicKeyBundleToKeys: (publicKeyBundle) =>
|
||||
publicKeyBundleToKeys(publicKeyBundle),
|
||||
|
||||
encrypt: async (key, input) => {
|
||||
const cipher = new bindings.WasmChaCha20Poly1305(key);
|
||||
try {
|
||||
return cipher.encrypt(input, new Uint8Array(0));
|
||||
} finally {
|
||||
cipher.free();
|
||||
}
|
||||
},
|
||||
|
||||
decrypt: async (key, input) => {
|
||||
const cipher = new bindings.WasmChaCha20Poly1305(key);
|
||||
try {
|
||||
return cipher.decrypt(input, new Uint8Array(0));
|
||||
} finally {
|
||||
cipher.free();
|
||||
}
|
||||
},
|
||||
|
||||
encryptText: async (key, plaintext) => {
|
||||
const cipher = new bindings.WasmChaCha20Poly1305(key);
|
||||
try {
|
||||
const ciphertext = cipher.encrypt(
|
||||
utf8Encode(plaintext),
|
||||
new Uint8Array(0),
|
||||
);
|
||||
return bytesToBase64(ciphertext);
|
||||
} finally {
|
||||
cipher.free();
|
||||
}
|
||||
},
|
||||
|
||||
decryptText: async (key, ciphertext) => {
|
||||
const cipher = new bindings.WasmChaCha20Poly1305(key);
|
||||
try {
|
||||
const decoded = base64ToBytes(ciphertext);
|
||||
const plaintext = cipher.decrypt(decoded, new Uint8Array(0));
|
||||
return utf8Decode(plaintext);
|
||||
} finally {
|
||||
cipher.free();
|
||||
}
|
||||
},
|
||||
|
||||
encapsulate: (otherPublicKey) =>
|
||||
bindings.wasm_kem_encapsulate(otherPublicKey),
|
||||
|
||||
decapsulate: (ownPrivateKey, ciphertext) =>
|
||||
bindings.wasm_kem_decapsulate(ownPrivateKey, ciphertext),
|
||||
};
|
||||
|
||||
export function encode(
|
||||
type: MTPCommunicationType,
|
||||
data: Record<string, unknown>,
|
||||
options?: MTPCodecOptions,
|
||||
): Uint8Array {
|
||||
const limits: MTPEncodeLimits = {
|
||||
maxDepth: MAX_DATA_VALUE_DEPTH,
|
||||
maxValues: MAX_DATA_VALUE_VALUES,
|
||||
maxOutputSize: 16 * 1024 * 1024,
|
||||
};
|
||||
const maxOutputSize = limits.maxOutputSize ?? 16 * 1024 * 1024;
|
||||
validateMTPDataValue(data as MTPDataValueInput, limits);
|
||||
const bounded = (
|
||||
bindings as typeof bindings & {
|
||||
build_frame_with_limits?: (
|
||||
type: string,
|
||||
data: Record<string, unknown>,
|
||||
options: MTPCodecOptions,
|
||||
limits: MTPEncodeLimits,
|
||||
) => Uint8Array;
|
||||
}
|
||||
).build_frame_with_limits;
|
||||
if (!bounded) {
|
||||
throw new Error("bounded WASM frame encoding is unavailable; rebuild mtp-wasm");
|
||||
}
|
||||
const frame = bounded(type, data, options ?? {}, limits);
|
||||
if (frame.length > maxOutputSize) {
|
||||
throw new RangeError("MTP frame encoded output limit exceeded");
|
||||
}
|
||||
return frame;
|
||||
}
|
||||
|
||||
export function decode(frame: MTPBytesInput): ParsedFrame {
|
||||
return bindings.parse_frame(bytesFrom(frame, "frame"));
|
||||
}
|
||||
|
||||
export function decodeWithLimits(
|
||||
frame: MTPBytesInput,
|
||||
limits: MTPReceiveLimits,
|
||||
): ParsedFrame {
|
||||
const parse = (
|
||||
bindings as typeof bindings & {
|
||||
parse_frame_with_limits?: (
|
||||
frame: Uint8Array,
|
||||
limits: MTPReceiveLimits,
|
||||
) => ParsedFrame;
|
||||
}
|
||||
).parse_frame_with_limits;
|
||||
if (!parse) {
|
||||
throw new Error(
|
||||
"configured receive limits require a rebuilt bounded WASM package",
|
||||
);
|
||||
}
|
||||
return parse(bytesFrom(frame, "frame"), limits);
|
||||
}
|
||||
|
||||
export function decodeDataValueWithLimits(
|
||||
value: MTPBytesInput,
|
||||
limits: MTPReceiveLimits,
|
||||
): MTPDataValue {
|
||||
const parse = (
|
||||
bindings as typeof bindings & {
|
||||
parse_data_value_with_limits?: (
|
||||
value: Uint8Array,
|
||||
limits: MTPReceiveLimits,
|
||||
) => MTPDataValue;
|
||||
}
|
||||
).parse_data_value_with_limits;
|
||||
if (!parse) {
|
||||
throw new Error(
|
||||
"configured receive limits require a rebuilt bounded WASM package",
|
||||
);
|
||||
}
|
||||
return parse(bytesFrom(value, "data value"), limits);
|
||||
}
|
||||
|
||||
export function format(frame: MTPBytesInput): string {
|
||||
return bindings.format_frame(bytesFrom(frame, "frame"));
|
||||
}
|
||||
|
||||
export const codec: MTPCodec = { encode, decode, format };
|
||||
|
||||
export function isBytes(value: unknown): value is MTPBytesInput {
|
||||
return value instanceof Uint8Array || Array.isArray(value);
|
||||
}
|
||||
|
||||
export function bytesFrom(value: MTPBytesInput, name: string): Uint8Array {
|
||||
if (value instanceof Uint8Array) return value.slice();
|
||||
if (Array.isArray(value)) {
|
||||
for (const byte of value) {
|
||||
if (!Number.isInteger(byte) || byte < 0 || byte > 255) {
|
||||
throw new RangeError(`${name} contains a non-byte value`);
|
||||
}
|
||||
}
|
||||
return Uint8Array.from(value);
|
||||
}
|
||||
throw new TypeError(`${name} must be a Uint8Array or number[]`);
|
||||
}
|
||||
|
||||
export function strictHexDecode(value: string, name = "value"): Uint8Array {
|
||||
if (typeof value !== "string") throw new TypeError(`${name} must be a string`);
|
||||
const text = value.replace(/^0x/i, "");
|
||||
if (text.length % 2 !== 0 || !/^[0-9a-fA-F]*$/.test(text)) {
|
||||
throw new TypeError(`${name} must be an even-length hexadecimal string`);
|
||||
}
|
||||
const bytes = new Uint8Array(text.length / 2);
|
||||
for (let i = 0; i < bytes.length; i += 1) {
|
||||
bytes[i] = Number.parseInt(text.slice(i * 2, i * 2 + 2), 16);
|
||||
}
|
||||
return bytes;
|
||||
}
|
||||
|
||||
export function strictBase64Decode(value: string, name = "value"): Uint8Array {
|
||||
if (typeof value !== "string") throw new TypeError(`${name} must be a string`);
|
||||
if (value.length === 0) return new Uint8Array(0);
|
||||
if (
|
||||
value.length % 4 !== 0 ||
|
||||
!/^(?:[A-Za-z0-9+/]{4})*(?:[A-Za-z0-9+/]{2}==|[A-Za-z0-9+/]{3}=)?$/.test(
|
||||
value,
|
||||
)
|
||||
) {
|
||||
throw new TypeError(`${name} is not valid padded base64`);
|
||||
}
|
||||
|
||||
let bytes: Uint8Array;
|
||||
try {
|
||||
if (typeof atob === "function") {
|
||||
const binary = atob(value);
|
||||
bytes = new Uint8Array(binary.length);
|
||||
for (let i = 0; i < binary.length; i += 1) {
|
||||
bytes[i] = binary.charCodeAt(i);
|
||||
}
|
||||
} else if (typeof Buffer !== "undefined") {
|
||||
bytes = new Uint8Array(Buffer.from(value, "base64"));
|
||||
} else {
|
||||
throw new TypeError("base64 decoding is not available in this environment");
|
||||
}
|
||||
} catch (error) {
|
||||
throw new TypeError(`${name} is not valid base64`, { cause: error });
|
||||
}
|
||||
|
||||
if (bytesToBase64(bytes) !== value) {
|
||||
throw new TypeError(`${name} is not canonical padded base64`);
|
||||
}
|
||||
|
||||
return bytes;
|
||||
}
|
||||
|
||||
export function bytesFromEncodedString(
|
||||
value: string,
|
||||
encoding: "hex" | "base64",
|
||||
name: string,
|
||||
): Uint8Array {
|
||||
return encoding === "hex"
|
||||
? strictHexDecode(value, name)
|
||||
: strictBase64Decode(value, name);
|
||||
}
|
||||
|
||||
/*
|
||||
* Compatibility parser for the historical format-detecting API. New callers
|
||||
* should select `bytesFromEncodedString` explicitly so a value cannot change
|
||||
* meaning when it happens to contain only hexadecimal characters.
|
||||
*/
|
||||
/** @deprecated Use `bytesFromEncodedString(value, encoding, name)`. */
|
||||
export function bytesFromString(value: string, name: string): Uint8Array {
|
||||
const trimmed = value.trim();
|
||||
if (!trimmed) throw new TypeError(`${name} must not be empty`);
|
||||
|
||||
const hex = trimmed.replace(/^(0x)/i, "").replace(/[\s:_-]/g, "");
|
||||
if (/^[0-9a-fA-F]+$/.test(hex)) {
|
||||
return strictHexDecode(hex, name);
|
||||
}
|
||||
return strictBase64Decode(trimmed, name);
|
||||
}
|
||||
|
||||
const HEX_DIGITS = "0123456789abcdef";
|
||||
|
||||
function bytesToHex(bytes: Uint8Array): string {
|
||||
let out = "";
|
||||
for (let i = 0; i < bytes.length; i += 1) {
|
||||
out += HEX_DIGITS[(bytes[i] >> 4) & 0xf] + HEX_DIGITS[bytes[i] & 0xf];
|
||||
}
|
||||
return out;
|
||||
}
|
||||
|
||||
export function bytesToBase64(bytes: Uint8Array): string {
|
||||
if (typeof btoa === "function") {
|
||||
let binary = "";
|
||||
for (let i = 0; i < bytes.length; i += 1) {
|
||||
binary += String.fromCharCode(bytes[i]);
|
||||
}
|
||||
return btoa(binary);
|
||||
}
|
||||
if (typeof Buffer !== "undefined") {
|
||||
return Buffer.from(bytes).toString("base64");
|
||||
}
|
||||
throw new TypeError("base64 encoding is not available in this environment");
|
||||
}
|
||||
|
||||
export function base64ToBytes(input: string): Uint8Array {
|
||||
return strictBase64Decode(input, "base64");
|
||||
}
|
||||
|
||||
function utf8Decode(bytes: Uint8Array): string {
|
||||
if (typeof TextDecoder !== "undefined") {
|
||||
try {
|
||||
return new TextDecoder("utf-8", { fatal: true }).decode(bytes);
|
||||
} catch (error) {
|
||||
throw new TypeError("invalid UTF-8", { cause: error });
|
||||
}
|
||||
}
|
||||
let out = "";
|
||||
let i = 0;
|
||||
while (i < bytes.length) {
|
||||
const b = bytes[i];
|
||||
if (b < 0x80) {
|
||||
out += String.fromCharCode(b);
|
||||
i += 1;
|
||||
} else if (b >= 0xc2 && b <= 0xdf) {
|
||||
if (i + 1 >= bytes.length || (bytes[i + 1] & 0xc0) !== 0x80) {
|
||||
throw new TypeError("invalid UTF-8");
|
||||
}
|
||||
out += String.fromCharCode(((b & 0x1f) << 6) | (bytes[i + 1] & 0x3f));
|
||||
i += 2;
|
||||
} else if (b >= 0xe0 && b <= 0xef) {
|
||||
if (
|
||||
i + 2 >= bytes.length ||
|
||||
(bytes[i + 1] & 0xc0) !== 0x80 ||
|
||||
(bytes[i + 2] & 0xc0) !== 0x80 ||
|
||||
(b === 0xe0 && bytes[i + 1] < 0xa0) ||
|
||||
(b === 0xed && bytes[i + 1] >= 0xa0)
|
||||
) {
|
||||
throw new TypeError("invalid UTF-8");
|
||||
}
|
||||
out += String.fromCharCode(
|
||||
((b & 0x0f) << 12) |
|
||||
((bytes[i + 1] & 0x3f) << 6) |
|
||||
(bytes[i + 2] & 0x3f),
|
||||
);
|
||||
i += 3;
|
||||
} else if (b >= 0xf0 && b <= 0xf4) {
|
||||
if (
|
||||
i + 3 >= bytes.length ||
|
||||
(bytes[i + 1] & 0xc0) !== 0x80 ||
|
||||
(bytes[i + 2] & 0xc0) !== 0x80 ||
|
||||
(bytes[i + 3] & 0xc0) !== 0x80 ||
|
||||
(b === 0xf0 && bytes[i + 1] < 0x90) ||
|
||||
(b === 0xf4 && bytes[i + 1] >= 0x90)
|
||||
) {
|
||||
throw new TypeError("invalid UTF-8");
|
||||
}
|
||||
const cp =
|
||||
((b & 0x07) << 18) |
|
||||
((bytes[i + 1] & 0x3f) << 12) |
|
||||
((bytes[i + 2] & 0x3f) << 6) |
|
||||
(bytes[i + 3] & 0x3f);
|
||||
out += String.fromCodePoint(cp);
|
||||
i += 4;
|
||||
} else {
|
||||
throw new TypeError("invalid UTF-8");
|
||||
}
|
||||
}
|
||||
return out;
|
||||
}
|
||||
|
||||
function requiredSecretKeyLength(): number {
|
||||
const lengthBinding = (
|
||||
bindings as typeof bindings & {
|
||||
mtp_symmetric_key_length?: () => number;
|
||||
}
|
||||
).mtp_symmetric_key_length;
|
||||
if (!lengthBinding) return 32;
|
||||
try {
|
||||
return lengthBinding();
|
||||
} catch {
|
||||
// The generated WASM wrapper is callable only after initialization. Keep
|
||||
// the historical size as a pre-initialization validation fallback.
|
||||
return 32;
|
||||
}
|
||||
}
|
||||
|
||||
export function secretKeyFromBytes(value: MTPBytesInput): Uint8Array {
|
||||
const bytes = bytesFrom(value, "secret key");
|
||||
const requiredLength = requiredSecretKeyLength();
|
||||
if (bytes.length !== requiredLength) {
|
||||
throw new RangeError(`secret key must be exactly ${requiredLength} bytes`);
|
||||
}
|
||||
return bytes;
|
||||
}
|
||||
|
||||
export function secretKeyFromHex(value: string): Uint8Array {
|
||||
return secretKeyFromBytes(strictHexDecode(value, "secret key"));
|
||||
}
|
||||
|
||||
export function secretKeyFromBase64(value: string): Uint8Array {
|
||||
return secretKeyFromBytes(strictBase64Decode(value, "secret key"));
|
||||
}
|
||||
|
||||
/*
|
||||
* Compatibility entry point. It now accepts only explicitly encoded key
|
||||
* material; arbitrary strings are no longer silently treated as passphrases.
|
||||
*/
|
||||
/** @deprecated Use `secretKeyFromBytes`, `secretKeyFromHex`, or `secretKeyFromBase64`. */
|
||||
export function secretKeyFromString(secret: string): Uint8Array {
|
||||
if (typeof secret !== "string" || !secret.trim()) {
|
||||
throw new TypeError("secret must be a non-empty string");
|
||||
}
|
||||
const trimmed = secret.trim();
|
||||
const hex = trimmed.replace(/^(0x)/i, "");
|
||||
if (/^[0-9a-fA-F]+$/.test(hex)) return secretKeyFromHex(hex);
|
||||
return secretKeyFromBase64(trimmed);
|
||||
}
|
||||
|
||||
/**
|
||||
* Reproduce the pre-v1 implicit-HKDF derivation for data migration only.
|
||||
*
|
||||
* @deprecated Do not use for new secrets. Replace this with explicit key
|
||||
* material or `deriveKeyFromPassphrase` and persist a password-KDF salt.
|
||||
*/
|
||||
export function legacySecretKeyFromStringV1(secret: string): Uint8Array {
|
||||
if (typeof secret !== "string" || !secret.trim()) {
|
||||
throw new TypeError("secret must be a non-empty string");
|
||||
}
|
||||
const trimmed = secret.trim();
|
||||
const hex = trimmed.replace(/^(0x)/i, "").replace(/[\s:_-]/g, "");
|
||||
if (/^[0-9a-fA-F]+$/.test(hex) && hex.length === 64) {
|
||||
return strictHexDecode(hex, "legacy secret key");
|
||||
}
|
||||
try {
|
||||
const decoded = bytesFromString(trimmed, "legacy secret key");
|
||||
if (decoded.length === requiredSecretKeyLength()) return decoded;
|
||||
} catch {
|
||||
// Preserve the historical fallback to HKDF for non-encoded strings.
|
||||
}
|
||||
const context = utf8Encode("mtp-symmetric-key");
|
||||
return bindings.wasm_derive_encryption_key(
|
||||
utf8Encode(trimmed),
|
||||
context,
|
||||
context,
|
||||
);
|
||||
}
|
||||
|
||||
export interface PasswordKdfParameters {
|
||||
memoryKiB: number;
|
||||
iterations: number;
|
||||
lanes: number;
|
||||
}
|
||||
|
||||
function validatePasswordKdfInput(
|
||||
passphrase: string,
|
||||
salt: MTPBytesInput,
|
||||
parameters: PasswordKdfParameters,
|
||||
): { passphrase: string; salt: Uint8Array; parameters: PasswordKdfParameters } {
|
||||
if (typeof passphrase !== "string" || passphrase.length === 0) {
|
||||
throw new TypeError("passphrase must not be empty");
|
||||
}
|
||||
const saltBytes = bytesFrom(salt, "passphrase salt");
|
||||
if (saltBytes.length < 16) {
|
||||
throw new RangeError("passphrase salt must be at least 16 bytes");
|
||||
}
|
||||
if (
|
||||
!Number.isInteger(parameters.memoryKiB) ||
|
||||
parameters.memoryKiB < 8 * 1024 ||
|
||||
parameters.memoryKiB > 256 * 1024 ||
|
||||
!Number.isInteger(parameters.iterations) ||
|
||||
parameters.iterations < 1 ||
|
||||
parameters.iterations > 10 ||
|
||||
!Number.isInteger(parameters.lanes) ||
|
||||
parameters.lanes < 1 ||
|
||||
parameters.lanes > 8
|
||||
) {
|
||||
throw new RangeError("invalid Argon2id password-KDF parameters");
|
||||
}
|
||||
return { passphrase, salt: saltBytes, parameters };
|
||||
}
|
||||
|
||||
function deriveKeyFromPassphraseSyncImpl(
|
||||
passphrase: string,
|
||||
salt: MTPBytesInput,
|
||||
parameters: PasswordKdfParameters,
|
||||
): Uint8Array {
|
||||
const validated = validatePasswordKdfInput(passphrase, salt, parameters);
|
||||
const kdf = (bindings as unknown as {
|
||||
wasm_argon2id?: (
|
||||
passphrase: Uint8Array,
|
||||
salt: Uint8Array,
|
||||
memoryKiB: number,
|
||||
iterations: number,
|
||||
lanes: number,
|
||||
) => Uint8Array;
|
||||
}).wasm_argon2id;
|
||||
if (!kdf) {
|
||||
throw new Error("Argon2id password derivation is unavailable in this WASM build");
|
||||
}
|
||||
return kdf(
|
||||
utf8Encode(validated.passphrase),
|
||||
validated.salt,
|
||||
validated.parameters.memoryKiB,
|
||||
validated.parameters.iterations,
|
||||
validated.parameters.lanes,
|
||||
);
|
||||
}
|
||||
|
||||
/**
|
||||
* Derive a passphrase key without yielding. Prefer the asynchronous API in
|
||||
* browser applications; this form is retained for workers and synchronous
|
||||
* command-line migrations.
|
||||
*/
|
||||
/** @deprecated Use `deriveKeyFromPassphrase` in browser-facing code. */
|
||||
export function deriveKeyFromPassphraseSync(
|
||||
passphrase: string,
|
||||
salt: MTPBytesInput,
|
||||
parameters: PasswordKdfParameters,
|
||||
): Uint8Array {
|
||||
return deriveKeyFromPassphraseSyncImpl(passphrase, salt, parameters);
|
||||
}
|
||||
|
||||
/**
|
||||
* Derive a passphrase key off the browser main thread when workers are
|
||||
* available. The worker imports the same generated WASM binding, so the
|
||||
* Argon2id computation does not block UI/event-loop work.
|
||||
*/
|
||||
export function deriveKeyFromPassphrase(
|
||||
passphrase: string,
|
||||
salt: MTPBytesInput,
|
||||
parameters: PasswordKdfParameters,
|
||||
): Promise<Uint8Array> {
|
||||
const validated = validatePasswordKdfInput(passphrase, salt, parameters);
|
||||
if (typeof Worker === "undefined") {
|
||||
return initWasmOnce().then(
|
||||
() =>
|
||||
new Promise((resolve) => {
|
||||
setTimeout(
|
||||
() =>
|
||||
resolve(
|
||||
deriveKeyFromPassphraseSyncImpl(
|
||||
validated.passphrase,
|
||||
validated.salt,
|
||||
validated.parameters,
|
||||
),
|
||||
),
|
||||
0,
|
||||
);
|
||||
}),
|
||||
);
|
||||
}
|
||||
|
||||
const worker = new Worker(new URL("./passphrase-worker.js", import.meta.url), {
|
||||
type: "module",
|
||||
});
|
||||
return new Promise<Uint8Array>((resolve, reject) => {
|
||||
const cleanup = () => worker.terminate();
|
||||
worker.onmessage = (event: MessageEvent<Uint8Array | { error: string }>) => {
|
||||
cleanup();
|
||||
if (event.data && "error" in event.data) {
|
||||
reject(new Error(event.data.error));
|
||||
} else {
|
||||
resolve(new Uint8Array(event.data));
|
||||
}
|
||||
};
|
||||
worker.onerror = (event) => {
|
||||
cleanup();
|
||||
reject(new Error(event.message || "Argon2id worker failed"));
|
||||
};
|
||||
const passphraseBytes = utf8Encode(validated.passphrase);
|
||||
const saltBytes = validated.salt.slice();
|
||||
worker.postMessage(
|
||||
{
|
||||
passphrase: passphraseBytes,
|
||||
salt: saltBytes,
|
||||
parameters: validated.parameters,
|
||||
},
|
||||
[passphraseBytes.buffer, saltBytes.buffer],
|
||||
);
|
||||
});
|
||||
}
|
||||
|
||||
export function normalizeBytes(
|
||||
value: string | MTPBytesInput | MTPEncodedBytesInput,
|
||||
name: string,
|
||||
encoding?: "hex" | "base64",
|
||||
): Uint8Array {
|
||||
if (typeof value === "string") {
|
||||
if (!encoding) {
|
||||
throw new TypeError(
|
||||
`${name} string input requires an explicit 'hex' or 'base64' encoding`,
|
||||
);
|
||||
}
|
||||
return bytesFromEncodedString(value, encoding, name);
|
||||
}
|
||||
if (
|
||||
value !== null &&
|
||||
typeof value === "object" &&
|
||||
!(value instanceof Uint8Array) &&
|
||||
!Array.isArray(value)
|
||||
) {
|
||||
const encoded = value as Partial<MTPEncodedBytesInput>;
|
||||
if (
|
||||
typeof encoded.value !== "string" ||
|
||||
(encoded.encoding !== "hex" && encoded.encoding !== "base64")
|
||||
) {
|
||||
throw new TypeError(
|
||||
`${name} must be bytes or { value: string, encoding: 'hex' | 'base64' }`,
|
||||
);
|
||||
}
|
||||
return bytesFromEncodedString(encoded.value, encoded.encoding, name);
|
||||
}
|
||||
return bytesFrom(value, name);
|
||||
}
|
||||
|
||||
export function inputU64(value: bigint | number | string, name: string): bigint {
|
||||
if (typeof value === "number" && !Number.isSafeInteger(value)) {
|
||||
throw new RangeError(
|
||||
`${name} must be a safe integer number, bigint, or integer string`,
|
||||
);
|
||||
}
|
||||
let result: bigint;
|
||||
try {
|
||||
result = BigInt(value);
|
||||
} catch (error) {
|
||||
throw new RangeError(`${name} must be an integer`, { cause: error });
|
||||
}
|
||||
if (result < 0n || result > 0xffff_ffff_ffff_ffffn) {
|
||||
throw new RangeError(`${name} must be a u64`);
|
||||
}
|
||||
return result;
|
||||
}
|
||||
|
||||
export function toBigInt(
|
||||
value: bigint | string | number | null | undefined,
|
||||
): bigint | null {
|
||||
if (value == null || value === "") return null;
|
||||
return inputU64(value, "clientId");
|
||||
}
|
||||
|
||||
const KEM_PUBLIC_KEY_LEN = 1216;
|
||||
const SIG_PQ_PUBLIC_KEY_LEN = 1952;
|
||||
const SIG_CL_PUBLIC_KEY_LEN = 32;
|
||||
|
||||
export function keyringToKeys(keyring: MTPKeyMaterialInput): MTPKeyringKeys {
|
||||
const bytes = normalizeBytes(keyring, "keyring");
|
||||
if (bytes.length < 12) {
|
||||
throw new TypeError("keyring data is too short to contain 6 keys");
|
||||
}
|
||||
|
||||
let offset = 0;
|
||||
const readKey = () => {
|
||||
if (offset + 2 > bytes.length) throw new TypeError("keyring is truncated");
|
||||
const len = (bytes[offset] << 8) | bytes[offset + 1];
|
||||
offset += 2;
|
||||
if (offset + len > bytes.length) throw new TypeError("keyring is truncated");
|
||||
const key = bytes.slice(offset, offset + len);
|
||||
offset += len;
|
||||
return key;
|
||||
};
|
||||
|
||||
const result = {
|
||||
kemPublicKey: readKey(),
|
||||
kemSecretKey: readKey(),
|
||||
sigPqPublicKey: readKey(),
|
||||
sigPqSecretKey: readKey(),
|
||||
sigClPublicKey: readKey(),
|
||||
sigClSecretKey: readKey(),
|
||||
};
|
||||
if (offset !== bytes.length) throw new TypeError("keyring has trailing data");
|
||||
return result;
|
||||
}
|
||||
|
||||
export function publicKeyBundleToKeys(
|
||||
publicKeyBundle: MTPKeyMaterialInput,
|
||||
): MTPPublicKeyBundleKeys {
|
||||
const bytes = normalizeBytes(publicKeyBundle, "publicKeyBundle");
|
||||
if (bytes.length < 6) {
|
||||
throw new TypeError("public key bundle data is too short to contain 3 keys");
|
||||
}
|
||||
|
||||
let offset = 0;
|
||||
const readKey = () => {
|
||||
if (offset + 2 > bytes.length) {
|
||||
throw new TypeError("public key bundle is truncated");
|
||||
}
|
||||
const len = (bytes[offset] << 8) | bytes[offset + 1];
|
||||
offset += 2;
|
||||
if (offset + len > bytes.length) {
|
||||
throw new TypeError("public key bundle is truncated");
|
||||
}
|
||||
const key = bytes.slice(offset, offset + len);
|
||||
offset += len;
|
||||
return key;
|
||||
};
|
||||
|
||||
const result = {
|
||||
kemPublicKey: readKey(),
|
||||
sigPqPublicKey: readKey(),
|
||||
sigClPublicKey: readKey(),
|
||||
};
|
||||
if (offset !== bytes.length) {
|
||||
throw new TypeError("public key bundle has trailing data");
|
||||
}
|
||||
if (
|
||||
result.kemPublicKey.length !== KEM_PUBLIC_KEY_LEN ||
|
||||
result.sigPqPublicKey.length !== SIG_PQ_PUBLIC_KEY_LEN ||
|
||||
result.sigClPublicKey.length !== SIG_CL_PUBLIC_KEY_LEN
|
||||
) {
|
||||
throw new TypeError("public key bundle contains invalid suite key lengths");
|
||||
}
|
||||
return result;
|
||||
}
|
||||
|
||||
export function cloneParsedValue(value: unknown): unknown {
|
||||
if (value instanceof Uint8Array) return value.slice();
|
||||
if (Array.isArray(value)) return value.map(cloneParsedValue);
|
||||
if (value !== null && typeof value === "object") {
|
||||
return Object.fromEntries(
|
||||
Object.entries(value).map(([key, entry]) => [key, cloneParsedValue(entry)]),
|
||||
);
|
||||
}
|
||||
return value;
|
||||
}
|
||||
|
||||
export function cloneParsedFrame(frame: ParsedFrame): ParsedFrame {
|
||||
return cloneParsedValue(frame) as ParsedFrame;
|
||||
}
|
||||
|
||||
function parsedDataObject(
|
||||
data: ParsedFrame["data"] | null | undefined,
|
||||
): Record<string, unknown> {
|
||||
if (
|
||||
data === null ||
|
||||
typeof data !== "object" ||
|
||||
Array.isArray(data) ||
|
||||
data instanceof Uint8Array
|
||||
) {
|
||||
return {};
|
||||
}
|
||||
const object = data as Record<string, unknown>;
|
||||
if (object.kind === "encrypted" || object.kind === "signed") return {};
|
||||
return object;
|
||||
}
|
||||
|
||||
export function errorMessage(
|
||||
frame: Pick<ParsedFrame, "type" | "data"> | null | undefined,
|
||||
): string {
|
||||
const data = parsedDataObject(frame?.data);
|
||||
return String(
|
||||
data.ErrorMessage ??
|
||||
data.Error ??
|
||||
data.Description ??
|
||||
`Received ${frame?.type ?? "error"} frame`,
|
||||
);
|
||||
}
|
||||
|
||||
export function parseProtectedFrame(
|
||||
frame: MTPProtectedFrameInput,
|
||||
limits?: MTPReceiveLimits,
|
||||
): ParsedFrame {
|
||||
const parse = (bytes: Uint8Array): ParsedFrame =>
|
||||
limits ? decodeWithLimits(bytes, limits) : bindings.parse_frame(bytes);
|
||||
if (isBytes(frame)) return parse(bytesFrom(frame, "frame"));
|
||||
if (
|
||||
frame === null ||
|
||||
typeof frame !== "object" ||
|
||||
typeof frame.type !== "string"
|
||||
) {
|
||||
throw new TypeError("frame must be a parsed MTP frame or serialized bytes");
|
||||
}
|
||||
if (frame.raw instanceof Uint8Array) return parse(frame.raw);
|
||||
return frame;
|
||||
}
|
||||
|
||||
export function assertKnownCommunicationType(frame: ParsedFrame): void {
|
||||
if (!frame.type || /^[0-9]+$/.test(frame.type)) {
|
||||
throw new Error(`Unknown communication type: ${frame.type || "unknown"}`);
|
||||
}
|
||||
try {
|
||||
bindings.build_frame(frame.type, null, {});
|
||||
} catch (error) {
|
||||
throw new Error(`Unknown communication type: ${frame.type}`, { cause: error });
|
||||
}
|
||||
}
|
||||
|
||||
export function protectedFrameBytes(
|
||||
frame: ParsedFrame,
|
||||
limits?: MTPReceiveLimits,
|
||||
): Uint8Array {
|
||||
if (frame.raw instanceof Uint8Array) return frame.raw.slice();
|
||||
const data =
|
||||
frame.data !== null &&
|
||||
typeof frame.data === "object" &&
|
||||
!Array.isArray(frame.data) &&
|
||||
!(frame.data instanceof Uint8Array)
|
||||
? (frame.data as Record<string, unknown>)
|
||||
: null;
|
||||
const encoded = data?.encoded;
|
||||
if (data?.kind !== "encrypted" || !(encoded instanceof Uint8Array)) {
|
||||
throw new Error("protected frame payload is not encrypted");
|
||||
}
|
||||
const options = {
|
||||
id: frame.id,
|
||||
...(frame.sender == null ? {} : { sender: frame.sender }),
|
||||
...(frame.receiver == null ? {} : { receiver: frame.receiver }),
|
||||
};
|
||||
const bounded = (
|
||||
bindings as typeof bindings & {
|
||||
build_frame_with_payload_with_limits?: (
|
||||
type: string,
|
||||
payload: Uint8Array,
|
||||
options: MTPCodecOptions,
|
||||
limits: MTPReceiveLimits,
|
||||
) => Uint8Array;
|
||||
}
|
||||
).build_frame_with_payload_with_limits;
|
||||
if (!bounded) {
|
||||
throw new Error("bounded WASM frame encoding is unavailable; rebuild mtp-wasm");
|
||||
}
|
||||
return bounded(frame.type, encoded, options, limits ?? {});
|
||||
}
|
||||
|
||||
export function assertApplicationCommunicationType(type: string): string {
|
||||
if (typeof type !== "string" || !type.trim() || /^[0-9]+$/.test(type)) {
|
||||
throw new Error(`Unknown communication type: ${type || "unknown"}`);
|
||||
}
|
||||
if (Object.prototype.hasOwnProperty.call(RESERVED_COMMUNICATION_TYPE_IDS, type)) {
|
||||
throw new Error(
|
||||
`MTP control communication type ${type} cannot be used as application content`,
|
||||
);
|
||||
}
|
||||
try {
|
||||
bindings.build_frame(type, null, {});
|
||||
} catch (error) {
|
||||
throw new Error(`Unknown communication type: ${type}`, { cause: error });
|
||||
}
|
||||
return type;
|
||||
}
|
||||
|
||||
export const MAX_DATA_VALUE_DEPTH = 64;
|
||||
export const MAX_DATA_VALUE_VALUES = 65_536;
|
||||
|
||||
const DEFAULT_ENCODE_LIMITS: Required<MTPEncodeLimits> = {
|
||||
maxDepth: MAX_DATA_VALUE_DEPTH,
|
||||
maxValues: MAX_DATA_VALUE_VALUES,
|
||||
maxOutputSize: 16 * 1024 * 1024,
|
||||
};
|
||||
|
||||
function normalizedEncodeLimits(
|
||||
limits: MTPEncodeLimits | undefined,
|
||||
): Required<MTPEncodeLimits> {
|
||||
const result = { ...DEFAULT_ENCODE_LIMITS, ...(limits ?? {}) };
|
||||
for (const [key, value] of Object.entries(result)) {
|
||||
if (!Number.isSafeInteger(value) || value < 0) {
|
||||
throw new TypeError(`encode limits ${key} must be a non-negative safe integer`);
|
||||
}
|
||||
}
|
||||
return result as Required<MTPEncodeLimits>;
|
||||
}
|
||||
|
||||
/** Validate a JS DataValue before crossing into the recursive WASM parser. */
|
||||
export function validateMTPDataValue(
|
||||
value: MTPDataValueInput,
|
||||
limits?: MTPEncodeLimits,
|
||||
): void {
|
||||
const effective = normalizedEncodeLimits(limits);
|
||||
const ancestors = new WeakSet<object>();
|
||||
let values = 0;
|
||||
const validate = (candidate: unknown, depth: number): void => {
|
||||
values += 1;
|
||||
if (values > effective.maxValues) {
|
||||
throw new RangeError("MTP DataValue value-count limit exceeded");
|
||||
}
|
||||
if (depth > effective.maxDepth) {
|
||||
throw new RangeError("MTP DataValue nesting-depth limit exceeded");
|
||||
}
|
||||
if (
|
||||
candidate === null ||
|
||||
typeof candidate === "boolean" ||
|
||||
typeof candidate === "string" ||
|
||||
typeof candidate === "bigint" ||
|
||||
candidate instanceof Uint8Array
|
||||
) {
|
||||
return;
|
||||
}
|
||||
if (typeof candidate === "number") {
|
||||
if (Number.isInteger(candidate) && !Number.isSafeInteger(candidate)) {
|
||||
throw new TypeError("unsafe integral MTP DataValue inputs must use bigint");
|
||||
}
|
||||
return;
|
||||
}
|
||||
if (typeof candidate !== "object") {
|
||||
throw new TypeError(`unsupported MTP DataValue input: ${typeof candidate}`);
|
||||
}
|
||||
const object = candidate as object;
|
||||
if (ancestors.has(object)) throw new TypeError("MTP DataValue input must not be cyclic");
|
||||
if (
|
||||
!Array.isArray(candidate) &&
|
||||
Object.getPrototypeOf(candidate) !== Object.prototype &&
|
||||
Object.getPrototypeOf(candidate) !== null
|
||||
) {
|
||||
throw new TypeError("MTP DataValue containers must be plain objects");
|
||||
}
|
||||
ancestors.add(object);
|
||||
const entries = Array.isArray(candidate)
|
||||
? candidate
|
||||
: Object.values(candidate as Record<string, unknown>);
|
||||
try {
|
||||
for (const entry of entries) validate(entry, depth + 1);
|
||||
} finally {
|
||||
ancestors.delete(object);
|
||||
}
|
||||
};
|
||||
validate(value, 0);
|
||||
}
|
||||
|
||||
export function encodeMTPDataValue(
|
||||
value: MTPDataValueInput,
|
||||
limits?: MTPEncodeLimits,
|
||||
): Uint8Array {
|
||||
const effective = normalizedEncodeLimits(limits);
|
||||
validateMTPDataValue(value, effective);
|
||||
const bounded = (
|
||||
bindings as typeof bindings & {
|
||||
encode_data_value_with_limits?: (
|
||||
value: MTPDataValueInput,
|
||||
limits: MTPEncodeLimits,
|
||||
) => Uint8Array;
|
||||
}
|
||||
).encode_data_value_with_limits;
|
||||
if (!bounded) {
|
||||
throw new Error("bounded WASM DataValue encoding is unavailable; rebuild mtp-wasm");
|
||||
}
|
||||
const encoded = bounded(value, effective);
|
||||
if (encoded.length > effective.maxOutputSize) {
|
||||
throw new RangeError("MTP DataValue encoded output limit exceeded");
|
||||
}
|
||||
return encoded;
|
||||
}
|
||||
|
||||
export function inputDataValueBigInt(value: unknown, name: string): bigint {
|
||||
try {
|
||||
if (typeof value === "bigint") return value;
|
||||
if (typeof value === "number" && Number.isSafeInteger(value)) return BigInt(value);
|
||||
if (typeof value === "string" && value.length > 0) return BigInt(value);
|
||||
} catch {
|
||||
// Normalize malformed protected metadata below.
|
||||
}
|
||||
throw new Error(`protected metadata field ${name} is not an integer`);
|
||||
}
|
||||
|
||||
export function inputDataValueString(value: unknown, name: string): string {
|
||||
if (typeof value === "string" && value.length > 0) return value;
|
||||
throw new Error(`protected metadata field ${name} is not a non-empty string`);
|
||||
}
|
||||
|
||||
export function signatureSuiteValue(
|
||||
suite: MTPProtectionSignatureSuite,
|
||||
): number {
|
||||
return suite === "dual"
|
||||
? bindings.mtp_protection_signature_suite_dual()
|
||||
: bindings.mtp_protection_signature_suite_ed25519();
|
||||
}
|
||||
|
||||
export function formatDataValue(value: MTPDataValue): MTPDataValue {
|
||||
return cloneParsedValue(value) as MTPDataValue;
|
||||
}
|
||||
|
|
@ -1,26 +0,0 @@
|
|||
import type { MTPClientCredentials } from "./index.js";
|
||||
|
||||
export type InternalCredentials = {
|
||||
clientId: bigint | null;
|
||||
keyring: Uint8Array;
|
||||
hostPublicKey?: Uint8Array;
|
||||
};
|
||||
|
||||
export function publicCredentials(
|
||||
credentials: InternalCredentials | null,
|
||||
): MTPClientCredentials | null {
|
||||
if (!credentials) {
|
||||
return null;
|
||||
}
|
||||
return {
|
||||
clientId: credentials.clientId,
|
||||
keyring: credentials.keyring.slice(),
|
||||
hostPublicKey: credentials.hostPublicKey?.slice(),
|
||||
};
|
||||
}
|
||||
|
||||
export function zeroCredentials(credentials: InternalCredentials | null): void {
|
||||
// The host public key is intentionally not wiped: it is public configuration
|
||||
// and may also be retained by the connection options.
|
||||
credentials?.keyring.fill(0);
|
||||
}
|
||||
3151
src/sdk/index.ts
3151
src/sdk/index.ts
File diff suppressed because it is too large
Load diff
|
|
@ -1,33 +0,0 @@
|
|||
import initWasm, * as bindings from "mtp/raw";
|
||||
|
||||
interface PasswordKdfWorkerRequest {
|
||||
passphrase: Uint8Array;
|
||||
salt: Uint8Array;
|
||||
parameters: {
|
||||
memoryKiB: number;
|
||||
iterations: number;
|
||||
lanes: number;
|
||||
};
|
||||
}
|
||||
|
||||
const scope = globalThis as unknown as {
|
||||
onmessage: ((event: MessageEvent<PasswordKdfWorkerRequest>) => void) | null;
|
||||
postMessage(message: Uint8Array | { error: string }, transfer?: Transferable[]): void;
|
||||
};
|
||||
|
||||
scope.onmessage = async (event) => {
|
||||
try {
|
||||
await initWasm();
|
||||
const { passphrase, salt, parameters } = event.data;
|
||||
const key = bindings.wasm_argon2id(
|
||||
passphrase,
|
||||
salt,
|
||||
parameters.memoryKiB,
|
||||
parameters.iterations,
|
||||
parameters.lanes,
|
||||
);
|
||||
scope.postMessage(key, [key.buffer]);
|
||||
} catch (error) {
|
||||
scope.postMessage({ error: String(error) });
|
||||
}
|
||||
};
|
||||
|
|
@ -1,258 +0,0 @@
|
|||
import type { InternalCredentials } from "./credentials.js";
|
||||
import {
|
||||
inputU64,
|
||||
keyringToKeys,
|
||||
normalizeBytes,
|
||||
publicKeyBundleToKeys,
|
||||
signatureSuiteValue,
|
||||
} from "./codec.js";
|
||||
import {
|
||||
MTPSignatureVerificationError,
|
||||
signerKeysUnavailable,
|
||||
} from "./signature-policy.js";
|
||||
import type { MTPSignatureVerificationPolicy } from "./signature-policy.js";
|
||||
import type {
|
||||
MTPDecryptionIdentity,
|
||||
MTPProtectionIdentity,
|
||||
MTPProtectionSignatureSuite,
|
||||
MTPReplayGuard,
|
||||
MTPSignerKeyResolver,
|
||||
MTPBytesInput,
|
||||
MTPKeyMaterialInput,
|
||||
} from "./client.js";
|
||||
|
||||
export class InMemoryReplayGuard implements MTPReplayGuard {
|
||||
#accepted = new Set<string>();
|
||||
readonly #capacity = 10_000;
|
||||
|
||||
accept(signerId: bigint, messageId: string, _createdAt: bigint): boolean {
|
||||
const key = `${signerId}:${messageId}`;
|
||||
if (this.#accepted.has(key)) return false;
|
||||
this.#accepted.add(key);
|
||||
if (this.#accepted.size > this.#capacity) {
|
||||
const oldest = this.#accepted.values().next().value;
|
||||
if (oldest !== undefined) this.#accepted.delete(oldest);
|
||||
}
|
||||
return true;
|
||||
}
|
||||
}
|
||||
|
||||
export class MTPReplayError extends Error {
|
||||
readonly signerId: bigint;
|
||||
readonly messageId: string;
|
||||
|
||||
constructor(signerId: bigint, messageId: string) {
|
||||
super(`message ${messageId} from signer ${signerId} was already accepted`);
|
||||
this.name = "MTPReplayError";
|
||||
this.signerId = signerId;
|
||||
this.messageId = messageId;
|
||||
}
|
||||
}
|
||||
|
||||
export class MTPMissingProtectedVersionError extends Error {
|
||||
constructor() {
|
||||
super("protected message does not declare a protected version");
|
||||
this.name = "MTPMissingProtectedVersionError";
|
||||
}
|
||||
}
|
||||
|
||||
export class MTPUnsupportedProtectedVersionError extends Error {
|
||||
readonly protectedVersion: bigint;
|
||||
|
||||
constructor(protectedVersion: bigint) {
|
||||
super(`unsupported protected message version ${protectedVersion}`);
|
||||
this.name = "MTPUnsupportedProtectedVersionError";
|
||||
this.protectedVersion = protectedVersion;
|
||||
}
|
||||
}
|
||||
|
||||
export class MTPResourceLimitError extends Error {
|
||||
constructor(message = "MTP receive resource limit exceeded") {
|
||||
super(message);
|
||||
this.name = "MTPResourceLimitError";
|
||||
}
|
||||
}
|
||||
|
||||
export interface ResolvedProtectionIdentity {
|
||||
signerId: bigint;
|
||||
keyring: Uint8Array;
|
||||
}
|
||||
|
||||
export interface ResolvedDecryptionIdentity {
|
||||
id?: bigint;
|
||||
keyrings: Uint8Array[];
|
||||
}
|
||||
|
||||
export interface SignerResolutionOptions {
|
||||
expectedSignerId?: bigint | number | string;
|
||||
resolveSignerPublicKeys?: MTPSignerKeyResolver;
|
||||
}
|
||||
|
||||
export function protectionSignatureSuiteValue(
|
||||
suite: MTPProtectionSignatureSuite,
|
||||
): number {
|
||||
return signatureSuiteValue(suite);
|
||||
}
|
||||
|
||||
export function effectiveProtectionSignatureSuite(
|
||||
keyring: Uint8Array,
|
||||
requested?: MTPProtectionSignatureSuite,
|
||||
): MTPProtectionSignatureSuite {
|
||||
const keys = keyringToKeys(keyring);
|
||||
const hasPqPublicKey = keys.sigPqPublicKey.length > 0;
|
||||
const hasPqSecretKey = keys.sigPqSecretKey.length > 0;
|
||||
const suite = requested ?? "ed25519";
|
||||
if (suite !== "ed25519" && suite !== "dual") {
|
||||
throw new Error("signatureSuite must be 'ed25519' or 'dual'");
|
||||
}
|
||||
if (suite === "dual" && !(hasPqPublicKey && hasPqSecretKey)) {
|
||||
throw new Error(
|
||||
"dual protected signatures require a complete ML-DSA key pair; choose 'ed25519' for a partial keyring",
|
||||
);
|
||||
}
|
||||
return suite;
|
||||
}
|
||||
|
||||
function sameBytes(left: Uint8Array, right: Uint8Array): boolean {
|
||||
if (left.length !== right.length) return false;
|
||||
for (let index = 0; index < left.length; index += 1) {
|
||||
if (left[index] !== right[index]) return false;
|
||||
}
|
||||
return true;
|
||||
}
|
||||
|
||||
function normalizeDecryptionKeyrings(
|
||||
identity: MTPDecryptionIdentity,
|
||||
): Uint8Array[] {
|
||||
const current = normalizeBytes(identity.keyring, "recipient.keyring");
|
||||
if (current.length === 0) throw new Error("recipient.keyring must not be empty");
|
||||
if (
|
||||
identity.keyringHistory !== undefined &&
|
||||
!Array.isArray(identity.keyringHistory)
|
||||
) {
|
||||
throw new TypeError("recipient.keyringHistory must be an array");
|
||||
}
|
||||
|
||||
const keyrings: Uint8Array[] = [];
|
||||
const add = (value: MTPKeyMaterialInput, name: string): void => {
|
||||
const bytes = normalizeBytes(value, name);
|
||||
if (bytes.length === 0) throw new Error(`${name} must not be empty`);
|
||||
if (!keyrings.some((existing) => sameBytes(existing, bytes))) {
|
||||
keyrings.push(bytes.slice());
|
||||
}
|
||||
};
|
||||
add(current, "recipient.keyring");
|
||||
for (const [index, history] of (identity.keyringHistory ?? []).entries()) {
|
||||
add(history, `recipient.keyringHistory[${index}]`);
|
||||
}
|
||||
if (keyrings.length === 0) throw new Error("recipient must contain at least one keyring");
|
||||
return keyrings;
|
||||
}
|
||||
|
||||
export function normalizeRecipientBundles(
|
||||
recipients: MTPKeyMaterialInput[],
|
||||
name: string,
|
||||
): Uint8Array[] {
|
||||
if (!Array.isArray(recipients) || recipients.length === 0) {
|
||||
throw new TypeError(`${name} must contain at least one public key bundle`);
|
||||
}
|
||||
return recipients.map((value, index) => {
|
||||
const bundle = normalizeBytes(value, `${name}[${index}]`);
|
||||
publicKeyBundleToKeys(bundle);
|
||||
return bundle.slice();
|
||||
});
|
||||
}
|
||||
|
||||
export function resolveProtectionIdentity(
|
||||
explicit: MTPProtectionIdentity | undefined,
|
||||
stored: InternalCredentials | null,
|
||||
): ResolvedProtectionIdentity {
|
||||
if (explicit) {
|
||||
return {
|
||||
signerId: inputU64(explicit.signerId, "identity.signerId"),
|
||||
keyring: normalizeBytes(explicit.keyring, "identity.keyring").slice(),
|
||||
};
|
||||
}
|
||||
if (stored?.clientId != null && stored.keyring.length > 0) {
|
||||
return { signerId: stored.clientId, keyring: stored.keyring.slice() };
|
||||
}
|
||||
throw new Error(
|
||||
"protected send requires an explicit protection identity or stored registered credentials",
|
||||
);
|
||||
}
|
||||
|
||||
export function resolveDecryptionIdentity(
|
||||
explicit: MTPDecryptionIdentity | undefined,
|
||||
stored: InternalCredentials | null,
|
||||
): ResolvedDecryptionIdentity {
|
||||
if (explicit) {
|
||||
return {
|
||||
id: explicit.id == null ? undefined : inputU64(explicit.id, "recipient.id"),
|
||||
keyrings: normalizeDecryptionKeyrings(explicit),
|
||||
};
|
||||
}
|
||||
if (stored?.clientId != null && stored.keyring.length > 0) {
|
||||
return {
|
||||
id: stored.clientId,
|
||||
keyrings: normalizeDecryptionKeyrings({
|
||||
id: stored.clientId,
|
||||
keyring: stored.keyring,
|
||||
}),
|
||||
};
|
||||
}
|
||||
throw new Error(
|
||||
"protected receive requires an explicit decryption identity or stored registered credentials",
|
||||
);
|
||||
}
|
||||
|
||||
export function protectedOpeningError(error: unknown, signerId?: bigint): Error {
|
||||
if (error !== null && typeof error === "object") {
|
||||
const structured = error as { code?: unknown; protectedVersion?: unknown };
|
||||
if (typeof structured.code === "string") {
|
||||
switch (structured.code) {
|
||||
case "missing-protected-version":
|
||||
return new MTPMissingProtectedVersionError();
|
||||
case "unsupported-protected-version":
|
||||
if (
|
||||
typeof structured.protectedVersion === "bigint" ||
|
||||
typeof structured.protectedVersion === "number" ||
|
||||
typeof structured.protectedVersion === "string"
|
||||
) {
|
||||
return new MTPUnsupportedProtectedVersionError(
|
||||
inputU64(structured.protectedVersion, "protectedVersion"),
|
||||
);
|
||||
}
|
||||
break;
|
||||
case "no-matching-recipient":
|
||||
return new Error("Unable to decrypt protected value with supplied recipient keyrings");
|
||||
case "reserved-application-type":
|
||||
return new Error("MTP control communication types cannot be used as application content");
|
||||
case "signature-policy-mismatch":
|
||||
return new MTPSignatureVerificationError("policy-rejected", signerId);
|
||||
case "unsupported-signature-suite":
|
||||
return new MTPSignatureVerificationError("unsupported-suite", signerId);
|
||||
case "invalid-signature":
|
||||
return new MTPSignatureVerificationError("invalid-signature", signerId);
|
||||
case "signer-id-mismatch":
|
||||
return new Error("protected signer ID mismatch");
|
||||
case "receiver-id-mismatch":
|
||||
return new Error("protected frame receiver ID mismatch");
|
||||
case "message-type-mismatch":
|
||||
return new Error("protected message type does not match outer routing");
|
||||
case "final-recipient-mismatch":
|
||||
return new Error("protected final recipient does not match outer routing receiver");
|
||||
case "sender-id-mismatch":
|
||||
return new Error("protected frame sender does not match authenticated signer");
|
||||
case "signer-key-not-found":
|
||||
return signerKeysUnavailable(signerId);
|
||||
case "replay":
|
||||
return new Error("protected message was already accepted");
|
||||
case "resource-limit":
|
||||
return new MTPResourceLimitError();
|
||||
}
|
||||
}
|
||||
}
|
||||
return error instanceof Error ? error : new Error(String(error));
|
||||
}
|
||||
|
||||
export type { MTPSignatureVerificationPolicy };
|
||||
204
src/sdk/relay.ts
204
src/sdk/relay.ts
|
|
@ -1,204 +0,0 @@
|
|||
import type * as RawBindings from "../raw/index";
|
||||
import { cloneParsedFrame, cloneParsedValue, inputU64 } from "./codec.js";
|
||||
import {
|
||||
MTPSignatureVerificationError,
|
||||
signerKeysUnavailable,
|
||||
} from "./signature-policy.js";
|
||||
import { MTPResourceLimitError } from "./protection.js";
|
||||
import type { MTPSignatureVerificationPolicy } from "./signature-policy.js";
|
||||
import type {
|
||||
MTPDataValue,
|
||||
MTPReceiveLimits,
|
||||
MTPVerifiedRelayContent,
|
||||
ParsedFrame,
|
||||
} from "./client.js";
|
||||
|
||||
export class MTPMissingRelayVersionError extends Error {
|
||||
constructor() {
|
||||
super("relay frame does not declare a relay version");
|
||||
this.name = "MTPMissingRelayVersionError";
|
||||
}
|
||||
}
|
||||
|
||||
export class MTPUnsupportedRelayVersionError extends Error {
|
||||
readonly relayVersion: bigint;
|
||||
|
||||
constructor(relayVersion: bigint) {
|
||||
super(`unsupported relay version ${relayVersion}`);
|
||||
this.name = "MTPUnsupportedRelayVersionError";
|
||||
this.relayVersion = relayVersion;
|
||||
}
|
||||
}
|
||||
|
||||
export function relayOpeningError(error: unknown, signerId?: bigint): Error {
|
||||
if (error !== null && typeof error === "object") {
|
||||
const structured = error as { code?: unknown; relayVersion?: unknown };
|
||||
if (typeof structured.code === "string") {
|
||||
switch (structured.code) {
|
||||
case "missing-relay-version":
|
||||
return new MTPMissingRelayVersionError();
|
||||
case "unsupported-relay-version":
|
||||
if (
|
||||
typeof structured.relayVersion === "bigint" ||
|
||||
typeof structured.relayVersion === "number" ||
|
||||
typeof structured.relayVersion === "string"
|
||||
) {
|
||||
return new MTPUnsupportedRelayVersionError(
|
||||
inputU64(structured.relayVersion, "relayVersion"),
|
||||
);
|
||||
}
|
||||
break;
|
||||
case "no-matching-recipient":
|
||||
return new Error("Unable to decrypt protected value with supplied recipient keyrings");
|
||||
case "not-final-recipient":
|
||||
return new Error("relay content is addressed to a different final recipient");
|
||||
case "reserved-application-type":
|
||||
return new Error("relay application message type is reserved for MTP control");
|
||||
case "signature-policy-mismatch":
|
||||
return new MTPSignatureVerificationError("policy-rejected", signerId);
|
||||
case "unsupported-signature-suite":
|
||||
return new MTPSignatureVerificationError("unsupported-suite", signerId);
|
||||
case "invalid-signature":
|
||||
return new MTPSignatureVerificationError("invalid-signature", signerId);
|
||||
case "signer-id-mismatch":
|
||||
return new Error("relay signer ID mismatch");
|
||||
case "purpose-mismatch":
|
||||
return new Error("relay protection purpose mismatch");
|
||||
case "signer-key-not-found":
|
||||
return signerKeysUnavailable(signerId);
|
||||
case "replay":
|
||||
return new Error("relay message was already accepted");
|
||||
case "resource-limit":
|
||||
return new MTPResourceLimitError();
|
||||
}
|
||||
}
|
||||
}
|
||||
return error instanceof Error ? error : new Error(String(error));
|
||||
}
|
||||
|
||||
export interface MTPRelayMetadataState {
|
||||
frame: ParsedFrame;
|
||||
native: RawBindings.WasmVerifiedRelayMetadata;
|
||||
relayVersion: number;
|
||||
signerId: bigint;
|
||||
finalRecipientId: bigint;
|
||||
messageId: string;
|
||||
createdAt: bigint;
|
||||
hasMetadata: boolean;
|
||||
metadata?: MTPDataValue;
|
||||
encryptedContent: Uint8Array;
|
||||
signerPublicKeys: Uint8Array[];
|
||||
matchedSignerKeyIndex: number;
|
||||
signaturePolicy: MTPSignatureVerificationPolicy;
|
||||
receiveLimits?: MTPReceiveLimits;
|
||||
receiveLimitsExplicit: boolean;
|
||||
disposed: boolean;
|
||||
finalizerToken: object;
|
||||
}
|
||||
|
||||
export const relayMetadataState = new WeakMap<
|
||||
MTPVerifiedRelayMetadata,
|
||||
MTPRelayMetadataState
|
||||
>();
|
||||
|
||||
const relayMetadataFinalizer = new FinalizationRegistry<
|
||||
RawBindings.WasmVerifiedRelayMetadata
|
||||
>((native) => {
|
||||
try {
|
||||
native.free();
|
||||
} catch {
|
||||
// The WASM instance may already have been torn down during page unload.
|
||||
}
|
||||
});
|
||||
|
||||
export const RELAY_METADATA_TOKEN = Symbol("mtp-authenticated-relay-metadata");
|
||||
|
||||
export class MTPVerifiedRelayMetadata {
|
||||
constructor(
|
||||
token: typeof RELAY_METADATA_TOKEN,
|
||||
state: MTPRelayMetadataState,
|
||||
) {
|
||||
if (token !== RELAY_METADATA_TOKEN) {
|
||||
throw new Error("relay metadata must be created by authenticated opening");
|
||||
}
|
||||
relayMetadataState.set(this, state);
|
||||
}
|
||||
|
||||
private get state(): MTPRelayMetadataState {
|
||||
const state = relayMetadataState.get(this);
|
||||
if (!state) throw new Error("relay metadata authentication state is missing");
|
||||
if (state.disposed) throw new Error("relay metadata has been disposed");
|
||||
return state;
|
||||
}
|
||||
|
||||
dispose(): void {
|
||||
const state = relayMetadataState.get(this);
|
||||
if (!state || state.disposed) return;
|
||||
state.disposed = true;
|
||||
relayMetadataFinalizer.unregister(state.finalizerToken);
|
||||
try {
|
||||
state.native.free();
|
||||
} catch {
|
||||
// The WASM instance may already have been torn down during page unload.
|
||||
}
|
||||
}
|
||||
|
||||
free(): void {
|
||||
this.dispose();
|
||||
}
|
||||
|
||||
[Symbol.dispose](): void {
|
||||
this.dispose();
|
||||
}
|
||||
|
||||
get frame(): ParsedFrame {
|
||||
return cloneParsedFrame(this.state.frame);
|
||||
}
|
||||
get signerId(): bigint {
|
||||
return this.state.signerId;
|
||||
}
|
||||
get relayVersion(): number {
|
||||
return this.state.relayVersion;
|
||||
}
|
||||
get finalRecipientId(): bigint {
|
||||
return this.state.finalRecipientId;
|
||||
}
|
||||
get messageId(): string {
|
||||
return this.state.messageId;
|
||||
}
|
||||
get createdAt(): bigint {
|
||||
return this.state.createdAt;
|
||||
}
|
||||
get metadata(): MTPDataValue | undefined {
|
||||
return this.state.hasMetadata
|
||||
? (cloneParsedValue(this.state.metadata) as MTPDataValue)
|
||||
: undefined;
|
||||
}
|
||||
get encryptedContent(): Uint8Array {
|
||||
return this.state.encryptedContent.slice();
|
||||
}
|
||||
get signerPublicKeys(): Uint8Array[] {
|
||||
return this.state.signerPublicKeys.map((bundle) => bundle.slice());
|
||||
}
|
||||
get matchedSignerKeyIndex(): number {
|
||||
return this.state.matchedSignerKeyIndex;
|
||||
}
|
||||
get matchedSignerPublicKey(): Uint8Array {
|
||||
const key = this.state.signerPublicKeys[this.state.matchedSignerKeyIndex];
|
||||
if (!key) throw new Error("relay verification matched an unavailable signer key");
|
||||
return key.slice();
|
||||
}
|
||||
get signaturePolicy(): MTPSignatureVerificationPolicy {
|
||||
return this.state.signaturePolicy;
|
||||
}
|
||||
}
|
||||
|
||||
export function registerRelayMetadata(
|
||||
metadata: MTPVerifiedRelayMetadata,
|
||||
native: RawBindings.WasmVerifiedRelayMetadata,
|
||||
finalizerToken: object,
|
||||
): void {
|
||||
relayMetadataFinalizer.register(metadata, native, finalizerToken);
|
||||
}
|
||||
|
||||
export type { MTPVerifiedRelayContent };
|
||||
|
|
@ -1,234 +0,0 @@
|
|||
import type { MTPRequestOptions, ParsedFrame, Unsubscribe } from "./client.js";
|
||||
import type { MTPCommunicationType } from "../type-map/index.js";
|
||||
|
||||
export interface MTPSchema<Input = unknown, Output = Input> {
|
||||
readonly _input: Input;
|
||||
readonly _output: Output;
|
||||
parseAsync(value: unknown): Promise<Output>;
|
||||
}
|
||||
|
||||
export interface MTPSchemaPair<
|
||||
Request extends MTPSchema = MTPSchema,
|
||||
Response extends MTPSchema = MTPSchema,
|
||||
> {
|
||||
request: Request;
|
||||
response: Response;
|
||||
}
|
||||
|
||||
export type MTPSchemaRegistry = Record<string, MTPSchemaPair>;
|
||||
export type MTPNoSchemas = Record<never, never>;
|
||||
|
||||
export type MTPSchemaInput<Schema extends MTPSchema> = Schema["_input"];
|
||||
export type MTPSchemaOutput<Schema extends MTPSchema> = Schema["_output"];
|
||||
export type MTPMessageType<Registry extends MTPSchemaRegistry> =
|
||||
keyof Registry & string;
|
||||
|
||||
export type MTPFrame<Data = ParsedFrame["data"]> = {
|
||||
id?: number;
|
||||
type: string;
|
||||
data: Data;
|
||||
sender?: ParsedFrame["sender"];
|
||||
receiver?: ParsedFrame["receiver"];
|
||||
raw?: ParsedFrame["raw"];
|
||||
};
|
||||
|
||||
export type MTPTypedFrame<Data = ParsedFrame["data"]> = MTPFrame<Data>;
|
||||
|
||||
export type MTPResponseFrame<
|
||||
Registry extends MTPSchemaRegistry,
|
||||
Type extends MTPMessageType<Registry>,
|
||||
> = MTPTypedFrame<MTPSchemaOutput<Registry[Type]["response"]>>;
|
||||
|
||||
export type MTPRequestData<
|
||||
Registry extends MTPSchemaRegistry,
|
||||
Type extends MTPMessageType<Registry>,
|
||||
> = MTPSchemaInput<Registry[Type]["request"]>;
|
||||
|
||||
export type MTPRequestFunction<Registry extends MTPSchemaRegistry> = <
|
||||
Type extends MTPMessageType<Registry>,
|
||||
>(
|
||||
type: Type,
|
||||
data?: MTPRequestData<Registry, Type>,
|
||||
options?: MTPRequestOptions,
|
||||
) => Promise<MTPResponseFrame<Registry, Type>>;
|
||||
|
||||
export type MTPSubscriptionFunction<Registry extends MTPSchemaRegistry> = <
|
||||
Type extends MTPMessageType<Registry>,
|
||||
>(
|
||||
type: Type,
|
||||
handler: (message: MTPResponseFrame<Registry, Type>) => void | Promise<void>,
|
||||
) => Unsubscribe;
|
||||
|
||||
export class MTPValidationError extends Error {
|
||||
readonly phase: "request" | "response" | "subscription";
|
||||
readonly messageType: string;
|
||||
readonly frame?: MTPFrame;
|
||||
|
||||
constructor(
|
||||
phase: MTPValidationError["phase"],
|
||||
messageType: string,
|
||||
cause: unknown,
|
||||
frame?: MTPFrame,
|
||||
) {
|
||||
super(`${phase} validation failed for ${messageType}`, { cause });
|
||||
this.name = "MTPValidationError";
|
||||
this.phase = phase;
|
||||
this.messageType = messageType;
|
||||
this.frame = frame;
|
||||
}
|
||||
}
|
||||
|
||||
export class MTPProtocolError extends Error {
|
||||
readonly type: string;
|
||||
readonly id: number | undefined;
|
||||
readonly communicationType: string;
|
||||
readonly requestId: number | undefined;
|
||||
readonly errorType: string | undefined;
|
||||
readonly frame: MTPFrame;
|
||||
|
||||
constructor(frame: MTPFrame) {
|
||||
const errorType =
|
||||
frame.data &&
|
||||
typeof frame.data === "object" &&
|
||||
!Array.isArray(frame.data) &&
|
||||
typeof (frame.data as Record<string, unknown>).ErrorType === "string"
|
||||
? ((frame.data as Record<string, unknown>).ErrorType as string)
|
||||
: undefined;
|
||||
super(errorType ? `${frame.type}: ${errorType}` : frame.type);
|
||||
this.name = "MTPProtocolError";
|
||||
this.type = frame.type;
|
||||
this.id = frame.id;
|
||||
this.communicationType = frame.type;
|
||||
this.requestId = frame.id;
|
||||
this.errorType = errorType;
|
||||
this.frame = frame;
|
||||
}
|
||||
}
|
||||
|
||||
export interface MTPProtocolOptions<Registry extends MTPSchemaRegistry> {
|
||||
schemas: Registry;
|
||||
throwProtocolErrors?: boolean;
|
||||
onValidationError?: (error: MTPValidationError) => void;
|
||||
}
|
||||
|
||||
function isErrorFrame(frame: MTPFrame): boolean {
|
||||
return frame.type.startsWith("Error");
|
||||
}
|
||||
|
||||
export class MTPProtocol<Registry extends MTPSchemaRegistry> {
|
||||
readonly schemas: Registry;
|
||||
readonly #throwProtocolErrors: boolean;
|
||||
readonly #onValidationError:
|
||||
| ((error: MTPValidationError) => void)
|
||||
| undefined;
|
||||
|
||||
constructor(options: MTPProtocolOptions<Registry>) {
|
||||
this.schemas = options.schemas;
|
||||
this.#throwProtocolErrors = options.throwProtocolErrors ?? false;
|
||||
this.#onValidationError = options.onValidationError;
|
||||
}
|
||||
|
||||
async parseRequest<Type extends MTPMessageType<Registry>>(
|
||||
type: Type,
|
||||
data: MTPRequestData<Registry, Type> | undefined,
|
||||
): Promise<MTPSchemaOutput<Registry[Type]["request"]>> {
|
||||
try {
|
||||
return await this.schemas[type].request.parseAsync(data);
|
||||
} catch (error) {
|
||||
throw new MTPValidationError("request", type, error);
|
||||
}
|
||||
}
|
||||
|
||||
async parseResponse<Type extends MTPMessageType<Registry>>(
|
||||
requestedType: Type,
|
||||
frame: MTPFrame,
|
||||
phase: "response" | "subscription" = "response",
|
||||
): Promise<MTPResponseFrame<Registry, Type>> {
|
||||
if (isErrorFrame(frame)) {
|
||||
if (phase === "response" && this.#throwProtocolErrors) {
|
||||
throw new MTPProtocolError(frame);
|
||||
}
|
||||
return frame as MTPResponseFrame<Registry, Type>;
|
||||
}
|
||||
|
||||
const schema =
|
||||
this.schemas[frame.type]?.response ??
|
||||
this.schemas[requestedType].response;
|
||||
try {
|
||||
const data = await schema.parseAsync(frame.data);
|
||||
return { ...frame, data } as MTPResponseFrame<Registry, Type>;
|
||||
} catch (error) {
|
||||
throw new MTPValidationError(
|
||||
phase,
|
||||
frame.type || requestedType,
|
||||
error,
|
||||
frame,
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
reportValidationError(error: unknown): void {
|
||||
if (error instanceof MTPValidationError) {
|
||||
this.#onValidationError?.(error);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
export interface MTPProxyAdapter {
|
||||
request(
|
||||
type: MTPCommunicationType,
|
||||
data: Record<string, unknown>,
|
||||
options?: MTPRequestOptions,
|
||||
): Promise<MTPFrame>;
|
||||
subscribe(
|
||||
type: MTPCommunicationType,
|
||||
handler: (message: MTPFrame) => void,
|
||||
): Unsubscribe;
|
||||
}
|
||||
|
||||
export class MTPProxyConnection<Registry extends MTPSchemaRegistry> {
|
||||
readonly #adapter: MTPProxyAdapter;
|
||||
readonly #protocol: MTPProtocol<Registry>;
|
||||
|
||||
constructor(adapter: MTPProxyAdapter, options: MTPProtocolOptions<Registry>) {
|
||||
this.#adapter = adapter;
|
||||
this.#protocol = new MTPProtocol(options);
|
||||
}
|
||||
|
||||
async request<Type extends MTPMessageType<Registry>>(
|
||||
type: Type,
|
||||
data?: MTPRequestData<Registry, Type>,
|
||||
options?: MTPRequestOptions,
|
||||
): Promise<MTPResponseFrame<Registry, Type>> {
|
||||
const parsed = await this.#protocol.parseRequest(type, data);
|
||||
const response = await this.#adapter.request(
|
||||
type,
|
||||
(parsed ?? {}) as Record<string, unknown>,
|
||||
options,
|
||||
);
|
||||
return await this.#protocol.parseResponse(type, response);
|
||||
}
|
||||
|
||||
subscribe<Type extends MTPMessageType<Registry>>(
|
||||
type: Type,
|
||||
handler: (
|
||||
message: MTPResponseFrame<Registry, Type>,
|
||||
) => void | Promise<void>,
|
||||
): Unsubscribe {
|
||||
let active = true;
|
||||
const unsubscribe = this.#adapter.subscribe(type, (message) => {
|
||||
void this.#protocol.parseResponse(type, message, "subscription").then(
|
||||
(parsed) => {
|
||||
if (active) void handler(parsed);
|
||||
},
|
||||
(error) => {
|
||||
this.#protocol.reportValidationError(error);
|
||||
},
|
||||
);
|
||||
});
|
||||
return () => {
|
||||
active = false;
|
||||
unsubscribe();
|
||||
};
|
||||
}
|
||||
}
|
||||
|
|
@ -92,14 +92,8 @@ export function signatureVerificationPolicyValue(
|
|||
return bindings.mtp_protection_signature_suite_ed25519();
|
||||
case "dual":
|
||||
return bindings.mtp_protection_signature_suite_dual();
|
||||
case "any-supported": {
|
||||
const compatibility = (
|
||||
bindings as typeof bindings & {
|
||||
mtp_protection_signature_suite_any_supported?: () => number;
|
||||
}
|
||||
).mtp_protection_signature_suite_any_supported;
|
||||
return compatibility?.() ?? 0;
|
||||
}
|
||||
case "any-supported":
|
||||
return 0;
|
||||
}
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -1,27 +0,0 @@
|
|||
export async function withTimeout<T>(
|
||||
promise: Promise<T>,
|
||||
timeoutMs: number | undefined,
|
||||
message: string,
|
||||
cancel?: () => void,
|
||||
): Promise<T> {
|
||||
if (!timeoutMs) {
|
||||
return await promise;
|
||||
}
|
||||
|
||||
let timeoutId: ReturnType<typeof setTimeout> | undefined;
|
||||
try {
|
||||
return await Promise.race([
|
||||
promise,
|
||||
new Promise<never>((_resolve, reject) => {
|
||||
timeoutId = setTimeout(() => {
|
||||
cancel?.();
|
||||
reject(new Error(message));
|
||||
}, timeoutMs);
|
||||
}),
|
||||
]);
|
||||
} finally {
|
||||
if (timeoutId !== undefined) {
|
||||
clearTimeout(timeoutId);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
|
@ -1,27 +0,0 @@
|
|||
import initWasm from "mtp/raw";
|
||||
|
||||
type WasmInitInput = Parameters<typeof initWasm>[0];
|
||||
type WasmExports = Awaited<ReturnType<typeof initWasm>>;
|
||||
type WasmInitializer = (input?: WasmInitInput) => Promise<WasmExports>;
|
||||
|
||||
export function createWasmInitializer(
|
||||
initialize: WasmInitializer = initWasm,
|
||||
): WasmInitializer {
|
||||
let wasmInitPromise: Promise<WasmExports> | undefined;
|
||||
|
||||
/**
|
||||
* Keep the successful WASM singleton, but make a failed attempt retryable.
|
||||
* A rejected promise is never retained in the module cache.
|
||||
*/
|
||||
return (input?: WasmInitInput): Promise<WasmExports> => {
|
||||
if (!wasmInitPromise) {
|
||||
wasmInitPromise = initialize(input).catch((error) => {
|
||||
wasmInitPromise = undefined;
|
||||
throw error;
|
||||
});
|
||||
}
|
||||
return wasmInitPromise;
|
||||
}
|
||||
}
|
||||
|
||||
export const initWasmOnce = createWasmInitializer();
|
||||
|
|
@ -1,20 +0,0 @@
|
|||
import assert from "node:assert/strict";
|
||||
import { test } from "node:test";
|
||||
import { createWasmInitializer } from "../dist/sdk/wasm-init.js";
|
||||
|
||||
test("WASM initialization can retry after a rejected attempt", async () => {
|
||||
let attempts = 0;
|
||||
const expected = { initialized: true };
|
||||
const init = createWasmInitializer(async () => {
|
||||
attempts += 1;
|
||||
if (attempts === 1) {
|
||||
throw new Error("initialization failed");
|
||||
}
|
||||
return expected;
|
||||
});
|
||||
|
||||
await assert.rejects(init(), /initialization failed/);
|
||||
assert.equal(await init(), expected);
|
||||
assert.equal(await init(), expected);
|
||||
assert.equal(attempts, 2);
|
||||
});
|
||||
|
|
@ -1,18 +1,12 @@
|
|||
use crate::ConnectionHandle;
|
||||
use crate::framing::RetryClassifier;
|
||||
#[cfg(feature = "pipes")]
|
||||
use crate::pipe::PipeReader;
|
||||
use mtp_codec::{CommunicationValue, DecodeError, DecodeLimits, EncodeLimits, TypeMap};
|
||||
use mtp_codec::{CommunicationValue, DecodeLimits, TypeMap};
|
||||
use mtp_common::CommunicationError;
|
||||
#[cfg(feature = "pipes")]
|
||||
use mtp_common::{FirstFrameDisposition, classify_first_frame};
|
||||
#[cfg(feature = "pipes")]
|
||||
use std::collections::HashSet;
|
||||
use std::ops::Deref;
|
||||
use std::sync::Arc;
|
||||
use std::sync::atomic::{AtomicU64, Ordering};
|
||||
use tokio::sync::{Mutex, Notify, RwLock, Semaphore, mpsc};
|
||||
use tokio::time::{Duration, Instant, sleep, timeout, timeout_at};
|
||||
use tokio::time::{Duration, sleep, timeout};
|
||||
use tracing::{debug, info, instrument, trace, warn};
|
||||
use wtransport::Connection;
|
||||
|
||||
|
|
@ -25,63 +19,6 @@ pub enum TransportEvent<R = wtransport::RecvStream> {
|
|||
|
||||
const APPLICATION_CLOSE_REASON: &str = "mtp-close";
|
||||
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
|
||||
pub enum DecodeRejectionClass {
|
||||
Malformed,
|
||||
ResourceLimit,
|
||||
DuplicateField,
|
||||
}
|
||||
|
||||
pub fn classify_decode_error(error: &DecodeError) -> DecodeRejectionClass {
|
||||
match error {
|
||||
DecodeError::MalformedEncoding => DecodeRejectionClass::Malformed,
|
||||
DecodeError::DepthLimit
|
||||
| DecodeError::ValueCountLimit
|
||||
| DecodeError::BlobLimit
|
||||
| DecodeError::AllocationLimit
|
||||
| DecodeError::RecipientLimit => DecodeRejectionClass::ResourceLimit,
|
||||
DecodeError::DuplicateField => DecodeRejectionClass::DuplicateField,
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Default)]
|
||||
pub struct DecodeRejectionCounters {
|
||||
malformed: AtomicU64,
|
||||
resource_limit: AtomicU64,
|
||||
duplicate_field: AtomicU64,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Copy, Default, PartialEq, Eq)]
|
||||
pub struct DecodeRejectionCounts {
|
||||
pub malformed: u64,
|
||||
pub resource_limit: u64,
|
||||
pub duplicate_field: u64,
|
||||
}
|
||||
|
||||
impl DecodeRejectionCounters {
|
||||
pub(crate) fn record(&self, error: &DecodeError) {
|
||||
match classify_decode_error(error) {
|
||||
DecodeRejectionClass::Malformed => {
|
||||
self.malformed.fetch_add(1, Ordering::Relaxed);
|
||||
}
|
||||
DecodeRejectionClass::ResourceLimit => {
|
||||
self.resource_limit.fetch_add(1, Ordering::Relaxed);
|
||||
}
|
||||
DecodeRejectionClass::DuplicateField => {
|
||||
self.duplicate_field.fetch_add(1, Ordering::Relaxed);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) fn snapshot(&self) -> DecodeRejectionCounts {
|
||||
DecodeRejectionCounts {
|
||||
malformed: self.malformed.load(Ordering::Relaxed),
|
||||
resource_limit: self.resource_limit.load(Ordering::Relaxed),
|
||||
duplicate_field: self.duplicate_field.load(Ordering::Relaxed),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
||||
pub enum SendMode {
|
||||
PersistentStream,
|
||||
|
|
@ -198,62 +135,27 @@ impl Policy {
|
|||
}
|
||||
}
|
||||
|
||||
/// A validated policy snapshot used after a public [`Policy`] crosses into a
|
||||
/// transport implementation. `Policy` intentionally remains a plain public
|
||||
/// struct for source compatibility, so callers can construct it directly and
|
||||
/// bypass builder methods. Every transport constructor takes this snapshot
|
||||
/// before creating channels or semaphores.
|
||||
#[derive(Debug, Clone, Copy)]
|
||||
pub(crate) struct RuntimePolicy(Policy);
|
||||
|
||||
impl RuntimePolicy {
|
||||
pub(crate) fn from_public(policy: &Policy) -> Self {
|
||||
let mut policy = *policy;
|
||||
policy.receiver_queue_capacity = policy.receiver_queue_capacity.max(1);
|
||||
policy.max_concurrent_stream_tasks = policy.max_concurrent_stream_tasks.max(1);
|
||||
Self(policy)
|
||||
}
|
||||
}
|
||||
|
||||
impl Deref for RuntimePolicy {
|
||||
type Target = Policy;
|
||||
|
||||
fn deref(&self) -> &Self::Target {
|
||||
&self.0
|
||||
}
|
||||
}
|
||||
|
||||
enum ReceivedFrame {
|
||||
Message(CommunicationValue),
|
||||
ClosedByPeer,
|
||||
Idle,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
||||
enum SenderState {
|
||||
Open,
|
||||
Closing,
|
||||
Closed,
|
||||
}
|
||||
|
||||
#[derive(Clone)]
|
||||
pub struct Sender {
|
||||
send_guard: Arc<Mutex<()>>,
|
||||
stream_guard: Arc<Mutex<Option<wtransport::SendStream>>>,
|
||||
state: Arc<Mutex<SenderState>>,
|
||||
handle: Arc<ConnectionHandle>,
|
||||
connection: Connection,
|
||||
policy: Arc<RuntimePolicy>,
|
||||
policy: Arc<Policy>,
|
||||
type_map: Arc<RwLock<TypeMap>>,
|
||||
}
|
||||
|
||||
impl Sender {
|
||||
pub fn new(connection: Connection, handle: Arc<ConnectionHandle>, policy: Arc<Policy>) -> Self {
|
||||
let policy = Arc::new(RuntimePolicy::from_public(&policy));
|
||||
Self {
|
||||
send_guard: Arc::new(Mutex::new(())),
|
||||
stream_guard: Arc::new(Mutex::new(None)),
|
||||
state: Arc::new(Mutex::new(SenderState::Open)),
|
||||
handle,
|
||||
connection,
|
||||
policy,
|
||||
|
|
@ -272,11 +174,7 @@ impl Sender {
|
|||
data: &CommunicationValue,
|
||||
policy: &Policy,
|
||||
) -> Result<(), CommunicationError> {
|
||||
let bytes = data
|
||||
.to_bytes_with_limits(EncodeLimits::for_transport_message_size(
|
||||
policy.max_message_size,
|
||||
))
|
||||
.map_err(|_| CommunicationError::Encode)?;
|
||||
let bytes = data.to_bytes().map_err(|_| CommunicationError::Encode)?;
|
||||
if bytes.len() as u64 > policy.max_message_size
|
||||
|| bytes.len() as u64 >= policy.close_frame_len as u64
|
||||
{
|
||||
|
|
@ -292,15 +190,15 @@ impl Sender {
|
|||
Ok(Ok(())) => Ok(()),
|
||||
Ok(Err(wtransport::error::StreamWriteError::Stopped(code))) => {
|
||||
warn!("[Sender] write failed: peer sent STOP_SENDING (error code {code})");
|
||||
Err(CommunicationError::DeliveryUnknown)
|
||||
Err(CommunicationError::StreamClosed)
|
||||
}
|
||||
Ok(Err(other)) => {
|
||||
warn!("[Sender] write failed: {other}");
|
||||
Err(CommunicationError::DeliveryUnknown)
|
||||
Err(CommunicationError::StreamError)
|
||||
}
|
||||
Err(_) => {
|
||||
warn!("[Sender] write timed out (len={})", bytes.len());
|
||||
Err(CommunicationError::DeliveryUnknown)
|
||||
Err(CommunicationError::StreamError)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
|
@ -371,11 +269,11 @@ impl Sender {
|
|||
return Ok(());
|
||||
}
|
||||
|
||||
let err = match res {
|
||||
Ok(()) => return Ok(()),
|
||||
Err(error) => error,
|
||||
};
|
||||
if !RetryClassifier::retry_persistent_stream(&err) {
|
||||
let err = res.err().unwrap_or(CommunicationError::StreamError);
|
||||
if !matches!(
|
||||
err,
|
||||
CommunicationError::StreamError | CommunicationError::StreamClosed
|
||||
) {
|
||||
return Err(err);
|
||||
}
|
||||
*stream_opt = None;
|
||||
|
|
@ -402,15 +300,15 @@ impl Sender {
|
|||
Ok(Ok(())) => Ok(()),
|
||||
Ok(Err(wtransport::error::StreamWriteError::Stopped(code))) => {
|
||||
warn!("[Sender] finish failed: peer sent STOP_SENDING (error code {code})");
|
||||
Err(CommunicationError::DeliveryUnknown)
|
||||
Err(CommunicationError::StreamClosed)
|
||||
}
|
||||
Ok(Err(other)) => {
|
||||
warn!("[Sender] finish failed: {other}");
|
||||
Err(CommunicationError::DeliveryUnknown)
|
||||
Err(CommunicationError::StreamError)
|
||||
}
|
||||
Err(_) => {
|
||||
warn!("[Sender] finish timed out");
|
||||
Err(CommunicationError::DeliveryUnknown)
|
||||
Err(CommunicationError::StreamError)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
|
@ -458,24 +356,20 @@ impl Sender {
|
|||
|
||||
#[instrument(skip(self, data), level = "trace")]
|
||||
pub async fn send(&self, data: &CommunicationValue) -> Result<(), CommunicationError> {
|
||||
let _send_lock = self.send_guard.lock().await;
|
||||
|
||||
{
|
||||
let state = self.state.lock().await;
|
||||
if *state != SenderState::Open {
|
||||
if self.handle.is_closed() {
|
||||
return Err(self
|
||||
.handle
|
||||
.close_reason()
|
||||
.unwrap_or(CommunicationError::StreamClosed));
|
||||
}
|
||||
.unwrap_or(CommunicationError::UseAfterClosed));
|
||||
}
|
||||
|
||||
let _send_lock = self.send_guard.lock().await;
|
||||
|
||||
if self.connection.quic_connection().close_reason().is_some() {
|
||||
let reason = self
|
||||
.handle
|
||||
.close_reason()
|
||||
.unwrap_or(CommunicationError::StreamClosed);
|
||||
*self.state.lock().await = SenderState::Closed;
|
||||
self.handle.close(Some(reason.clone()));
|
||||
return Err(reason);
|
||||
}
|
||||
|
|
@ -508,7 +402,6 @@ impl Sender {
|
|||
if self.connection.quic_connection().close_reason().is_some()
|
||||
|| matches!(normalized, CommunicationError::StreamClosed)
|
||||
{
|
||||
*self.state.lock().await = SenderState::Closed;
|
||||
self.handle.close(Some(normalized.clone()));
|
||||
}
|
||||
|
||||
|
|
@ -520,12 +413,6 @@ impl Sender {
|
|||
#[instrument(skip(self), level = "trace")]
|
||||
pub async fn finish_stream(&self) -> Result<(), CommunicationError> {
|
||||
let _send_lock = self.send_guard.lock().await;
|
||||
if *self.state.lock().await != SenderState::Open {
|
||||
return Err(self
|
||||
.handle
|
||||
.close_reason()
|
||||
.unwrap_or(CommunicationError::StreamClosed));
|
||||
}
|
||||
let mut stream_opt = self.stream_guard.lock().await;
|
||||
if let Some(mut stream) = stream_opt.take() {
|
||||
match timeout(self.policy.write_timeout, stream.finish()).await {
|
||||
|
|
@ -558,12 +445,11 @@ impl Sender {
|
|||
pipe_id: u32,
|
||||
description: &str,
|
||||
) -> Result<crate::pipe::PipeWriter, CommunicationError> {
|
||||
let _send_lock = self.send_guard.lock().await;
|
||||
if *self.state.lock().await != SenderState::Open {
|
||||
if self.handle.is_closed() {
|
||||
return Err(self
|
||||
.handle
|
||||
.close_reason()
|
||||
.unwrap_or(CommunicationError::StreamClosed));
|
||||
.unwrap_or(CommunicationError::UseAfterClosed));
|
||||
}
|
||||
|
||||
if self.connection.quic_connection().close_reason().is_some() {
|
||||
|
|
@ -600,26 +486,13 @@ impl Sender {
|
|||
let handle = self.handle.clone();
|
||||
let policy = self.policy.clone();
|
||||
let stream_guard = self.stream_guard.clone();
|
||||
let send_guard = self.send_guard.clone();
|
||||
let state = self.state.clone();
|
||||
|
||||
tokio::spawn(async move {
|
||||
let _send_lock = send_guard.lock().await;
|
||||
{
|
||||
let mut sender_state = state.lock().await;
|
||||
if *sender_state != SenderState::Open {
|
||||
return;
|
||||
}
|
||||
*sender_state = SenderState::Closing;
|
||||
}
|
||||
|
||||
if connection.quic_connection().close_reason().is_some() || handle.is_closed() {
|
||||
*state.lock().await = SenderState::Closed;
|
||||
handle.close(Some(CommunicationError::StreamClosed));
|
||||
return;
|
||||
}
|
||||
|
||||
{
|
||||
if let Some(mut stream) = stream_guard.lock().await.take() {
|
||||
match timeout(policy.write_timeout, stream.finish()).await {
|
||||
Ok(Ok(())) => {}
|
||||
|
|
@ -632,15 +505,12 @@ impl Sender {
|
|||
Err(_) => warn!("[Sender] persistent stream finish timed out"),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
let _ = Self::send_close_frame(&connection, &policy).await;
|
||||
|
||||
*state.lock().await = SenderState::Closed;
|
||||
handle.close(Some(CommunicationError::StreamClosed));
|
||||
info!(target = "mtp.transport", "connection closed");
|
||||
|
||||
drop(_send_lock);
|
||||
sleep(policy.force_close_delay).await;
|
||||
if connection.quic_connection().close_reason().is_none() {
|
||||
connection.quic_connection().close(
|
||||
|
|
@ -658,23 +528,13 @@ impl Sender {
|
|||
let connection = self.connection.clone();
|
||||
let handle = self.handle.clone();
|
||||
let policy = self.policy.clone();
|
||||
let _send_lock = self.send_guard.lock().await;
|
||||
{
|
||||
let mut state = self.state.lock().await;
|
||||
if *state != SenderState::Open {
|
||||
return;
|
||||
}
|
||||
*state = SenderState::Closing;
|
||||
}
|
||||
let mut stream_opt = self.stream_guard.lock().await;
|
||||
|
||||
if connection.quic_connection().close_reason().is_some() || handle.is_closed() {
|
||||
*self.state.lock().await = SenderState::Closed;
|
||||
handle.close(Some(CommunicationError::StreamClosed));
|
||||
return;
|
||||
}
|
||||
|
||||
{
|
||||
let mut stream_opt = self.stream_guard.lock().await;
|
||||
if let Some(mut stream) = stream_opt.take() {
|
||||
let close_bytes = policy.close_frame_len.to_be_bytes();
|
||||
let close_write = async {
|
||||
|
|
@ -693,13 +553,10 @@ impl Sender {
|
|||
} else {
|
||||
let _ = Self::send_close_frame(&connection, &policy).await;
|
||||
}
|
||||
}
|
||||
|
||||
*self.state.lock().await = SenderState::Closed;
|
||||
handle.close(Some(CommunicationError::StreamClosed));
|
||||
info!(target = "mtp.transport", "connection closed");
|
||||
|
||||
drop(_send_lock);
|
||||
sleep(policy.force_close_delay).await;
|
||||
if connection.quic_connection().close_reason().is_none() {
|
||||
connection.quic_connection().close(
|
||||
|
|
@ -748,9 +605,6 @@ struct ReceiverInner {
|
|||
queue_notify: Arc<Notify>,
|
||||
max_message_size: Arc<AtomicU64>,
|
||||
type_map: Arc<RwLock<TypeMap>>,
|
||||
decode_rejections: Arc<DecodeRejectionCounters>,
|
||||
#[cfg(feature = "pipes")]
|
||||
expected_pipes: Arc<std::sync::Mutex<HashSet<u32>>>,
|
||||
}
|
||||
|
||||
impl Clone for Receiver {
|
||||
|
|
@ -772,26 +626,12 @@ impl Drop for Receiver {
|
|||
#[derive(Clone, Default)]
|
||||
struct PingControl {
|
||||
pong_sender: Option<Sender>,
|
||||
pong_observer: Option<mpsc::Sender<CommunicationValue>>,
|
||||
expected_pong_id: Option<u32>,
|
||||
}
|
||||
|
||||
impl PingControl {
|
||||
fn accepts_pong(&mut self, id: Option<u32>) -> bool {
|
||||
if self.expected_pong_id == id && id.is_some() {
|
||||
self.expected_pong_id = None;
|
||||
true
|
||||
} else {
|
||||
false
|
||||
}
|
||||
}
|
||||
pong_observer: Option<mpsc::UnboundedSender<CommunicationValue>>,
|
||||
}
|
||||
|
||||
impl Receiver {
|
||||
pub fn new(connection: Connection, handle: Arc<ConnectionHandle>, policy: Arc<Policy>) -> Self {
|
||||
let policy = Arc::new(RuntimePolicy::from_public(&policy));
|
||||
let max_message_size = policy.max_message_size;
|
||||
Self::new_with_max_message_size(connection, handle, policy, max_message_size)
|
||||
Self::new_with_max_message_size(connection, handle, policy.clone(), policy.max_message_size)
|
||||
}
|
||||
|
||||
#[cfg(feature = "host")]
|
||||
|
|
@ -800,7 +640,6 @@ impl Receiver {
|
|||
handle: Arc<ConnectionHandle>,
|
||||
policy: Arc<Policy>,
|
||||
) -> Self {
|
||||
let policy = Arc::new(RuntimePolicy::from_public(&policy));
|
||||
let initial_max = policy
|
||||
.handshake_max_message_size
|
||||
.min(policy.max_message_size);
|
||||
|
|
@ -810,7 +649,7 @@ impl Receiver {
|
|||
fn new_with_max_message_size(
|
||||
connection: Connection,
|
||||
handle: Arc<ConnectionHandle>,
|
||||
policy: Arc<RuntimePolicy>,
|
||||
policy: Arc<Policy>,
|
||||
initial_max_message_size: u64,
|
||||
) -> Self {
|
||||
#[cfg(feature = "pipes")]
|
||||
|
|
@ -835,12 +674,6 @@ impl Receiver {
|
|||
let accept_max_message_size = max_message_size.clone();
|
||||
let type_map = Arc::new(RwLock::new(TypeMap::latest()));
|
||||
let accept_type_map = type_map.clone();
|
||||
let decode_rejections = Arc::new(DecodeRejectionCounters::default());
|
||||
let accept_decode_rejections = decode_rejections.clone();
|
||||
#[cfg(feature = "pipes")]
|
||||
let expected_pipes = Arc::new(std::sync::Mutex::new(HashSet::new()));
|
||||
#[cfg(feature = "pipes")]
|
||||
let accept_expected_pipes = expected_pipes.clone();
|
||||
let stream_limit = Arc::new(Semaphore::new(policy.max_concurrent_stream_tasks.max(1)));
|
||||
let accept_stream_limit = stream_limit.clone();
|
||||
debug!(
|
||||
|
|
@ -909,9 +742,6 @@ impl Receiver {
|
|||
let stream_ping_control = accept_ping_control.clone();
|
||||
let stream_max_message_size = accept_max_message_size.clone();
|
||||
let stream_type_map = accept_type_map.clone();
|
||||
let stream_decode_rejections = accept_decode_rejections.clone();
|
||||
#[cfg(feature = "pipes")]
|
||||
let stream_expected_pipes = accept_expected_pipes.clone();
|
||||
|
||||
tokio::spawn(async move {
|
||||
let _permit = permit;
|
||||
|
|
@ -931,14 +761,7 @@ impl Receiver {
|
|||
}
|
||||
|
||||
let frame_limit = stream_max_message_size.load(Ordering::Relaxed);
|
||||
match Self::read_one_frame(
|
||||
&mut s,
|
||||
&stream_policy,
|
||||
frame_limit,
|
||||
&stream_decode_rejections,
|
||||
)
|
||||
.await
|
||||
{
|
||||
match Self::read_one_frame(&mut s, &stream_policy, frame_limit).await {
|
||||
Ok(ReceivedFrame::Message(mut msg)) => {
|
||||
let negotiated_type_map =
|
||||
stream_type_map.read().await.clone();
|
||||
|
|
@ -947,32 +770,17 @@ impl Receiver {
|
|||
|
||||
#[cfg(feature = "pipes")]
|
||||
{
|
||||
if frame_count == 1 {
|
||||
let is_pipe_request = msg.is_type(
|
||||
mtp_codec::CommunicationType::PipeRequest,
|
||||
if msg.is_type(mtp_codec::CommunicationType::PipeRequest)
|
||||
&& frame_count == 1
|
||||
{
|
||||
let Some(pipe_id) = msg.id().filter(|id| *id != 0) else {
|
||||
let error = CommunicationError::Other(
|
||||
"PipeRequest frame must contain a non-zero id".into(),
|
||||
);
|
||||
let pipe_id = msg.id().filter(|id| *id != 0);
|
||||
let pipe_is_expected = is_pipe_request && pipe_id.is_some_and(|pipe_id| {
|
||||
stream_expected_pipes
|
||||
.lock()
|
||||
.is_ok_and(|mut expected| expected.remove(&pipe_id))
|
||||
});
|
||||
let disposition = match classify_first_frame(
|
||||
is_pipe_request,
|
||||
msg.id(),
|
||||
pipe_is_expected,
|
||||
) {
|
||||
Ok(disposition) => disposition,
|
||||
Err(error) => {
|
||||
let _ = msg_tx_stream
|
||||
.send(Err(error.clone()))
|
||||
.await;
|
||||
let _ = msg_tx_stream.send(Err(error.clone())).await;
|
||||
stream_handle.close(Some(error));
|
||||
break;
|
||||
}
|
||||
};
|
||||
|
||||
if let FirstFrameDisposition::Pipe(pipe_id) = disposition {
|
||||
let description = msg
|
||||
.get_str(mtp_codec::DataType::Description)
|
||||
.unwrap_or("")
|
||||
|
|
@ -996,19 +804,25 @@ impl Receiver {
|
|||
break;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
let pong_sender = if msg.is_type(mtp_codec::CommunicationType::Ping) {
|
||||
stream_ping_control
|
||||
.read()
|
||||
.await
|
||||
let control = {
|
||||
let control = stream_ping_control.read().await;
|
||||
if msg.is_type(mtp_codec::CommunicationType::Ping) {
|
||||
control
|
||||
.pong_sender
|
||||
.clone()
|
||||
.map(|sender| (Some(sender), None))
|
||||
} else if msg.is_type(mtp_codec::CommunicationType::Pong) {
|
||||
control
|
||||
.pong_observer
|
||||
.clone()
|
||||
.map(|observer| (None, Some(observer)))
|
||||
} else {
|
||||
None
|
||||
}
|
||||
};
|
||||
|
||||
if let Some(sender) = pong_sender {
|
||||
if let Some((Some(sender), _)) = control {
|
||||
let mut pong = CommunicationValue::new_with_type_map(
|
||||
mtp_codec::CommunicationType::Pong,
|
||||
&negotiated_type_map,
|
||||
|
|
@ -1030,18 +844,8 @@ impl Receiver {
|
|||
continue;
|
||||
}
|
||||
|
||||
if msg.is_type(mtp_codec::CommunicationType::Pong) {
|
||||
let observer = {
|
||||
let mut control = stream_ping_control.write().await;
|
||||
if control.accepts_pong(msg.id()) {
|
||||
control.pong_observer.clone()
|
||||
} else {
|
||||
None
|
||||
}
|
||||
};
|
||||
if let Some(observer) = observer {
|
||||
let _ = observer.try_send(msg);
|
||||
}
|
||||
if let Some((_, Some(observer))) = control {
|
||||
let _ = observer.send(msg);
|
||||
continue;
|
||||
}
|
||||
|
||||
|
|
@ -1138,9 +942,6 @@ impl Receiver {
|
|||
queue_notify,
|
||||
max_message_size,
|
||||
type_map,
|
||||
decode_rejections,
|
||||
#[cfg(feature = "pipes")]
|
||||
expected_pipes,
|
||||
}),
|
||||
}
|
||||
}
|
||||
|
|
@ -1157,34 +958,6 @@ impl Receiver {
|
|||
*self.inner.type_map.write().await = type_map.clone();
|
||||
}
|
||||
|
||||
#[cfg(feature = "pipes")]
|
||||
pub fn expect_pipe(&self, pipe_id: u32) -> Result<(), CommunicationError> {
|
||||
if pipe_id == 0 {
|
||||
return Err(CommunicationError::Other("pipe id must be non-zero".into()));
|
||||
}
|
||||
self.inner
|
||||
.expected_pipes
|
||||
.lock()
|
||||
.map_err(|_| CommunicationError::Other("expected pipe state is unavailable".into()))?
|
||||
.insert(pipe_id);
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[cfg(feature = "pipes")]
|
||||
pub fn cancel_expected_pipe(&self, pipe_id: u32) {
|
||||
if let Ok(mut expected) = self.inner.expected_pipes.lock() {
|
||||
expected.remove(&pipe_id);
|
||||
}
|
||||
}
|
||||
|
||||
/// Return local counts for frames rejected by the structured decoder.
|
||||
///
|
||||
/// These counters are intentionally local-only; peers continue to receive
|
||||
/// the generic protocol parse failure.
|
||||
pub fn decode_rejection_counts(&self) -> DecodeRejectionCounts {
|
||||
self.inner.decode_rejections.snapshot()
|
||||
}
|
||||
|
||||
/* Respond to reserved Ping frames without exposing them to application I/O. */
|
||||
pub fn respond_to_pings(&self, sender: Sender) {
|
||||
if let Ok(mut control) = self.inner.ping_control.try_write() {
|
||||
|
|
@ -1194,21 +967,16 @@ impl Receiver {
|
|||
}
|
||||
}
|
||||
|
||||
/* Route only the currently expected reserved Pong through a bounded observer. */
|
||||
pub async fn observe_pongs_bounded(&self, observer: mpsc::Sender<CommunicationValue>) {
|
||||
/* Route reserved Pong frames to a connection-level observer. */
|
||||
pub async fn observe_pongs(&self, observer: mpsc::UnboundedSender<CommunicationValue>) {
|
||||
self.inner.ping_control.write().await.pong_observer = Some(observer);
|
||||
}
|
||||
|
||||
pub async fn set_expected_pong_id(&self, expected_pong_id: Option<u32>) {
|
||||
self.inner.ping_control.write().await.expected_pong_id = expected_pong_id;
|
||||
}
|
||||
|
||||
#[instrument(skip(stream, policy, decode_rejections), level = "trace")]
|
||||
#[instrument(skip(stream, policy), level = "trace")]
|
||||
async fn read_one_frame(
|
||||
stream: &mut wtransport::RecvStream,
|
||||
policy: &RuntimePolicy,
|
||||
policy: &Policy,
|
||||
max_message_size: u64,
|
||||
decode_rejections: &DecodeRejectionCounters,
|
||||
) -> Result<ReceivedFrame, CommunicationError> {
|
||||
use wtransport::error::{StreamReadError, StreamReadExactError};
|
||||
|
||||
|
|
@ -1243,7 +1011,6 @@ impl Receiver {
|
|||
if len == policy.close_frame_len {
|
||||
return Ok(ReceivedFrame::ClosedByPeer);
|
||||
}
|
||||
let deadline = Instant::now() + policy.read_timeout;
|
||||
|
||||
let body_len = len as usize;
|
||||
let frame_len = body_len
|
||||
|
|
@ -1253,29 +1020,29 @@ impl Receiver {
|
|||
return Err(CommunicationError::MessageTooLarge);
|
||||
}
|
||||
|
||||
// The length has already been checked against the admitted frame
|
||||
// limit, so reserve one bounded framing buffer and decode it without a
|
||||
// second prefix-plus-body allocation/copy.
|
||||
let mut frame = Vec::new();
|
||||
frame
|
||||
.try_reserve_exact(frame_len)
|
||||
// Grow in bounded chunks instead of trusting the peer's length prefix
|
||||
// enough to allocate the complete frame up front.
|
||||
let mut buf = Vec::new();
|
||||
buf.try_reserve(body_len.min(16 * 1024))
|
||||
.map_err(|_| CommunicationError::MessageTooLarge)?;
|
||||
frame.extend_from_slice(&len_buf);
|
||||
frame.resize(frame_len, 0);
|
||||
let mut body_offset = 4usize;
|
||||
while body_offset < frame_len {
|
||||
let chunk_len = (frame_len - body_offset).min(16 * 1024);
|
||||
match timeout_at(
|
||||
deadline,
|
||||
stream.read_exact(&mut frame[body_offset..body_offset + chunk_len]),
|
||||
while buf.len() < body_len {
|
||||
let chunk_len = (body_len - buf.len()).min(16 * 1024);
|
||||
let mut chunk = [0u8; 16 * 1024];
|
||||
match timeout(
|
||||
policy.read_timeout,
|
||||
stream.read_exact(&mut chunk[..chunk_len]),
|
||||
)
|
||||
.await
|
||||
{
|
||||
Ok(Ok(())) => body_offset += chunk_len,
|
||||
Ok(Ok(())) => {
|
||||
buf.try_reserve(chunk_len)
|
||||
.map_err(|_| CommunicationError::MessageTooLarge)?;
|
||||
buf.extend_from_slice(&chunk[..chunk_len]);
|
||||
}
|
||||
Ok(Err(StreamReadExactError::FinishedEarly(n))) => {
|
||||
warn!(
|
||||
"[Receiver] body read ended early ({}/{body_len} bytes): stream closed by peer",
|
||||
body_offset.saturating_sub(4) + n
|
||||
buf.len() + n
|
||||
);
|
||||
return Err(CommunicationError::StreamError);
|
||||
}
|
||||
|
|
@ -1296,21 +1063,14 @@ impl Receiver {
|
|||
}
|
||||
}
|
||||
|
||||
let message = match CommunicationValue::try_from_bytes_with_limits(
|
||||
let mut frame = Vec::with_capacity(frame_len);
|
||||
frame.extend_from_slice(&len_buf);
|
||||
frame.extend_from_slice(&buf);
|
||||
let message = CommunicationValue::from_bytes_with_limits(
|
||||
&frame,
|
||||
DecodeLimits::for_transport_message_size(max_message_size),
|
||||
) {
|
||||
Ok(message) => message,
|
||||
Err(error) => {
|
||||
decode_rejections.record(&error);
|
||||
warn!(
|
||||
?error,
|
||||
class = ?classify_decode_error(&error),
|
||||
"[Receiver] rejected frame during bounded decode"
|
||||
);
|
||||
return Err(CommunicationError::ParseCommunicationValue);
|
||||
}
|
||||
};
|
||||
)
|
||||
.map_err(|_| CommunicationError::ParseCommunicationValue)?;
|
||||
|
||||
Ok(ReceivedFrame::Message(message))
|
||||
}
|
||||
|
|
@ -1507,63 +1267,4 @@ mod tests {
|
|||
let debug_str = format!("{:?}", p);
|
||||
assert!(debug_str.contains("Policy"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn runtime_policy_normalizes_zero_channel_and_task_limits() {
|
||||
let policy = Policy {
|
||||
receiver_queue_capacity: 0,
|
||||
max_concurrent_stream_tasks: 0,
|
||||
..Policy::default()
|
||||
};
|
||||
|
||||
let runtime = RuntimePolicy::from_public(&policy);
|
||||
|
||||
assert_eq!(runtime.receiver_queue_capacity, 1);
|
||||
assert_eq!(runtime.max_concurrent_stream_tasks, 1);
|
||||
assert_eq!(policy.receiver_queue_capacity, 0);
|
||||
assert_eq!(policy.max_concurrent_stream_tasks, 0);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn ping_control_accepts_only_the_current_expected_id() {
|
||||
let mut control = PingControl {
|
||||
expected_pong_id: Some(7),
|
||||
..PingControl::default()
|
||||
};
|
||||
|
||||
assert!(!control.accepts_pong(Some(6)));
|
||||
assert_eq!(control.expected_pong_id, Some(7));
|
||||
assert!(control.accepts_pong(Some(7)));
|
||||
assert_eq!(control.expected_pong_id, None);
|
||||
assert!(!control.accepts_pong(Some(7)));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn decode_rejection_classes_are_stable_and_counted() {
|
||||
assert_eq!(
|
||||
classify_decode_error(&DecodeError::MalformedEncoding),
|
||||
DecodeRejectionClass::Malformed
|
||||
);
|
||||
assert_eq!(
|
||||
classify_decode_error(&DecodeError::AllocationLimit),
|
||||
DecodeRejectionClass::ResourceLimit
|
||||
);
|
||||
assert_eq!(
|
||||
classify_decode_error(&DecodeError::DuplicateField),
|
||||
DecodeRejectionClass::DuplicateField
|
||||
);
|
||||
|
||||
let counters = DecodeRejectionCounters::default();
|
||||
counters.record(&DecodeError::MalformedEncoding);
|
||||
counters.record(&DecodeError::DepthLimit);
|
||||
counters.record(&DecodeError::DuplicateField);
|
||||
assert_eq!(
|
||||
counters.snapshot(),
|
||||
DecodeRejectionCounts {
|
||||
malformed: 1,
|
||||
resource_limit: 1,
|
||||
duplicate_field: 1,
|
||||
}
|
||||
);
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -2,26 +2,22 @@ use mtp_common::CommunicationError;
|
|||
use std::net::SocketAddr;
|
||||
use std::sync::{
|
||||
Arc,
|
||||
atomic::{AtomicBool, AtomicU64, Ordering},
|
||||
atomic::{AtomicBool, Ordering},
|
||||
};
|
||||
use tokio::sync::watch;
|
||||
|
||||
#[derive(Debug)]
|
||||
pub struct ConnectionHandle {
|
||||
connection_id: u64,
|
||||
closed: AtomicBool,
|
||||
close_tx: watch::Sender<Option<CommunicationError>>,
|
||||
close_rx: watch::Receiver<Option<CommunicationError>>,
|
||||
remote_addr: Option<SocketAddr>,
|
||||
}
|
||||
|
||||
static NEXT_CONNECTION_ID: AtomicU64 = AtomicU64::new(1);
|
||||
|
||||
impl ConnectionHandle {
|
||||
pub fn new() -> Self {
|
||||
let (close_tx, close_rx) = watch::channel(None);
|
||||
Self {
|
||||
connection_id: NEXT_CONNECTION_ID.fetch_add(1, Ordering::Relaxed).max(1),
|
||||
closed: AtomicBool::new(false),
|
||||
close_tx,
|
||||
close_rx,
|
||||
|
|
@ -39,11 +35,6 @@ impl ConnectionHandle {
|
|||
self.remote_addr
|
||||
}
|
||||
|
||||
/// Stable process-local identifier for authentication-rate-limit scopes.
|
||||
pub fn connection_id(&self) -> u64 {
|
||||
self.connection_id
|
||||
}
|
||||
|
||||
pub fn is_open(&self) -> bool {
|
||||
!self.closed.load(Ordering::SeqCst)
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1,21 +1,7 @@
|
|||
use crate::{Policy, TransportSendStream};
|
||||
use mtp_codec::{CommunicationValue, EncodeLimits};
|
||||
use mtp_codec::CommunicationValue;
|
||||
use mtp_common::CommunicationError;
|
||||
|
||||
/// Classifies failures that may be recovered by replacing a persistent
|
||||
/// application stream. Encoding and frame-size failures are deterministic and
|
||||
/// must reach the caller without opening more streams.
|
||||
pub(crate) struct RetryClassifier;
|
||||
|
||||
impl RetryClassifier {
|
||||
pub(crate) fn retry_persistent_stream(error: &CommunicationError) -> bool {
|
||||
matches!(
|
||||
error,
|
||||
CommunicationError::StreamError | CommunicationError::StreamClosed
|
||||
)
|
||||
}
|
||||
}
|
||||
|
||||
/// Writes the canonical self-framed MTP value used by every transport.
|
||||
///
|
||||
/// `CommunicationValue` already begins with the four-byte body length. The
|
||||
|
|
@ -26,11 +12,7 @@ pub(crate) async fn write_frame<S: TransportSendStream>(
|
|||
value: &CommunicationValue,
|
||||
policy: &Policy,
|
||||
) -> Result<(), CommunicationError> {
|
||||
let bytes = value
|
||||
.to_bytes_with_limits(EncodeLimits::for_transport_message_size(
|
||||
policy.max_message_size,
|
||||
))
|
||||
.map_err(|_| CommunicationError::Encode)?;
|
||||
let bytes = value.to_bytes().map_err(|_| CommunicationError::Encode)?;
|
||||
if bytes.len() as u64 > policy.max_message_size
|
||||
|| bytes.len() as u64 >= policy.close_frame_len as u64
|
||||
{
|
||||
|
|
@ -82,10 +64,6 @@ mod tests {
|
|||
async fn finish(&mut self) -> Result<(), CommunicationError> {
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn reset(&mut self, _code: u32) -> Result<(), CommunicationError> {
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
|
|
|
|||
|
|
@ -5,27 +5,21 @@
|
|||
//! wrappers while the framing implementation below is shared by adapters.
|
||||
|
||||
use crate::{
|
||||
Policy, TransportConnection, TransportRecvStream, TransportSendStream,
|
||||
connection::{DecodeRejectionCounters, RuntimePolicy, classify_decode_error},
|
||||
framing::{RetryClassifier, write_frame},
|
||||
Policy, TransportConnection, TransportRecvStream, TransportSendStream, framing::write_frame,
|
||||
};
|
||||
use mtp_codec::{CommunicationValue, DataType, DecodeLimits, TypeMap};
|
||||
use mtp_common::{CommunicationError, FirstFrameDisposition, classify_first_frame};
|
||||
#[cfg(feature = "pipes")]
|
||||
use std::collections::HashSet;
|
||||
use mtp_codec::{CommunicationValue, DecodeLimits, TypeMap};
|
||||
use mtp_common::CommunicationError;
|
||||
use std::sync::Arc;
|
||||
#[cfg(feature = "pipes")]
|
||||
use std::sync::Mutex as StdMutex;
|
||||
use std::sync::atomic::{AtomicU64, Ordering};
|
||||
use tokio::sync::{Mutex, Notify, RwLock, Semaphore, mpsc};
|
||||
use tokio::time::{Instant, timeout, timeout_at};
|
||||
use tokio::sync::{Mutex, RwLock, Semaphore, mpsc};
|
||||
use tokio::time::timeout;
|
||||
|
||||
#[cfg(feature = "pipes")]
|
||||
use crate::pipe::{PipeReader, PipeWriter};
|
||||
|
||||
pub struct GenericSender<C: TransportConnection> {
|
||||
connection: C,
|
||||
policy: Arc<RuntimePolicy>,
|
||||
policy: Arc<Policy>,
|
||||
persistent: Arc<Mutex<Option<C::SendStream>>>,
|
||||
send_lock: Arc<Mutex<()>>,
|
||||
type_map: Arc<RwLock<TypeMap>>,
|
||||
|
|
@ -45,7 +39,6 @@ impl<C: TransportConnection> Clone for GenericSender<C> {
|
|||
|
||||
impl<C: TransportConnection> GenericSender<C> {
|
||||
pub fn new(connection: C, policy: Arc<Policy>) -> Self {
|
||||
let policy = Arc::new(RuntimePolicy::from_public(&policy));
|
||||
Self {
|
||||
connection,
|
||||
policy,
|
||||
|
|
@ -71,15 +64,6 @@ impl<C: TransportConnection> GenericSender<C> {
|
|||
if self.connection.close_reason().is_some() {
|
||||
return Err(CommunicationError::StreamClosed);
|
||||
}
|
||||
if let Some(version) = value.get_str(DataType::Version) {
|
||||
tracing::debug!(
|
||||
message_type = ?value.get_type(),
|
||||
version,
|
||||
connected = ?value.get_data(DataType::Connected),
|
||||
client_id = ?value.get_data(DataType::Id),
|
||||
"sending MTP handshake response frame"
|
||||
);
|
||||
}
|
||||
match self.policy.send_mode {
|
||||
crate::SendMode::SingleStreamPerMessage => {
|
||||
let mut stream = self.open().await?;
|
||||
|
|
@ -89,10 +73,9 @@ impl<C: TransportConnection> GenericSender<C> {
|
|||
)
|
||||
.await
|
||||
.map_err(|_| CommunicationError::StreamError)??;
|
||||
match timeout(self.policy.write_timeout, stream.finish()).await {
|
||||
Ok(Ok(())) => Ok(()),
|
||||
Ok(Err(_)) | Err(_) => Err(CommunicationError::DeliveryUnknown),
|
||||
}
|
||||
timeout(self.policy.write_timeout, stream.finish())
|
||||
.await
|
||||
.map_err(|_| CommunicationError::StreamError)?
|
||||
}
|
||||
crate::SendMode::PersistentStream => {
|
||||
let mut stream = self.persistent.lock().await;
|
||||
|
|
@ -101,30 +84,20 @@ impl<C: TransportConnection> GenericSender<C> {
|
|||
if stream.is_none() {
|
||||
*stream = Some(self.open().await?);
|
||||
}
|
||||
let result = match stream.as_mut() {
|
||||
Some(stream) => timeout(
|
||||
let result = timeout(
|
||||
self.policy.write_timeout,
|
||||
write_frame(stream, value, &self.policy),
|
||||
write_frame(stream.as_mut().unwrap(), value, &self.policy),
|
||||
)
|
||||
.await
|
||||
.map_err(|_| CommunicationError::StreamError)
|
||||
.and_then(|result| result),
|
||||
None => Err(CommunicationError::StreamError),
|
||||
};
|
||||
.and_then(|r| r);
|
||||
if result.is_ok() {
|
||||
return Ok(());
|
||||
}
|
||||
let error = match result {
|
||||
Ok(()) => return Ok(()),
|
||||
Err(error) => error,
|
||||
};
|
||||
if !RetryClassifier::retry_persistent_stream(&error) {
|
||||
return Err(error);
|
||||
return result;
|
||||
}
|
||||
*stream = None;
|
||||
attempts += 1;
|
||||
if attempts > self.policy.persistent_stream_max_retries {
|
||||
return Err(error);
|
||||
return result;
|
||||
}
|
||||
tokio::time::sleep(
|
||||
self.policy.persistent_stream_retry_backoff * attempts as u32,
|
||||
|
|
@ -141,7 +114,6 @@ impl<C: TransportConnection> GenericSender<C> {
|
|||
pipe_id: u32,
|
||||
description: &str,
|
||||
) -> Result<PipeWriter<C::SendStream>, CommunicationError> {
|
||||
let _send_lock = self.send_lock.lock().await;
|
||||
if self.connection.close_reason().is_some() {
|
||||
return Err(CommunicationError::StreamClosed);
|
||||
}
|
||||
|
|
@ -207,10 +179,6 @@ pub struct GenericReceiver<C: TransportConnection> {
|
|||
ping_sender: Arc<RwLock<Option<GenericSender<C>>>>,
|
||||
max_message_size: Arc<AtomicU64>,
|
||||
type_map: Arc<RwLock<TypeMap>>,
|
||||
queue_notify: Arc<Notify>,
|
||||
decode_rejections: Arc<DecodeRejectionCounters>,
|
||||
#[cfg(feature = "pipes")]
|
||||
expected_pipes: Arc<StdMutex<HashSet<u32>>>,
|
||||
_accept_task: Arc<tokio::task::JoinHandle<()>>,
|
||||
}
|
||||
|
||||
|
|
@ -224,10 +192,6 @@ impl<C: TransportConnection> Clone for GenericReceiver<C> {
|
|||
ping_sender: self.ping_sender.clone(),
|
||||
max_message_size: self.max_message_size.clone(),
|
||||
type_map: self.type_map.clone(),
|
||||
queue_notify: self.queue_notify.clone(),
|
||||
decode_rejections: self.decode_rejections.clone(),
|
||||
#[cfg(feature = "pipes")]
|
||||
expected_pipes: self.expected_pipes.clone(),
|
||||
_accept_task: self._accept_task.clone(),
|
||||
}
|
||||
}
|
||||
|
|
@ -243,7 +207,6 @@ impl<C: TransportConnection> Drop for GenericReceiver<C> {
|
|||
|
||||
impl<C: TransportConnection> GenericReceiver<C> {
|
||||
pub fn new(connection: C, policy: Arc<Policy>) -> Self {
|
||||
let policy = Arc::new(RuntimePolicy::from_public(&policy));
|
||||
let (tx, rx) = mpsc::channel(policy.receiver_queue_capacity);
|
||||
#[cfg(feature = "pipes")]
|
||||
let (pipe_tx, pipe_rx) = mpsc::channel(policy.receiver_queue_capacity);
|
||||
|
|
@ -259,14 +222,6 @@ impl<C: TransportConnection> GenericReceiver<C> {
|
|||
let task_max_message_size = max_message_size.clone();
|
||||
let type_map = Arc::new(RwLock::new(TypeMap::latest()));
|
||||
let task_type_map = type_map.clone();
|
||||
let queue_notify = Arc::new(Notify::new());
|
||||
let task_queue_notify = queue_notify.clone();
|
||||
let decode_rejections = Arc::new(DecodeRejectionCounters::default());
|
||||
let task_decode_rejections = decode_rejections.clone();
|
||||
#[cfg(feature = "pipes")]
|
||||
let expected_pipes = Arc::new(StdMutex::new(HashSet::new()));
|
||||
#[cfg(feature = "pipes")]
|
||||
let task_expected_pipes = expected_pipes.clone();
|
||||
let task_accept_task_tx = tx.clone();
|
||||
#[cfg(feature = "pipes")]
|
||||
let task_accept_task_pipe_tx = pipe_tx.clone();
|
||||
|
|
@ -282,10 +237,8 @@ impl<C: TransportConnection> GenericReceiver<C> {
|
|||
#[cfg(not(feature = "pipes"))]
|
||||
let cap_full = task_accept_task_tx.capacity() == 0;
|
||||
|
||||
let notified = task_queue_notify.notified();
|
||||
tokio::pin!(notified);
|
||||
if cap_full {
|
||||
notified.await;
|
||||
tokio::time::sleep(std::time::Duration::from_millis(1)).await;
|
||||
continue;
|
||||
}
|
||||
|
||||
|
|
@ -324,14 +277,11 @@ impl<C: TransportConnection> GenericReceiver<C> {
|
|||
let ping_sender = task_ping_sender.clone();
|
||||
let connection = task_connection.clone();
|
||||
let type_map = task_type_map.clone();
|
||||
let decode_rejections = task_decode_rejections.clone();
|
||||
#[cfg(feature = "pipes")]
|
||||
let expected_pipes = task_expected_pipes.clone();
|
||||
tokio::spawn(async move {
|
||||
let _permit = permit;
|
||||
let mut stream = stream;
|
||||
let mut frames = 0usize;
|
||||
'stream: loop {
|
||||
loop {
|
||||
if policy
|
||||
.max_frames_per_stream
|
||||
.is_some_and(|max| frames >= max)
|
||||
|
|
@ -354,21 +304,9 @@ impl<C: TransportConnection> GenericReceiver<C> {
|
|||
policy.application_close_code,
|
||||
b"frame header read error",
|
||||
);
|
||||
break 'stream;
|
||||
break;
|
||||
}
|
||||
Err(_) => {
|
||||
if frames == 0 {
|
||||
tracing::warn!(
|
||||
timeout = ?policy.read_timeout,
|
||||
"MTP receive stream timed out before its first complete frame"
|
||||
);
|
||||
} else {
|
||||
tracing::debug!(
|
||||
frames,
|
||||
timeout = ?policy.read_timeout,
|
||||
"MTP receive stream idle timeout"
|
||||
);
|
||||
}
|
||||
break;
|
||||
}
|
||||
}
|
||||
|
|
@ -376,7 +314,6 @@ impl<C: TransportConnection> GenericReceiver<C> {
|
|||
if len == policy.close_frame_len {
|
||||
break;
|
||||
}
|
||||
let deadline = Instant::now() + policy.read_timeout;
|
||||
let frame_limit = max_message_size.load(Ordering::Relaxed);
|
||||
let body_len = len as usize;
|
||||
let frame_len = match body_len.checked_add(4) {
|
||||
|
|
@ -393,28 +330,29 @@ impl<C: TransportConnection> GenericReceiver<C> {
|
|||
connection.close(policy.application_close_code, b"frame too large");
|
||||
break;
|
||||
}
|
||||
let mut frame = Vec::new();
|
||||
if frame.try_reserve_exact(frame_len).is_err() {
|
||||
tracing::warn!(frame_len, "MTP receive stream could not reserve frame");
|
||||
let target_len = body_len;
|
||||
let mut body = Vec::new();
|
||||
if body.try_reserve(target_len.min(16 * 1024)).is_err() {
|
||||
tracing::warn!(
|
||||
target_len,
|
||||
"MTP receive stream could not reserve frame body"
|
||||
);
|
||||
let _ = tx.send(Err(CommunicationError::MessageTooLarge)).await;
|
||||
connection
|
||||
.close(policy.application_close_code, b"frame allocation failed");
|
||||
break;
|
||||
}
|
||||
frame.extend_from_slice(&len.to_be_bytes());
|
||||
frame.resize(frame_len, 0);
|
||||
let mut body_offset = 4usize;
|
||||
while body_offset < frame_len {
|
||||
let chunk_len = (frame_len - body_offset).min(16 * 1024);
|
||||
let body_read = timeout_at(
|
||||
deadline,
|
||||
stream.read_exact(&mut frame[body_offset..body_offset + chunk_len]),
|
||||
while body.len() < target_len {
|
||||
let chunk_len = (target_len - body.len()).min(16 * 1024);
|
||||
let mut chunk = [0u8; 16 * 1024];
|
||||
let body_read = tokio::time::timeout(
|
||||
policy.read_timeout,
|
||||
stream.read_exact(&mut chunk[..chunk_len]),
|
||||
)
|
||||
.await;
|
||||
if !matches!(&body_read, Ok(Ok(()))) {
|
||||
if matches!(&body_read, Ok(Err(CommunicationError::StreamClosed))) {
|
||||
break 'stream;
|
||||
}
|
||||
if !matches!(&body_read, Ok(Ok(())))
|
||||
|| body.try_reserve(chunk_len).is_err()
|
||||
{
|
||||
tracing::warn!(
|
||||
pipe_chunk_len = chunk_len,
|
||||
?body_read,
|
||||
|
|
@ -423,23 +361,24 @@ impl<C: TransportConnection> GenericReceiver<C> {
|
|||
let _ = tx.send(Err(CommunicationError::StreamError)).await;
|
||||
connection
|
||||
.close(policy.application_close_code, b"frame body read error");
|
||||
break 'stream;
|
||||
break;
|
||||
}
|
||||
body_offset += chunk_len;
|
||||
body.extend_from_slice(&chunk[..chunk_len]);
|
||||
}
|
||||
if body.len() != target_len {
|
||||
break;
|
||||
}
|
||||
frames += 1;
|
||||
let mut message = match CommunicationValue::try_from_bytes_with_limits(
|
||||
let mut frame = Vec::with_capacity(frame_len);
|
||||
frame.extend_from_slice(&len.to_be_bytes());
|
||||
frame.extend_from_slice(&body);
|
||||
let mut message = match CommunicationValue::from_bytes_with_limits(
|
||||
&frame,
|
||||
DecodeLimits::for_transport_message_size(frame_limit),
|
||||
) {
|
||||
Ok(message) => message,
|
||||
Err(error) => {
|
||||
tracing::warn!(
|
||||
?error,
|
||||
class = ?classify_decode_error(&error),
|
||||
"MTP receive stream rejected by bounded decode"
|
||||
);
|
||||
decode_rejections.record(&error);
|
||||
Err(_) => {
|
||||
tracing::warn!("MTP receive stream contained an invalid frame");
|
||||
let _ = tx
|
||||
.send(Err(CommunicationError::ParseCommunicationValue))
|
||||
.await;
|
||||
|
|
@ -447,44 +386,25 @@ impl<C: TransportConnection> GenericReceiver<C> {
|
|||
break;
|
||||
}
|
||||
};
|
||||
tracing::debug!(
|
||||
frames,
|
||||
frame_len,
|
||||
message_type = ?message.get_type(),
|
||||
"decoded MTP receive frame"
|
||||
);
|
||||
let negotiated_type_map = type_map.read().await.clone();
|
||||
message.set_type_map(&negotiated_type_map);
|
||||
|
||||
#[cfg(feature = "pipes")]
|
||||
{
|
||||
if frames == 1 {
|
||||
let is_pipe_request =
|
||||
message.is_type(mtp_codec::CommunicationType::PipeRequest);
|
||||
let pipe_id = message.id().filter(|id| *id != 0);
|
||||
let pipe_is_expected = is_pipe_request
|
||||
&& pipe_id.is_some_and(|pipe_id| {
|
||||
expected_pipes
|
||||
.lock()
|
||||
.is_ok_and(|mut expected| expected.remove(&pipe_id))
|
||||
});
|
||||
let disposition = match classify_first_frame(
|
||||
is_pipe_request,
|
||||
message.id(),
|
||||
pipe_is_expected,
|
||||
) {
|
||||
Ok(disposition) => disposition,
|
||||
Err(error) => {
|
||||
if message.is_type(mtp_codec::CommunicationType::PipeRequest)
|
||||
&& frames == 1
|
||||
{
|
||||
let Some(pipe_id) = message.id().filter(|id| *id != 0) else {
|
||||
let error = CommunicationError::Other(
|
||||
"PipeRequest frame must contain a non-zero id".into(),
|
||||
);
|
||||
let _ = tx.send(Err(error.clone())).await;
|
||||
connection.close(
|
||||
policy.application_close_code,
|
||||
b"pipe request missing id",
|
||||
);
|
||||
break;
|
||||
}
|
||||
};
|
||||
|
||||
if let FirstFrameDisposition::Pipe(pipe_id) = disposition {
|
||||
let description = message
|
||||
.get_str(mtp_codec::DataType::Description)
|
||||
.unwrap_or("")
|
||||
|
|
@ -504,7 +424,6 @@ impl<C: TransportConnection> GenericReceiver<C> {
|
|||
return;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
if message.is_type(mtp_codec::CommunicationType::Ping) {
|
||||
if let Some(sender) = ping_sender.read().await.clone() {
|
||||
|
|
@ -544,10 +463,6 @@ impl<C: TransportConnection> GenericReceiver<C> {
|
|||
ping_sender,
|
||||
max_message_size,
|
||||
type_map,
|
||||
queue_notify,
|
||||
decode_rejections,
|
||||
#[cfg(feature = "pipes")]
|
||||
expected_pipes,
|
||||
_accept_task: Arc::new(accept_task),
|
||||
}
|
||||
}
|
||||
|
|
@ -555,25 +470,6 @@ impl<C: TransportConnection> GenericReceiver<C> {
|
|||
*self.ping_sender.write().await = Some(sender);
|
||||
}
|
||||
|
||||
#[cfg(feature = "pipes")]
|
||||
pub fn expect_pipe(&self, pipe_id: u32) -> Result<(), CommunicationError> {
|
||||
if pipe_id == 0 {
|
||||
return Err(CommunicationError::Other("pipe id must be non-zero".into()));
|
||||
}
|
||||
self.expected_pipes
|
||||
.lock()
|
||||
.map_err(|_| CommunicationError::Other("expected pipe state is unavailable".into()))?
|
||||
.insert(pipe_id);
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[cfg(feature = "pipes")]
|
||||
pub fn cancel_expected_pipe(&self, pipe_id: u32) {
|
||||
if let Ok(mut expected) = self.expected_pipes.lock() {
|
||||
expected.remove(&pipe_id);
|
||||
}
|
||||
}
|
||||
|
||||
/// Switch from the handshake frame limit to the application frame limit.
|
||||
pub fn set_max_message_size(&self, max_message_size: u64) {
|
||||
self.max_message_size
|
||||
|
|
@ -584,23 +480,13 @@ impl<C: TransportConnection> GenericReceiver<C> {
|
|||
pub async fn set_type_map(&self, type_map: &TypeMap) {
|
||||
*self.type_map.write().await = type_map.clone();
|
||||
}
|
||||
|
||||
/// Return local counts for frames rejected by the structured decoder.
|
||||
pub fn decode_rejection_counts(&self) -> crate::DecodeRejectionCounts {
|
||||
self.decode_rejections.snapshot()
|
||||
}
|
||||
pub async fn receive(&self) -> Result<CommunicationValue, CommunicationError> {
|
||||
let result = self
|
||||
.incoming
|
||||
self.incoming
|
||||
.lock()
|
||||
.await
|
||||
.recv()
|
||||
.await
|
||||
.unwrap_or(Err(CommunicationError::StreamClosed));
|
||||
if result.is_ok() {
|
||||
self.queue_notify.notify_one();
|
||||
}
|
||||
result
|
||||
.unwrap_or(Err(CommunicationError::StreamClosed))
|
||||
}
|
||||
|
||||
#[cfg(feature = "pipes")]
|
||||
|
|
@ -612,10 +498,7 @@ impl<C: TransportConnection> GenericReceiver<C> {
|
|||
tokio::select! {
|
||||
msg = incoming.recv() => {
|
||||
match msg {
|
||||
Some(Ok(val)) => {
|
||||
self.queue_notify.notify_one();
|
||||
Ok(crate::TransportEvent::Message(val))
|
||||
}
|
||||
Some(Ok(val)) => Ok(crate::TransportEvent::Message(val)),
|
||||
Some(Err(e)) => Err(e),
|
||||
None => Err(self
|
||||
.connection
|
||||
|
|
@ -625,10 +508,7 @@ impl<C: TransportConnection> GenericReceiver<C> {
|
|||
}
|
||||
pipe = pipes.recv() => {
|
||||
match pipe {
|
||||
Some(reader) => {
|
||||
self.queue_notify.notify_one();
|
||||
Ok(crate::TransportEvent::Pipe(reader))
|
||||
}
|
||||
Some(reader) => Ok(crate::TransportEvent::Pipe(reader)),
|
||||
None => Err(self
|
||||
.connection
|
||||
.close_reason()
|
||||
|
|
@ -640,17 +520,12 @@ impl<C: TransportConnection> GenericReceiver<C> {
|
|||
|
||||
#[cfg(feature = "pipes")]
|
||||
pub async fn receive_pipe(&self) -> Result<PipeReader<C::RecvStream>, CommunicationError> {
|
||||
let result = self
|
||||
.pipes
|
||||
self.pipes
|
||||
.lock()
|
||||
.await
|
||||
.recv()
|
||||
.await
|
||||
.ok_or(CommunicationError::StreamClosed);
|
||||
if result.is_ok() {
|
||||
self.queue_notify.notify_one();
|
||||
}
|
||||
result
|
||||
.ok_or(CommunicationError::StreamClosed)
|
||||
}
|
||||
|
||||
#[cfg(feature = "pipes")]
|
||||
|
|
@ -659,10 +534,7 @@ impl<C: TransportConnection> GenericReceiver<C> {
|
|||
) -> Result<Option<PipeReader<C::RecvStream>>, CommunicationError> {
|
||||
match self.pipes.try_lock() {
|
||||
Ok(mut rx) => match rx.try_recv() {
|
||||
Ok(reader) => {
|
||||
self.queue_notify.notify_one();
|
||||
Ok(Some(reader))
|
||||
}
|
||||
Ok(reader) => Ok(Some(reader)),
|
||||
Err(mpsc::error::TryRecvError::Empty) => Ok(None),
|
||||
Err(mpsc::error::TryRecvError::Disconnected) => {
|
||||
Err(CommunicationError::StreamClosed)
|
||||
|
|
|
|||
|
|
@ -11,10 +11,7 @@ pub mod encrypted_pipe;
|
|||
#[cfg(feature = "pipes")]
|
||||
pub mod pipe;
|
||||
|
||||
pub use connection::{
|
||||
DecodeRejectionClass, DecodeRejectionCounters, DecodeRejectionCounts, Policy, Receiver,
|
||||
SendMode, Sender, classify_decode_error,
|
||||
};
|
||||
pub use connection::{Policy, Receiver, SendMode, Sender};
|
||||
pub use generic_connection::{GenericReceiver, GenericSender};
|
||||
|
||||
#[cfg(feature = "pipes")]
|
||||
|
|
|
|||
|
|
@ -17,7 +17,6 @@ use mtp_common::CommunicationError;
|
|||
pub trait TransportSendStream: tokio::io::AsyncWrite + Send + Sync {
|
||||
async fn write_all(&mut self, buf: &[u8]) -> Result<(), CommunicationError>;
|
||||
async fn finish(&mut self) -> Result<(), CommunicationError>;
|
||||
fn reset(&mut self, code: u32) -> Result<(), CommunicationError>;
|
||||
}
|
||||
|
||||
/// A readable unidirectional stream suitable for MTP frames.
|
||||
|
|
@ -29,9 +28,6 @@ pub trait TransportSendStream: tokio::io::AsyncWrite + Send + Sync {
|
|||
pub trait TransportRecvStream: tokio::io::AsyncRead + Send + Sync {
|
||||
async fn read_exact(&mut self, buf: &mut [u8]) -> Result<(), CommunicationError>;
|
||||
async fn read_chunk(&mut self, max: usize) -> Result<Option<Vec<u8>>, CommunicationError>;
|
||||
fn stop(self, code: u32) -> Result<(), CommunicationError>
|
||||
where
|
||||
Self: Sized;
|
||||
}
|
||||
|
||||
/// A QUIC/WebTransport connection that provides MTP's unidirectional streams.
|
||||
|
|
@ -51,7 +47,7 @@ impl TransportSendStream for wtransport::SendStream {
|
|||
async fn write_all(&mut self, buf: &[u8]) -> Result<(), CommunicationError> {
|
||||
wtransport::SendStream::write_all(self, buf)
|
||||
.await
|
||||
.map_err(|_| CommunicationError::DeliveryUnknown)
|
||||
.map_err(|_| CommunicationError::StreamError)
|
||||
}
|
||||
|
||||
async fn finish(&mut self) -> Result<(), CommunicationError> {
|
||||
|
|
@ -59,23 +55,14 @@ impl TransportSendStream for wtransport::SendStream {
|
|||
.await
|
||||
.map_err(|_| CommunicationError::StreamError)
|
||||
}
|
||||
|
||||
fn reset(&mut self, code: u32) -> Result<(), CommunicationError> {
|
||||
wtransport::SendStream::reset(self, wtransport::VarInt::from_u32(code))
|
||||
.map_err(|_| CommunicationError::StreamClosed)
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl TransportRecvStream for wtransport::RecvStream {
|
||||
async fn read_exact(&mut self, buf: &mut [u8]) -> Result<(), CommunicationError> {
|
||||
match wtransport::RecvStream::read_exact(self, buf).await {
|
||||
Ok(()) => Ok(()),
|
||||
Err(wtransport::error::StreamReadExactError::FinishedEarly(0)) => {
|
||||
Err(CommunicationError::StreamClosed)
|
||||
}
|
||||
Err(_) => Err(CommunicationError::StreamError),
|
||||
}
|
||||
wtransport::RecvStream::read_exact(self, buf)
|
||||
.await
|
||||
.map_err(|_| CommunicationError::StreamError)
|
||||
}
|
||||
|
||||
async fn read_chunk(&mut self, max: usize) -> Result<Option<Vec<u8>>, CommunicationError> {
|
||||
|
|
@ -89,11 +76,6 @@ impl TransportRecvStream for wtransport::RecvStream {
|
|||
Err(_) => Err(CommunicationError::StreamError),
|
||||
}
|
||||
}
|
||||
|
||||
fn stop(self, code: u32) -> Result<(), CommunicationError> {
|
||||
wtransport::RecvStream::stop(self, wtransport::VarInt::from_u32(code));
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
|
|
|
|||
|
|
@ -4,7 +4,7 @@ use async_trait::async_trait;
|
|||
use mtp_codec::CommunicationValue;
|
||||
use mtp_common::CommunicationError;
|
||||
use mtp_transport::{
|
||||
GenericReceiver, GenericSender, Policy, SendMode, TransportConnection, TransportEvent,
|
||||
GenericReceiver, GenericSender, Policy, TransportConnection, TransportEvent,
|
||||
TransportRecvStream, TransportSendStream,
|
||||
};
|
||||
use std::sync::Arc;
|
||||
|
|
@ -53,10 +53,6 @@ impl TransportSendStream for MockSendStream {
|
|||
.await
|
||||
.map_err(|_| CommunicationError::StreamError)
|
||||
}
|
||||
|
||||
fn reset(&mut self, _code: u32) -> Result<(), CommunicationError> {
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
|
||||
struct MockRecvStream {
|
||||
|
|
@ -93,10 +89,6 @@ impl TransportRecvStream for MockRecvStream {
|
|||
Err(_) => Err(CommunicationError::StreamError),
|
||||
}
|
||||
}
|
||||
|
||||
fn stop(self, _code: u32) -> Result<(), CommunicationError> {
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone)]
|
||||
|
|
@ -165,7 +157,6 @@ async fn test_open_pipe_and_receive_reader() -> Result<(), Box<dyn std::error::E
|
|||
let sender = GenericSender::new(conn_a, policy.clone());
|
||||
let receiver = GenericReceiver::new(conn_b, policy);
|
||||
|
||||
receiver.expect_pipe(42)?;
|
||||
let pipe_writer = sender.open_pipe(42, "test-pipe").await?;
|
||||
|
||||
let pipe_reader = receiver.receive_pipe().await?;
|
||||
|
|
@ -183,7 +174,6 @@ async fn test_pipe_raw_data_roundtrip() -> Result<(), Box<dyn std::error::Error>
|
|||
let sender = GenericSender::new(conn_a, policy.clone());
|
||||
let receiver = GenericReceiver::new(conn_b, policy);
|
||||
|
||||
receiver.expect_pipe(1)?;
|
||||
let mut pipe_writer = sender.open_pipe(1, "data-pipe").await?;
|
||||
|
||||
let data = b"hello through the pipe";
|
||||
|
|
@ -205,7 +195,6 @@ async fn test_pipe_large_payload() -> Result<(), Box<dyn std::error::Error>> {
|
|||
let sender = GenericSender::new(conn_a, policy.clone());
|
||||
let receiver = GenericReceiver::new(conn_b, policy);
|
||||
|
||||
receiver.expect_pipe(7)?;
|
||||
let mut pipe_writer = sender.open_pipe(7, "big-pipe").await?;
|
||||
|
||||
let data: Vec<u8> = (0..256 * 1024).map(|i| (i % 256) as u8).collect();
|
||||
|
|
@ -234,7 +223,6 @@ async fn test_receive_event_dispatches_pipe() -> Result<(), Box<dyn std::error::
|
|||
let sender = GenericSender::new(conn_a, policy.clone());
|
||||
let receiver = GenericReceiver::new(conn_b, policy);
|
||||
|
||||
receiver.expect_pipe(99)?;
|
||||
let mut pipe_writer = sender.open_pipe(99, "event-pipe").await?;
|
||||
|
||||
match receiver.receive_event().await? {
|
||||
|
|
@ -273,23 +261,22 @@ async fn test_try_receive_pipe_returns_none_when_empty() -> Result<(), Box<dyn s
|
|||
async fn test_regular_messages_still_work_alongside_pipes() -> Result<(), Box<dyn std::error::Error>>
|
||||
{
|
||||
let (conn_a, conn_b) = mock_connected_pair().await;
|
||||
let policy = Arc::new(Policy::default().with_send_mode(SendMode::SingleStreamPerMessage));
|
||||
let policy = Arc::new(Policy::default());
|
||||
let sender = GenericSender::new(conn_a, policy.clone());
|
||||
let receiver = GenericReceiver::new(conn_b, policy);
|
||||
|
||||
let request = CommunicationValue::new(mtp_codec::CommunicationType::PipeRequest)
|
||||
.with_id(1)
|
||||
.add_typed_default(
|
||||
mtp_codec::DataType::Description,
|
||||
mtp_codec::DataValue::Str("mixed-pipe".into()),
|
||||
);
|
||||
sender.send(&request).await?;
|
||||
let msg = CommunicationValue::new(mtp_codec::CommunicationType::Pong);
|
||||
sender.send(&msg).await?;
|
||||
|
||||
let _pipe_writer = sender.open_pipe(1, "mixed-pipe").await?;
|
||||
|
||||
let received = receiver.receive().await?;
|
||||
assert!(received.is_type(mtp_codec::CommunicationType::PipeRequest));
|
||||
|
||||
receiver.expect_pipe(1)?;
|
||||
let _pipe_writer = sender.open_pipe(1, "mixed-pipe").await?;
|
||||
assert_eq!(
|
||||
received.get_type(),
|
||||
mtp_codec::CommunicationType::Pong
|
||||
.try_to_id(&mtp_codec::TypeMap::latest())
|
||||
.unwrap()
|
||||
);
|
||||
|
||||
let pipe_reader = receiver.receive_pipe().await?;
|
||||
assert_eq!(pipe_reader.pipe_id(), 1);
|
||||
|
|
@ -304,8 +291,6 @@ async fn test_multiple_pipes() -> Result<(), Box<dyn std::error::Error>> {
|
|||
let sender = GenericSender::new(conn_a, policy.clone());
|
||||
let receiver = GenericReceiver::new(conn_b, policy);
|
||||
|
||||
receiver.expect_pipe(10)?;
|
||||
receiver.expect_pipe(20)?;
|
||||
let mut pw1 = sender.open_pipe(10, "first").await?;
|
||||
let mut pw2 = sender.open_pipe(20, "second").await?;
|
||||
|
||||
|
|
|
|||
|
|
@ -127,12 +127,12 @@ async fn test_send_receive_roundtrip() -> Result<(), Box<dyn std::error::Error>>
|
|||
assert_numbered_message(&received, CommunicationType::Ping, 42, &tm);
|
||||
|
||||
// Host sends a response
|
||||
let resp = numbered_message(CommunicationType::BadRequest, 99, &tm);
|
||||
let resp = numbered_message(CommunicationType::Pong, 99, &tm);
|
||||
host_tx.send(&resp).await?;
|
||||
|
||||
// Client receives it
|
||||
let client_received = client_rx.receive().await?;
|
||||
assert_numbered_message(&client_received, CommunicationType::BadRequest, 99, &tm);
|
||||
assert_numbered_message(&client_received, CommunicationType::Pong, 99, &tm);
|
||||
|
||||
// Close both sides
|
||||
client_tx.close().await;
|
||||
|
|
@ -179,13 +179,13 @@ async fn test_concurrent_messages() -> Result<(), Box<dyn std::error::Error>> {
|
|||
|
||||
// Send 3 responses back
|
||||
for i in 0..3u128 {
|
||||
let msg = numbered_message(CommunicationType::BadRequest, i * 10, &tm);
|
||||
let msg = numbered_message(CommunicationType::Pong, i * 10, &tm);
|
||||
client_tx.send(&msg).await?;
|
||||
}
|
||||
|
||||
for i in 0..3u128 {
|
||||
let received = host_rx.receive().await?;
|
||||
assert_numbered_message(&received, CommunicationType::BadRequest, i * 10, &tm);
|
||||
assert_numbered_message(&received, CommunicationType::Pong, i * 10, &tm);
|
||||
}
|
||||
|
||||
client_tx.close().await;
|
||||
|
|
@ -260,11 +260,11 @@ async fn test_drop_receiver_keeps_sender_alive() -> Result<(), Box<dyn std::erro
|
|||
|
||||
let tm = TypeMap::latest();
|
||||
|
||||
let resp = numbered_message(CommunicationType::BadRequest, 7, &tm);
|
||||
let resp = numbered_message(CommunicationType::Pong, 7, &tm);
|
||||
host_tx.send(&resp).await?;
|
||||
|
||||
let got = client_rx.receive().await?;
|
||||
assert_numbered_message(&got, CommunicationType::BadRequest, 7, &tm);
|
||||
assert_numbered_message(&got, CommunicationType::Pong, 7, &tm);
|
||||
|
||||
client_tx.close().await;
|
||||
host_tx.close().await;
|
||||
|
|
@ -284,10 +284,10 @@ async fn test_persistent_stream_reopens_after_local_finish()
|
|||
|
||||
client_tx.finish_stream().await?;
|
||||
|
||||
let msg2 = numbered_message(CommunicationType::BadRequest, 22, &tm);
|
||||
let msg2 = numbered_message(CommunicationType::Pong, 22, &tm);
|
||||
client_tx.send(&msg2).await?;
|
||||
let received2 = host_rx.receive().await?;
|
||||
assert_numbered_message(&received2, CommunicationType::BadRequest, 22, &tm);
|
||||
assert_numbered_message(&received2, CommunicationType::Pong, 22, &tm);
|
||||
|
||||
client_tx.close().await;
|
||||
host_tx.close().await;
|
||||
|
|
|
|||
|
|
@ -16,7 +16,3 @@ pipes = []
|
|||
[build-dependencies]
|
||||
serde = { version = "1", features = ["derive"] }
|
||||
serde_yaml = "0.9"
|
||||
|
||||
[package.metadata.cargo-machete]
|
||||
# cargo-machete does not inspect build.rs, where both build dependencies are used.
|
||||
ignored = ["serde", "serde_yaml"]
|
||||
|
|
|
|||
|
|
@ -27,7 +27,7 @@ getrandom-v04 = { package = "getrandom", version = "0.4.3", features = ["wasm_js
|
|||
mtp-common = { version = "0.3.0", path = "../common" }
|
||||
mtp-type-map = { version = "0.3.0", path = "../type-map" }
|
||||
mtp-codec = { version = "0.3.0", path = "../codec", features = ["crypto", "pipes", "registry"] }
|
||||
mtp-crypto = { version = "0.3.0", path = "../crypto", features = ["wasm", "password-kdf"] }
|
||||
mtp-crypto = { version = "0.3.0", path = "../crypto", features = ["wasm"] }
|
||||
zeroize = "1.9"
|
||||
wasm-bindgen-test = "0.3.76"
|
||||
|
||||
|
|
|
|||
1389
wasm/src/client.rs
Normal file
1389
wasm/src/client.rs
Normal file
File diff suppressed because it is too large
Load diff
|
|
@ -1,680 +0,0 @@
|
|||
use wasm_bindgen::prelude::*;
|
||||
|
||||
use mtp_codec::{CommunicationType, CommunicationValue, DataType, DataValue, PROTOCOL_VERSION};
|
||||
|
||||
use crate::auth;
|
||||
use crate::client::{ConnectionState, WasmClient};
|
||||
use crate::config::ConnectionConfig;
|
||||
use crate::error::js_error;
|
||||
use crate::transport::WasmTransport;
|
||||
|
||||
fn server_rejection_message(outcome: &CommunicationValue) -> Option<&str> {
|
||||
(outcome.get_data(DataType::Connected) == Some(&DataValue::BoolFalse)).then(|| {
|
||||
outcome
|
||||
.get_str(DataType::ErrorMessage)
|
||||
.unwrap_or("host rejected the connection")
|
||||
})
|
||||
}
|
||||
|
||||
#[wasm_bindgen]
|
||||
#[allow(deprecated)]
|
||||
impl WasmClient {
|
||||
pub async fn connect(&self, config: &ConnectionConfig) -> Result<(), JsValue> {
|
||||
self.connect_owned(config.clone()).await
|
||||
}
|
||||
|
||||
#[wasm_bindgen(js_name = connectOwned)]
|
||||
pub async fn connect_owned(&self, config: ConnectionConfig) -> Result<(), JsValue> {
|
||||
let generation = self.begin_connection();
|
||||
let transport = match WasmTransport::connect_with_limits(
|
||||
&config.url,
|
||||
config.server_certificate_hashes.clone(),
|
||||
config.max_message_size,
|
||||
self.receive_decode_limits(),
|
||||
)
|
||||
.await
|
||||
{
|
||||
Ok(transport) => transport,
|
||||
Err(error) => {
|
||||
self.set_state_if_current(generation, ConnectionState::Disconnected);
|
||||
return Err(error);
|
||||
}
|
||||
};
|
||||
if !self.install_attempt_transport(&transport, generation) {
|
||||
return Err(js_error("connection attempt superseded"));
|
||||
}
|
||||
|
||||
let result = async {
|
||||
let version_str = format!("{}", PROTOCOL_VERSION);
|
||||
let opening_codec = mtp_codec::registry::VersionedCodec::for_version(
|
||||
mtp_codec::registry::Registry::builtin(),
|
||||
PROTOCOL_VERSION,
|
||||
)
|
||||
.ok_or_else(|| js_error("client protocol version is not registered"))?;
|
||||
transport.set_type_map(opening_codec.type_map());
|
||||
let mut ident = CommunicationValue::new_with_type_map(
|
||||
CommunicationType::Identification,
|
||||
opening_codec.type_map(),
|
||||
)
|
||||
.add_typed_default(DataType::Version, DataValue::Str(version_str))
|
||||
.add_typed_default(
|
||||
DataType::Id,
|
||||
DataValue::UnsignedNumber(config.client_id as u128),
|
||||
);
|
||||
if let Some(desc) = &config.description {
|
||||
ident =
|
||||
ident.add_typed_default(DataType::Description, DataValue::Str(desc.clone()));
|
||||
}
|
||||
let ident_bytes = ident
|
||||
.to_bytes()
|
||||
.map_err(|e| js_error(format!("encode failed: {}", e)))?;
|
||||
transport.send_frame(&ident_bytes).await?;
|
||||
|
||||
let outcome_bytes = transport.read_one_frame().await?;
|
||||
let outcome = CommunicationValue::try_from_bytes_with_type_map_and_limits(
|
||||
&outcome_bytes,
|
||||
opening_codec.type_map(),
|
||||
transport.decode_limits(),
|
||||
)
|
||||
.map_err(|e| js_error(format!("parse handshake outcome: {e}")))?;
|
||||
if Some(outcome.get_type())
|
||||
== CommunicationType::ErrorBadVersion.try_to_id(opening_codec.type_map())
|
||||
{
|
||||
self.set_state_if_current(generation, ConnectionState::Disconnected);
|
||||
return Err(js_error(
|
||||
outcome
|
||||
.get_str(DataType::ErrorMessage)
|
||||
.unwrap_or("host does not support this protocol version"),
|
||||
));
|
||||
}
|
||||
|
||||
// Generic host rejections are IdentificationResponse frames with
|
||||
// Connected=false. They intentionally do not carry a negotiated
|
||||
// Version because negotiation never completed. Check this before
|
||||
// reading Version, otherwise a useful server error such as an
|
||||
// authentication timeout is reported as the misleading
|
||||
// "host omitted a valid negotiated protocol version".
|
||||
if let Some(message) = server_rejection_message(&outcome) {
|
||||
self.set_state_if_current(generation, ConnectionState::Disconnected);
|
||||
return Err(js_error(message));
|
||||
}
|
||||
|
||||
let missing_version = || {
|
||||
js_error(format!(
|
||||
"host omitted a valid negotiated protocol version (response_type={:?}, connected={:?}, frame_len={})",
|
||||
outcome.get_type(),
|
||||
outcome.get_data(DataType::Connected),
|
||||
outcome_bytes.len(),
|
||||
))
|
||||
};
|
||||
|
||||
let negotiated_version = match outcome.get_data(DataType::Version) {
|
||||
Some(DataValue::Str(version)) => mtp_codec::Version::parse(version)
|
||||
.ok_or_else(|| missing_version())?,
|
||||
_ => return Err(missing_version()),
|
||||
};
|
||||
if negotiated_version != PROTOCOL_VERSION {
|
||||
return Err(js_error(
|
||||
"host selected a protocol version the client did not offer",
|
||||
));
|
||||
}
|
||||
let codec = mtp_codec::registry::VersionedCodec::for_version(
|
||||
mtp_codec::registry::Registry::builtin(),
|
||||
negotiated_version,
|
||||
)
|
||||
.ok_or_else(|| js_error("host returned an unsupported negotiated protocol version"))?;
|
||||
transport.set_type_map(codec.type_map());
|
||||
let outcome = CommunicationValue::try_from_bytes_with_type_map_and_limits(
|
||||
&outcome_bytes,
|
||||
codec.type_map(),
|
||||
transport.decode_limits(),
|
||||
)
|
||||
.map_err(|e| js_error(format!("parse negotiated handshake outcome: {e}")))?;
|
||||
let tm = codec.type_map();
|
||||
let expected = CommunicationType::IdentificationResponse
|
||||
.try_to_id(&tm)
|
||||
.ok_or_else(|| js_error("IdentificationResponse is absent from the type map"))?;
|
||||
if outcome.get_type() != expected
|
||||
|| outcome.get_data(DataType::Connected) != Some(&DataValue::BoolTrue)
|
||||
{
|
||||
self.set_state_if_current(generation, ConnectionState::Disconnected);
|
||||
return Err(js_error(
|
||||
outcome
|
||||
.get_str(DataType::ErrorMessage)
|
||||
.unwrap_or("host rejected the connection"),
|
||||
));
|
||||
}
|
||||
let assigned_id = match outcome.get_data(DataType::Id) {
|
||||
Some(DataValue::UnsignedNumber(id)) => {
|
||||
u64::try_from(*id).map_err(|_| js_error("assigned ID is out of range"))?
|
||||
}
|
||||
_ => {
|
||||
self.set_state_if_current(generation, ConnectionState::Disconnected);
|
||||
return Err(js_error("host omitted the assigned client ID"));
|
||||
}
|
||||
};
|
||||
|
||||
if !self.start_receive_loop(transport.clone(), generation, assigned_id) {
|
||||
return Err(js_error("connection attempt superseded"));
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
.await;
|
||||
if let Err(error) = &result {
|
||||
self.abort_attempt(&transport, generation);
|
||||
let _ = error;
|
||||
}
|
||||
result
|
||||
}
|
||||
|
||||
#[wasm_bindgen]
|
||||
#[deprecated(
|
||||
note = "use the SDK authentication methods; this raw method remains for compatibility"
|
||||
)]
|
||||
pub async fn auth_connect(
|
||||
&self,
|
||||
config: &ConnectionConfig,
|
||||
host_public_key_bytes: &[u8],
|
||||
keyring_bytes: &[u8],
|
||||
client_id: u64,
|
||||
) -> Result<u64, JsValue> {
|
||||
self.auth_connect_owned(
|
||||
config.clone(),
|
||||
host_public_key_bytes.to_vec(),
|
||||
keyring_bytes.to_vec(),
|
||||
client_id,
|
||||
)
|
||||
.await
|
||||
}
|
||||
|
||||
#[wasm_bindgen(js_name = authConnectOwned)]
|
||||
pub async fn auth_connect_owned(
|
||||
&self,
|
||||
config: ConnectionConfig,
|
||||
host_public_key_bytes: Vec<u8>,
|
||||
keyring_bytes: Vec<u8>,
|
||||
client_id: u64,
|
||||
) -> Result<u64, JsValue> {
|
||||
let generation = self.begin_connection();
|
||||
|
||||
let host_pk = match mtp_crypto::PublicKeyBundle::from_bytes(&host_public_key_bytes) {
|
||||
Ok(value) => value,
|
||||
Err(error) => {
|
||||
let error = js_error(format!("invalid host public key: {}", error));
|
||||
self.set_state_if_current(generation, ConnectionState::Disconnected);
|
||||
return Err(error);
|
||||
}
|
||||
};
|
||||
let keyring = match mtp_crypto::Keyring::from_bytes(&keyring_bytes) {
|
||||
Ok(value) => value,
|
||||
Err(error) => {
|
||||
let error = js_error(format!("invalid keyring: {}", error));
|
||||
self.set_state_if_current(generation, ConnectionState::Disconnected);
|
||||
return Err(error);
|
||||
}
|
||||
};
|
||||
|
||||
let handshake_codec = mtp_codec::registry::VersionedCodec::for_version(
|
||||
mtp_codec::registry::Registry::builtin(),
|
||||
PROTOCOL_VERSION,
|
||||
)
|
||||
.ok_or_else(|| js_error("client protocol version is not registered"))?;
|
||||
let tm = handshake_codec.type_map().clone();
|
||||
let version_str = format!("{}", PROTOCOL_VERSION);
|
||||
let public_key_bytes = keyring
|
||||
.public_key_bundle()
|
||||
.try_as_bytes()
|
||||
.map_err(|error| js_error(format!("public key serialization failed: {error}")))?;
|
||||
|
||||
let transport = match WasmTransport::connect_with_limits(
|
||||
&config.url,
|
||||
config.server_certificate_hashes.clone(),
|
||||
config.max_message_size,
|
||||
self.receive_decode_limits(),
|
||||
)
|
||||
.await
|
||||
{
|
||||
Ok(transport) => transport,
|
||||
Err(error) => {
|
||||
self.set_state_if_current(generation, ConnectionState::Disconnected);
|
||||
return Err(error);
|
||||
}
|
||||
};
|
||||
transport.set_type_map(&tm);
|
||||
if !self.install_attempt_transport(&transport, generation) {
|
||||
return Err(js_error("connection attempt superseded"));
|
||||
}
|
||||
|
||||
let result = async {
|
||||
let mut hello =
|
||||
CommunicationValue::new_with_type_map(CommunicationType::Identification, &tm)
|
||||
.add_typed_default(DataType::Version, DataValue::Str(version_str.clone()))
|
||||
.add_typed_default(DataType::Id, DataValue::UnsignedNumber(client_id as u128))
|
||||
// Mark this as an authentication-capable opening so a
|
||||
// non-crypto host can reject it explicitly.
|
||||
.add_typed_default(
|
||||
DataType::PublicKeys,
|
||||
DataValue::Bytes(public_key_bytes.clone()),
|
||||
);
|
||||
if let Some(desc) = &config.description {
|
||||
hello =
|
||||
hello.add_typed_default(DataType::Description, DataValue::Str(desc.clone()));
|
||||
}
|
||||
let hello_bytes = hello
|
||||
.to_bytes()
|
||||
.map_err(|e| js_error(format!("encode failed: {}", e)))?;
|
||||
transport.send_frame(&hello_bytes).await?;
|
||||
|
||||
let server_challenge = self
|
||||
.read_verified_challenge(
|
||||
&transport,
|
||||
&tm,
|
||||
&host_pk,
|
||||
client_id,
|
||||
"auth_connect challenge",
|
||||
config.require_pq,
|
||||
!keyring.sig_pq_secret_key.as_bytes().is_empty(),
|
||||
generation,
|
||||
)
|
||||
.await?;
|
||||
|
||||
let client_nonce = auth::random_nonce()?;
|
||||
|
||||
let proof_payload = mtp_crypto::auth::login_proof_payload(
|
||||
&version_str,
|
||||
client_id,
|
||||
server_challenge,
|
||||
client_nonce,
|
||||
);
|
||||
let proof =
|
||||
auth::signed_challenge_response_bytes(&keyring, &proof_payload, client_nonce, &tm)?;
|
||||
transport.send_frame(&proof).await?;
|
||||
|
||||
let response = transport.read_one_frame().await?;
|
||||
let resp_comm = CommunicationValue::try_from_bytes_with_type_map_and_limits(
|
||||
&response,
|
||||
&tm,
|
||||
transport.decode_limits(),
|
||||
)
|
||||
.map_err(|e| js_error(format!("parse response: {}", e)))?;
|
||||
let negotiated_version = match resp_comm.get_data(DataType::Version) {
|
||||
Some(DataValue::Str(version)) => mtp_codec::Version::parse(version)
|
||||
.ok_or_else(|| js_error("host returned an invalid negotiated version"))?,
|
||||
_ => return Err(js_error("host omitted the negotiated version")),
|
||||
};
|
||||
if negotiated_version != PROTOCOL_VERSION {
|
||||
return Err(js_error(
|
||||
"host selected a protocol version the client did not offer",
|
||||
));
|
||||
}
|
||||
let codec = mtp_codec::registry::VersionedCodec::for_version(
|
||||
mtp_codec::registry::Registry::builtin(),
|
||||
negotiated_version,
|
||||
)
|
||||
.ok_or_else(|| js_error("host returned an unsupported negotiated version"))?;
|
||||
transport.set_type_map(codec.type_map());
|
||||
let resp_comm = CommunicationValue::try_from_bytes_with_type_map_and_limits(
|
||||
&response,
|
||||
codec.type_map(),
|
||||
transport.decode_limits(),
|
||||
)
|
||||
.map_err(|e| js_error(format!("parse negotiated response: {}", e)))?;
|
||||
let tm = codec.type_map();
|
||||
let resp_type = resp_comm.get_type();
|
||||
let expected_type = CommunicationType::IdentificationResponse
|
||||
.try_to_id(&tm)
|
||||
.ok_or_else(|| js_error("IdentificationResponse is absent from the type map"))?;
|
||||
if resp_type != expected_type {
|
||||
self.set_state_if_current(generation, ConnectionState::Disconnected);
|
||||
return Err(auth::unexpected_response_type_error(
|
||||
"auth_connect",
|
||||
expected_type,
|
||||
resp_type,
|
||||
&response,
|
||||
&resp_comm,
|
||||
));
|
||||
}
|
||||
|
||||
if resp_comm.get_data(DataType::Connected) != Some(&DataValue::BoolTrue) {
|
||||
self.set_state_if_current(generation, ConnectionState::Disconnected);
|
||||
return Err(js_error(
|
||||
resp_comm
|
||||
.get_str(DataType::ErrorMessage)
|
||||
.unwrap_or("host rejected authentication"),
|
||||
));
|
||||
}
|
||||
|
||||
if let Err(e) = auth::verify_host_final(
|
||||
&resp_comm,
|
||||
&tm,
|
||||
&host_pk,
|
||||
client_id,
|
||||
client_nonce,
|
||||
server_challenge,
|
||||
config.require_pq,
|
||||
) {
|
||||
self.set_state_if_current(generation, ConnectionState::Disconnected);
|
||||
return Err(e);
|
||||
}
|
||||
|
||||
let assigned_id = match resp_comm.get_data(DataType::Id) {
|
||||
Some(DataValue::UnsignedNumber(n)) => match u64::try_from(*n) {
|
||||
Ok(id) => id,
|
||||
Err(_) => {
|
||||
self.set_state_if_current(generation, ConnectionState::Disconnected);
|
||||
return Err(js_error("assigned ID is out of range"));
|
||||
}
|
||||
},
|
||||
_ => {
|
||||
self.set_state_if_current(generation, ConnectionState::Disconnected);
|
||||
return Err(js_error("missing assigned ID"));
|
||||
}
|
||||
};
|
||||
|
||||
if !self.start_receive_loop(transport.clone(), generation, assigned_id) {
|
||||
return Err(js_error("connection attempt superseded"));
|
||||
}
|
||||
|
||||
Ok(assigned_id)
|
||||
}
|
||||
.await;
|
||||
if let Err(error) = &result {
|
||||
self.abort_attempt(&transport, generation);
|
||||
let _ = error;
|
||||
}
|
||||
result
|
||||
}
|
||||
|
||||
#[wasm_bindgen]
|
||||
#[deprecated(
|
||||
note = "use the SDK registration methods; this raw method remains for compatibility"
|
||||
)]
|
||||
pub async fn auth_register(
|
||||
&self,
|
||||
config: &ConnectionConfig,
|
||||
host_public_key_bytes: &[u8],
|
||||
keyring_bytes: &[u8],
|
||||
) -> Result<u64, JsValue> {
|
||||
self.auth_register_owned(
|
||||
config.clone(),
|
||||
host_public_key_bytes.to_vec(),
|
||||
keyring_bytes.to_vec(),
|
||||
)
|
||||
.await
|
||||
}
|
||||
|
||||
#[wasm_bindgen(js_name = authRegisterOwned)]
|
||||
pub async fn auth_register_owned(
|
||||
&self,
|
||||
config: ConnectionConfig,
|
||||
host_public_key_bytes: Vec<u8>,
|
||||
keyring_bytes: Vec<u8>,
|
||||
) -> Result<u64, JsValue> {
|
||||
let generation = self.begin_connection();
|
||||
|
||||
let host_pk = match mtp_crypto::PublicKeyBundle::from_bytes(&host_public_key_bytes) {
|
||||
Ok(value) => value,
|
||||
Err(error) => {
|
||||
let error = js_error(format!("invalid host public key: {}", error));
|
||||
self.set_state_if_current(generation, ConnectionState::Disconnected);
|
||||
return Err(error);
|
||||
}
|
||||
};
|
||||
let keyring = match mtp_crypto::Keyring::from_bytes(&keyring_bytes) {
|
||||
Ok(value) => value,
|
||||
Err(error) => {
|
||||
let error = js_error(format!("invalid keyring: {}", error));
|
||||
self.set_state_if_current(generation, ConnectionState::Disconnected);
|
||||
return Err(error);
|
||||
}
|
||||
};
|
||||
|
||||
let handshake_codec = mtp_codec::registry::VersionedCodec::for_version(
|
||||
mtp_codec::registry::Registry::builtin(),
|
||||
PROTOCOL_VERSION,
|
||||
)
|
||||
.ok_or_else(|| js_error("client protocol version is not registered"))?;
|
||||
let tm = handshake_codec.type_map().clone();
|
||||
let version_str = format!("{}", PROTOCOL_VERSION);
|
||||
let pk_bytes = keyring
|
||||
.public_key_bundle()
|
||||
.try_as_bytes()
|
||||
.map_err(|error| js_error(format!("public key serialization failed: {error}")))?;
|
||||
|
||||
let transport = match WasmTransport::connect_with_limits(
|
||||
&config.url,
|
||||
config.server_certificate_hashes.clone(),
|
||||
config.max_message_size,
|
||||
self.receive_decode_limits(),
|
||||
)
|
||||
.await
|
||||
{
|
||||
Ok(transport) => transport,
|
||||
Err(error) => {
|
||||
self.set_state_if_current(generation, ConnectionState::Disconnected);
|
||||
return Err(error);
|
||||
}
|
||||
};
|
||||
transport.set_type_map(&tm);
|
||||
if !self.install_attempt_transport(&transport, generation) {
|
||||
return Err(js_error("connection attempt superseded"));
|
||||
}
|
||||
|
||||
let result = async {
|
||||
let mut hello = CommunicationValue::new_with_type_map(CommunicationType::Register, &tm)
|
||||
.add_typed_default(DataType::Version, DataValue::Str(version_str.clone()))
|
||||
.add_typed_default(DataType::PublicKeys, DataValue::Bytes(pk_bytes.clone()));
|
||||
if let Some(desc) = &config.description {
|
||||
hello =
|
||||
hello.add_typed_default(DataType::Description, DataValue::Str(desc.clone()));
|
||||
}
|
||||
let hello_bytes = hello
|
||||
.to_bytes()
|
||||
.map_err(|e| js_error(format!("encode failed: {}", e)))?;
|
||||
transport.send_frame(&hello_bytes).await?;
|
||||
|
||||
let server_challenge = self
|
||||
.read_verified_challenge(
|
||||
&transport,
|
||||
&tm,
|
||||
&host_pk,
|
||||
0,
|
||||
"auth_register challenge",
|
||||
config.require_pq,
|
||||
!keyring.sig_pq_secret_key.as_bytes().is_empty(),
|
||||
generation,
|
||||
)
|
||||
.await?;
|
||||
|
||||
let client_nonce = auth::random_nonce()?;
|
||||
|
||||
let proof_payload = mtp_crypto::auth::register_proof_payload(
|
||||
&version_str,
|
||||
&pk_bytes,
|
||||
server_challenge,
|
||||
client_nonce,
|
||||
);
|
||||
let proof =
|
||||
auth::signed_challenge_response_bytes(&keyring, &proof_payload, client_nonce, &tm)?;
|
||||
transport.send_frame(&proof).await?;
|
||||
|
||||
let response = transport.read_one_frame().await?;
|
||||
let resp_comm = CommunicationValue::try_from_bytes_with_type_map_and_limits(
|
||||
&response,
|
||||
&tm,
|
||||
transport.decode_limits(),
|
||||
)
|
||||
.map_err(|e| js_error(format!("parse response: {}", e)))?;
|
||||
let negotiated_version = match resp_comm.get_data(DataType::Version) {
|
||||
Some(DataValue::Str(version)) => mtp_codec::Version::parse(version)
|
||||
.ok_or_else(|| js_error("host returned an invalid negotiated version"))?,
|
||||
_ => return Err(js_error("host omitted the negotiated version")),
|
||||
};
|
||||
if negotiated_version != PROTOCOL_VERSION {
|
||||
return Err(js_error(
|
||||
"host selected a protocol version the client did not offer",
|
||||
));
|
||||
}
|
||||
let codec = mtp_codec::registry::VersionedCodec::for_version(
|
||||
mtp_codec::registry::Registry::builtin(),
|
||||
negotiated_version,
|
||||
)
|
||||
.ok_or_else(|| js_error("host returned an unsupported negotiated version"))?;
|
||||
transport.set_type_map(codec.type_map());
|
||||
let resp_comm = CommunicationValue::try_from_bytes_with_type_map_and_limits(
|
||||
&response,
|
||||
codec.type_map(),
|
||||
transport.decode_limits(),
|
||||
)
|
||||
.map_err(|e| js_error(format!("parse negotiated response: {}", e)))?;
|
||||
let tm = codec.type_map();
|
||||
let resp_type = resp_comm.get_type();
|
||||
let expected_type = CommunicationType::RegisterResponse
|
||||
.try_to_id(&tm)
|
||||
.ok_or_else(|| js_error("RegisterResponse is absent from the type map"))?;
|
||||
if resp_type != expected_type {
|
||||
self.set_state_if_current(generation, ConnectionState::Disconnected);
|
||||
return Err(auth::unexpected_response_type_error(
|
||||
"auth_register",
|
||||
expected_type,
|
||||
resp_type,
|
||||
&response,
|
||||
&resp_comm,
|
||||
));
|
||||
}
|
||||
|
||||
if resp_comm.get_data(DataType::Connected) != Some(&DataValue::BoolTrue) {
|
||||
self.set_state_if_current(generation, ConnectionState::Disconnected);
|
||||
return Err(js_error(
|
||||
resp_comm
|
||||
.get_str(DataType::ErrorMessage)
|
||||
.unwrap_or("host rejected registration"),
|
||||
));
|
||||
}
|
||||
|
||||
let assigned_id = match resp_comm.get_data(DataType::Id) {
|
||||
Some(DataValue::UnsignedNumber(n)) => match u64::try_from(*n) {
|
||||
Ok(id) => id,
|
||||
Err(_) => {
|
||||
self.set_state_if_current(generation, ConnectionState::Disconnected);
|
||||
return Err(js_error("assigned ID is out of range"));
|
||||
}
|
||||
},
|
||||
_ => {
|
||||
self.set_state_if_current(generation, ConnectionState::Disconnected);
|
||||
return Err(js_error("missing assigned ID"));
|
||||
}
|
||||
};
|
||||
|
||||
if let Err(e) = auth::verify_host_final(
|
||||
&resp_comm,
|
||||
&tm,
|
||||
&host_pk,
|
||||
assigned_id,
|
||||
client_nonce,
|
||||
server_challenge,
|
||||
config.require_pq,
|
||||
) {
|
||||
self.set_state_if_current(generation, ConnectionState::Disconnected);
|
||||
return Err(e);
|
||||
}
|
||||
|
||||
if !self.start_receive_loop(transport.clone(), generation, assigned_id) {
|
||||
return Err(js_error("connection attempt superseded"));
|
||||
}
|
||||
|
||||
Ok(assigned_id)
|
||||
}
|
||||
.await;
|
||||
if let Err(error) = &result {
|
||||
self.abort_attempt(&transport, generation);
|
||||
let _ = error;
|
||||
}
|
||||
result
|
||||
}
|
||||
|
||||
async fn read_verified_challenge(
|
||||
&self,
|
||||
transport: &WasmTransport,
|
||||
tm: &mtp_codec::TypeMap,
|
||||
host_pk: &mtp_crypto::PublicKeyBundle,
|
||||
bound_id: u64,
|
||||
context: &str,
|
||||
require_pq: bool,
|
||||
client_has_pq_key: bool,
|
||||
generation: u32,
|
||||
) -> Result<u128, JsValue> {
|
||||
let challenge_bytes = transport.read_one_frame().await?;
|
||||
let challenge = CommunicationValue::try_from_bytes_with_type_map_and_limits(
|
||||
&challenge_bytes,
|
||||
tm,
|
||||
transport.decode_limits(),
|
||||
)
|
||||
.map_err(|e| js_error(format!("parse challenge: {}", e)))?;
|
||||
let expected = CommunicationType::Challenge
|
||||
.try_to_id(tm)
|
||||
.ok_or_else(|| js_error("Challenge is absent from the type map"))?;
|
||||
if challenge.get_type() != expected {
|
||||
self.set_state_if_current(generation, ConnectionState::Disconnected);
|
||||
return Err(auth::unexpected_response_type_error(
|
||||
context,
|
||||
expected,
|
||||
challenge.get_type(),
|
||||
&challenge_bytes,
|
||||
&challenge,
|
||||
));
|
||||
}
|
||||
|
||||
let server_challenge = match challenge.get_data(DataType::ServerNonce) {
|
||||
Some(DataValue::UnsignedNumber(n)) => *n,
|
||||
_ => {
|
||||
self.set_state_if_current(generation, ConnectionState::Disconnected);
|
||||
return Err(js_error("missing server challenge"));
|
||||
}
|
||||
};
|
||||
|
||||
if challenge.get_data(DataType::RequirePq) == Some(&DataValue::BoolTrue)
|
||||
&& !client_has_pq_key
|
||||
{
|
||||
self.set_state_if_current(generation, ConnectionState::Disconnected);
|
||||
return Err(js_error(
|
||||
"host requires post-quantum authentication but the client PQ key is absent",
|
||||
));
|
||||
}
|
||||
|
||||
if let Err(e) = auth::verify_host_challenge(
|
||||
&challenge,
|
||||
tm,
|
||||
host_pk,
|
||||
bound_id,
|
||||
server_challenge,
|
||||
require_pq,
|
||||
) {
|
||||
self.set_state_if_current(generation, ConnectionState::Disconnected);
|
||||
return Err(e);
|
||||
}
|
||||
|
||||
Ok(server_challenge)
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::server_rejection_message;
|
||||
use mtp_codec::{CommunicationType, CommunicationValue, DataType, DataValue};
|
||||
|
||||
#[test]
|
||||
fn reports_rejection_reason_without_a_negotiated_version() {
|
||||
let response = CommunicationValue::new(CommunicationType::IdentificationResponse)
|
||||
.add_typed_default(DataType::Connected, DataValue::BoolFalse)
|
||||
.add_typed_default(
|
||||
DataType::ErrorMessage,
|
||||
DataValue::Str("authentication handshake timed out".into()),
|
||||
);
|
||||
|
||||
assert_eq!(
|
||||
server_rejection_message(&response),
|
||||
Some("authentication handshake timed out")
|
||||
);
|
||||
}
|
||||
}
|
||||
|
|
@ -1,102 +0,0 @@
|
|||
use wasm_bindgen::prelude::*;
|
||||
|
||||
use crate::client::{ConnectionState, WasmClient};
|
||||
use crate::client_pipe;
|
||||
use crate::transport::WasmTransport;
|
||||
|
||||
use super::dispatch::set_shared_state;
|
||||
|
||||
#[wasm_bindgen]
|
||||
impl WasmClient {
|
||||
pub fn disconnect(&self) {
|
||||
self.connection_generation
|
||||
.set(self.connection_generation.get().wrapping_add(1));
|
||||
self.stop_protocol_pings();
|
||||
if let Some(t) = self.transport.borrow_mut().take() {
|
||||
t.close();
|
||||
}
|
||||
if let Some(t) = self.attempt_transport.borrow_mut().take() {
|
||||
t.close();
|
||||
}
|
||||
self.subscriptions.borrow_mut().clear();
|
||||
self.reject_pending_requests("disconnected");
|
||||
client_pipe::reject_pending_pipe_creations(&self.pending_pipe_creations, "disconnected");
|
||||
self.expired_pipe_creations.borrow_mut().clear();
|
||||
client_pipe::reject_pending_pipes(&self.pending_pipes, "disconnected");
|
||||
self.connection_client_id.set(0);
|
||||
self.set_state(ConnectionState::Disconnected);
|
||||
}
|
||||
|
||||
pub(super) fn set_state(&self, new_state: ConnectionState) {
|
||||
set_shared_state(
|
||||
&self.state,
|
||||
&self.pending_state_callbacks,
|
||||
self.state_callback.as_ref(),
|
||||
new_state,
|
||||
);
|
||||
}
|
||||
|
||||
pub(super) fn set_state_if_current(&self, generation: u32, new_state: ConnectionState) {
|
||||
if self.connection_generation.get() == generation {
|
||||
self.set_state(new_state);
|
||||
}
|
||||
}
|
||||
|
||||
pub(super) fn install_attempt_transport(
|
||||
&self,
|
||||
transport: &WasmTransport,
|
||||
generation: u32,
|
||||
) -> bool {
|
||||
if self.connection_generation.get() != generation {
|
||||
transport.close();
|
||||
return false;
|
||||
}
|
||||
*self.attempt_transport.borrow_mut() = Some(transport.clone());
|
||||
true
|
||||
}
|
||||
|
||||
pub(super) fn abort_attempt(&self, transport: &WasmTransport, generation: u32) {
|
||||
transport.close();
|
||||
if self.connection_generation.get() != generation {
|
||||
return;
|
||||
}
|
||||
if let Some(current) = self.attempt_transport.borrow_mut().take() {
|
||||
current.close();
|
||||
}
|
||||
if let Some(current) = self.transport.borrow_mut().take() {
|
||||
current.close();
|
||||
}
|
||||
self.stop_protocol_pings();
|
||||
self.reject_pending_requests("connection failed");
|
||||
client_pipe::reject_pending_pipe_creations(
|
||||
&self.pending_pipe_creations,
|
||||
"connection failed",
|
||||
);
|
||||
self.expired_pipe_creations.borrow_mut().clear();
|
||||
client_pipe::reject_pending_pipes(&self.pending_pipes, "connection failed");
|
||||
self.connection_client_id.set(0);
|
||||
self.set_state(ConnectionState::Disconnected);
|
||||
}
|
||||
|
||||
pub(super) fn begin_connection(&self) -> u32 {
|
||||
let generation = self.connection_generation.get().wrapping_add(1);
|
||||
self.connection_generation.set(generation);
|
||||
self.stop_protocol_pings();
|
||||
if let Some(transport) = self.transport.borrow_mut().take() {
|
||||
transport.close();
|
||||
}
|
||||
if let Some(transport) = self.attempt_transport.borrow_mut().take() {
|
||||
transport.close();
|
||||
}
|
||||
self.reject_pending_requests("connection replaced");
|
||||
client_pipe::reject_pending_pipe_creations(
|
||||
&self.pending_pipe_creations,
|
||||
"connection replaced",
|
||||
);
|
||||
self.expired_pipe_creations.borrow_mut().clear();
|
||||
client_pipe::reject_pending_pipes(&self.pending_pipes, "connection replaced");
|
||||
self.connection_client_id.set(0);
|
||||
self.set_state(ConnectionState::Connecting);
|
||||
generation
|
||||
}
|
||||
}
|
||||
|
|
@ -1,184 +0,0 @@
|
|||
use std::cell::{Cell, RefCell};
|
||||
use std::collections::{HashMap, VecDeque};
|
||||
use std::rc::Rc;
|
||||
|
||||
use wasm_bindgen::prelude::*;
|
||||
|
||||
use crate::client::ConnectionState;
|
||||
use crate::client_pipe::{self, PendingRequest};
|
||||
|
||||
pub(super) struct PingTimer {
|
||||
pub(super) id: i32,
|
||||
pub(super) closure: Closure<dyn FnMut()>,
|
||||
}
|
||||
|
||||
pub(super) struct PendingPing {
|
||||
pub(super) generation: u32,
|
||||
pub(super) sent_at: f64,
|
||||
}
|
||||
|
||||
pub(super) fn frame_property(frame: &JsValue, key: &str) -> Option<JsValue> {
|
||||
js_sys::Reflect::get(frame, &JsValue::from_str(key))
|
||||
.ok()
|
||||
.filter(|value| !value.is_null() && !value.is_undefined())
|
||||
}
|
||||
|
||||
pub(super) fn frame_id(frame: &JsValue) -> Option<u32> {
|
||||
frame_property(frame, "id")
|
||||
.and_then(|value| value.as_f64())
|
||||
.filter(|value| {
|
||||
value.is_finite() && value.fract() == 0.0 && (0.0..=u32::MAX as f64).contains(value)
|
||||
})
|
||||
.and_then(|value| u32::try_from(value as u64).ok())
|
||||
}
|
||||
|
||||
pub(super) fn frame_type(frame: &JsValue) -> Option<String> {
|
||||
frame_property(frame, "type").and_then(|value| value.as_string())
|
||||
}
|
||||
|
||||
pub(super) fn route_incoming_frame(
|
||||
frame: &JsValue,
|
||||
generation: u32,
|
||||
on_message: &js_sys::Function,
|
||||
subscriptions: &Rc<RefCell<HashMap<u32, (String, js_sys::Function)>>>,
|
||||
pending_requests: &Rc<RefCell<HashMap<u32, PendingRequest>>>,
|
||||
expired_requests: &Rc<RefCell<HashMap<u32, f64>>>,
|
||||
pending_pings: &Rc<RefCell<HashMap<u32, PendingPing>>>,
|
||||
ping_ms: &Rc<Cell<Option<f64>>>,
|
||||
) {
|
||||
let message_type = frame_type(frame);
|
||||
|
||||
if message_type.as_deref() == Some("Pong")
|
||||
&& let Some(ping_id) = frame_id(frame)
|
||||
{
|
||||
let sent_at = pending_pings
|
||||
.borrow()
|
||||
.get(&ping_id)
|
||||
.filter(|ping| ping.generation == generation)
|
||||
.map(|ping| ping.sent_at);
|
||||
if let Some(sent_at) = sent_at {
|
||||
pending_pings.borrow_mut().remove(&ping_id);
|
||||
ping_ms.set(Some(js_sys::Date::now() - sent_at));
|
||||
return;
|
||||
}
|
||||
}
|
||||
|
||||
if let Some(request_id) = frame_id(frame) {
|
||||
let pending = {
|
||||
let mut requests = pending_requests.borrow_mut();
|
||||
if requests
|
||||
.get(&request_id)
|
||||
.is_some_and(|request| request.generation == generation)
|
||||
{
|
||||
requests.remove(&request_id)
|
||||
} else {
|
||||
None
|
||||
}
|
||||
};
|
||||
if let Some(pending) = pending {
|
||||
let type_matches = pending
|
||||
.response_type
|
||||
.as_ref()
|
||||
.zip(message_type.as_ref())
|
||||
.map(|(expected, actual)| expected == actual)
|
||||
.unwrap_or(true);
|
||||
if type_matches {
|
||||
let _ = pending.sender.send(Ok(frame.clone()));
|
||||
} else {
|
||||
let actual = message_type.clone().unwrap_or_else(|| "unknown".into());
|
||||
let _ = pending.sender.send(Err(crate::error::js_error(format!(
|
||||
"unexpected response type: expected {}, got {}",
|
||||
pending.response_type.unwrap_or_else(|| "unknown".into()),
|
||||
actual
|
||||
))));
|
||||
}
|
||||
return;
|
||||
}
|
||||
if client_pipe::consume_expired_request(expired_requests, request_id) {
|
||||
return;
|
||||
}
|
||||
}
|
||||
|
||||
let _ = on_message.call1(&JsValue::NULL, frame);
|
||||
let Some(message_type) = message_type else {
|
||||
return;
|
||||
};
|
||||
let callbacks: Vec<js_sys::Function> = subscriptions
|
||||
.borrow()
|
||||
.iter()
|
||||
.filter(|(_, (t, _))| t == &message_type)
|
||||
.map(|(_, (_, cb))| cb.clone())
|
||||
.collect();
|
||||
for callback in callbacks {
|
||||
let _ = callback.call1(&JsValue::NULL, frame);
|
||||
}
|
||||
}
|
||||
|
||||
pub(super) fn stop_ping_timer(ping_timer: &Rc<RefCell<Option<PingTimer>>>) {
|
||||
let Some(timer) = ping_timer.borrow_mut().take() else {
|
||||
return;
|
||||
};
|
||||
if let Ok(clear_interval) =
|
||||
js_sys::Reflect::get(&js_sys::global(), &JsValue::from_str("clearInterval"))
|
||||
.and_then(|value| value.dyn_into::<js_sys::Function>())
|
||||
{
|
||||
let _ = clear_interval.call1(&JsValue::NULL, &JsValue::from_f64(timer.id as f64));
|
||||
}
|
||||
drop(timer.closure);
|
||||
}
|
||||
|
||||
pub(super) fn reject_pending_requests(
|
||||
pending_requests: &Rc<RefCell<HashMap<u32, PendingRequest>>>,
|
||||
message: &str,
|
||||
) {
|
||||
let pending = std::mem::take(&mut *pending_requests.borrow_mut());
|
||||
for (_, pending) in pending {
|
||||
let _ = pending.sender.send(Err(crate::error::js_error(message)));
|
||||
}
|
||||
}
|
||||
|
||||
pub(super) async fn wait_for_timeout(timeout_ms: u32) -> Result<(), JsValue> {
|
||||
let promise = js_sys::Promise::new(&mut |resolve, reject| {
|
||||
let result = js_sys::Reflect::get(&js_sys::global(), &JsValue::from_str("setTimeout"))
|
||||
.and_then(|value| value.dyn_into::<js_sys::Function>())
|
||||
.and_then(|set_timeout| {
|
||||
set_timeout.call2(
|
||||
&JsValue::NULL,
|
||||
&resolve,
|
||||
&JsValue::from_f64(timeout_ms as f64),
|
||||
)
|
||||
});
|
||||
if let Err(error) = result {
|
||||
let _ = reject.call1(&JsValue::NULL, &error);
|
||||
}
|
||||
});
|
||||
wasm_bindgen_futures::JsFuture::from(promise).await?;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
pub(super) fn set_shared_state(
|
||||
state: &Rc<Cell<ConnectionState>>,
|
||||
pending_state_callbacks: &Rc<RefCell<VecDeque<ConnectionState>>>,
|
||||
state_callback: &JsValue,
|
||||
new_state: ConnectionState,
|
||||
) {
|
||||
state.set(new_state);
|
||||
pending_state_callbacks.borrow_mut().push_back(new_state);
|
||||
|
||||
let global = js_sys::global();
|
||||
let qmt = js_sys::Reflect::get(&global, &JsValue::from_str("queueMicrotask"))
|
||||
.and_then(|f| f.dyn_into::<js_sys::Function>());
|
||||
let scheduled = qmt
|
||||
.and_then(|qmt| qmt.call1(&global, state_callback))
|
||||
.is_ok();
|
||||
if !scheduled
|
||||
&& js_sys::Reflect::get(&global, &JsValue::from_str("setTimeout"))
|
||||
.and_then(|f| f.dyn_into::<js_sys::Function>())
|
||||
.and_then(|set_timeout| {
|
||||
set_timeout.call2(&global, state_callback, &JsValue::from_f64(0.0))
|
||||
})
|
||||
.is_err()
|
||||
{
|
||||
pending_state_callbacks.borrow_mut().pop_back();
|
||||
}
|
||||
}
|
||||
|
|
@ -1,299 +0,0 @@
|
|||
// WASM client facade. Lifecycle, authentication, receive dispatch, and pipes
|
||||
// live in private child modules below.
|
||||
use std::cell::{Cell, RefCell};
|
||||
use std::collections::HashMap;
|
||||
use std::collections::VecDeque;
|
||||
use std::rc::Rc;
|
||||
|
||||
use futures_channel::oneshot;
|
||||
use futures_util::{FutureExt, pin_mut, select};
|
||||
use wasm_bindgen::prelude::*;
|
||||
|
||||
use mtp_codec::{CommunicationValue, DecodeLimits, EncodeLimits};
|
||||
|
||||
use crate::client_pipe::{self, PendingRequest};
|
||||
use crate::error::js_error;
|
||||
use crate::transport::WasmTransport;
|
||||
|
||||
mod authentication;
|
||||
mod connection;
|
||||
mod dispatch;
|
||||
mod pipes;
|
||||
mod receive;
|
||||
use dispatch::{PendingPing, PingTimer, wait_for_timeout};
|
||||
|
||||
const DEFAULT_REQUEST_TIMEOUT_MS: u32 = 30_000;
|
||||
const MAX_SAFE_JS_INTEGER: f64 = 9_007_199_254_740_991.0;
|
||||
|
||||
fn decode_limit(value: &JsValue, key: &str, default: usize) -> Result<usize, JsValue> {
|
||||
if value.is_null() || value.is_undefined() {
|
||||
return Ok(default);
|
||||
}
|
||||
let value = js_sys::Reflect::get(value, &JsValue::from_str(key))?;
|
||||
if value.is_null() || value.is_undefined() {
|
||||
return Ok(default);
|
||||
}
|
||||
let Some(number) = value.as_f64() else {
|
||||
return Err(js_error(format!("{key} must be a number")));
|
||||
};
|
||||
if !number.is_finite() || number.fract() != 0.0 || number < 0.0 || number > MAX_SAFE_JS_INTEGER
|
||||
{
|
||||
return Err(js_error(format!("{key} must be a non-negative integer")));
|
||||
}
|
||||
usize::try_from(number as u64).map_err(|_| js_error(format!("{key} is out of range")))
|
||||
}
|
||||
|
||||
pub(crate) fn encode_limits_from_js(value: &JsValue) -> Result<EncodeLimits, JsValue> {
|
||||
let defaults = EncodeLimits::default();
|
||||
Ok(EncodeLimits {
|
||||
max_depth: decode_limit(value, "maxDepth", defaults.max_depth)?,
|
||||
max_values: decode_limit(value, "maxValues", defaults.max_values)?,
|
||||
max_output_size: decode_limit(value, "maxOutputSize", defaults.max_output_size)?,
|
||||
})
|
||||
}
|
||||
|
||||
pub(crate) fn decode_limits_from_js(value: &JsValue) -> Result<DecodeLimits, JsValue> {
|
||||
let defaults = DecodeLimits::default();
|
||||
Ok(DecodeLimits {
|
||||
max_depth: decode_limit(value, "maxDepth", defaults.max_depth)?,
|
||||
max_values: decode_limit(value, "maxValues", defaults.max_values)?,
|
||||
max_blob_size: decode_limit(value, "maxBlobSize", defaults.max_blob_size)?,
|
||||
max_recipients: decode_limit(value, "maxRecipients", defaults.max_recipients)?,
|
||||
max_allocated_bytes: decode_limit(
|
||||
value,
|
||||
"maxAllocatedBytes",
|
||||
defaults.max_allocated_bytes,
|
||||
)?,
|
||||
})
|
||||
}
|
||||
|
||||
#[wasm_bindgen]
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
||||
pub enum ConnectionState {
|
||||
Disconnected = 0,
|
||||
Connecting = 1,
|
||||
Connected = 2,
|
||||
Failed = 3,
|
||||
}
|
||||
|
||||
#[wasm_bindgen]
|
||||
pub struct WasmClient {
|
||||
transport: Rc<RefCell<Option<WasmTransport>>>,
|
||||
attempt_transport: Rc<RefCell<Option<WasmTransport>>>,
|
||||
connection_generation: Rc<Cell<u32>>,
|
||||
state: Rc<Cell<ConnectionState>>,
|
||||
pending_state_callbacks: Rc<RefCell<VecDeque<ConnectionState>>>,
|
||||
state_callback: Closure<dyn FnMut()>,
|
||||
pub(crate) on_message: js_sys::Function,
|
||||
pub(crate) on_error: js_sys::Function,
|
||||
subscriptions: Rc<RefCell<HashMap<u32, (String, js_sys::Function)>>>,
|
||||
next_subscription_id: Rc<Cell<u32>>,
|
||||
pending_requests: Rc<RefCell<HashMap<u32, PendingRequest>>>,
|
||||
expired_requests: Rc<RefCell<HashMap<u32, f64>>>,
|
||||
ping_timer: Rc<RefCell<Option<PingTimer>>>,
|
||||
pending_pings: Rc<RefCell<HashMap<u32, PendingPing>>>,
|
||||
ping_ms: Rc<Cell<Option<f64>>>,
|
||||
pending_pipe_creations: client_pipe::PendingPipeCreations,
|
||||
expired_pipe_creations: Rc<RefCell<HashMap<u32, f64>>>,
|
||||
pending_pipes: client_pipe::PendingPipes,
|
||||
connection_client_id: Rc<Cell<u64>>,
|
||||
on_pipe_request: Rc<RefCell<Option<js_sys::Function>>>,
|
||||
receive_decode_limits: Rc<RefCell<Option<DecodeLimits>>>,
|
||||
}
|
||||
|
||||
#[wasm_bindgen]
|
||||
#[allow(deprecated)]
|
||||
impl WasmClient {
|
||||
#[wasm_bindgen(constructor)]
|
||||
pub fn new(
|
||||
on_state_change: Option<js_sys::Function>,
|
||||
on_message: Option<js_sys::Function>,
|
||||
on_error: Option<js_sys::Function>,
|
||||
) -> Self {
|
||||
let noop = || js_sys::Function::new_no_args("");
|
||||
let on_state_change = on_state_change.unwrap_or_else(noop);
|
||||
let pending_state_callbacks = Rc::new(RefCell::new(VecDeque::new()));
|
||||
let callback_queue = pending_state_callbacks.clone();
|
||||
let callback = on_state_change.clone();
|
||||
let state_callback = Closure::wrap(Box::new(move || {
|
||||
let state = callback_queue.borrow_mut().pop_front();
|
||||
if let Some(state) = state {
|
||||
let _ = callback.call1(&JsValue::NULL, &JsValue::from(state as u8));
|
||||
}
|
||||
}) as Box<dyn FnMut()>);
|
||||
Self {
|
||||
transport: Rc::new(RefCell::new(None)),
|
||||
attempt_transport: Rc::new(RefCell::new(None)),
|
||||
connection_generation: Rc::new(Cell::new(0)),
|
||||
state: Rc::new(Cell::new(ConnectionState::Disconnected)),
|
||||
pending_state_callbacks,
|
||||
state_callback,
|
||||
on_message: on_message.unwrap_or_else(noop),
|
||||
on_error: on_error.unwrap_or_else(noop),
|
||||
subscriptions: Rc::new(RefCell::new(HashMap::new())),
|
||||
next_subscription_id: Rc::new(Cell::new(1)),
|
||||
pending_requests: Rc::new(RefCell::new(HashMap::new())),
|
||||
expired_requests: Rc::new(RefCell::new(HashMap::new())),
|
||||
ping_timer: Rc::new(RefCell::new(None)),
|
||||
pending_pings: Rc::new(RefCell::new(HashMap::new())),
|
||||
ping_ms: Rc::new(Cell::new(None)),
|
||||
pending_pipe_creations: Rc::new(RefCell::new(HashMap::new())),
|
||||
expired_pipe_creations: Rc::new(RefCell::new(HashMap::new())),
|
||||
pending_pipes: Rc::new(RefCell::new(HashMap::new())),
|
||||
connection_client_id: Rc::new(Cell::new(0)),
|
||||
on_pipe_request: Rc::new(RefCell::new(None)),
|
||||
receive_decode_limits: Rc::new(RefCell::new(None)),
|
||||
}
|
||||
}
|
||||
pub fn is_supported() -> bool {
|
||||
js_sys::Reflect::has(&js_sys::global(), &JsValue::from_str("WebTransport")).unwrap_or(false)
|
||||
}
|
||||
|
||||
#[wasm_bindgen(getter)]
|
||||
pub fn state(&self) -> u8 {
|
||||
self.state.get() as u8
|
||||
}
|
||||
|
||||
#[wasm_bindgen(getter)]
|
||||
pub fn ping_ms(&self) -> Option<f64> {
|
||||
self.ping_ms.get()
|
||||
}
|
||||
|
||||
#[wasm_bindgen(getter)]
|
||||
pub fn client_id(&self) -> u64 {
|
||||
self.connection_client_id.get()
|
||||
}
|
||||
|
||||
/// Apply one decoder policy to frames received by this raw WASM client.
|
||||
/// The high-level SDK calls this before authentication so handshake,
|
||||
/// transport, and protected opening share the same policy input.
|
||||
#[wasm_bindgen]
|
||||
pub fn set_receive_limits(&self, limits: JsValue) -> Result<(), JsValue> {
|
||||
let parsed = if limits.is_null() || limits.is_undefined() {
|
||||
None
|
||||
} else {
|
||||
Some(decode_limits_from_js(&limits)?)
|
||||
};
|
||||
*self.receive_decode_limits.borrow_mut() = parsed;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
pub(super) fn receive_decode_limits(&self) -> Option<DecodeLimits> {
|
||||
*self.receive_decode_limits.borrow()
|
||||
}
|
||||
|
||||
#[wasm_bindgen]
|
||||
#[deprecated(
|
||||
note = "use the SDK connection methods; this raw method remains for compatibility"
|
||||
)]
|
||||
pub async fn send(&self, frame: Vec<u8>) -> Result<(), JsValue> {
|
||||
if self.state.get() != ConnectionState::Connected {
|
||||
return Err(js_error("not connected"));
|
||||
}
|
||||
let transport = self.transport.borrow().clone();
|
||||
match transport {
|
||||
Some(t) => t.send_frame(&frame).await,
|
||||
None => Err(js_error("not connected")),
|
||||
}
|
||||
}
|
||||
|
||||
#[wasm_bindgen]
|
||||
pub async fn request(
|
||||
&self,
|
||||
frame: Vec<u8>,
|
||||
response_type: Option<String>,
|
||||
timeout_ms: Option<u32>,
|
||||
) -> Result<JsValue, JsValue> {
|
||||
if self.state.get() != ConnectionState::Connected {
|
||||
return Err(js_error("not connected"));
|
||||
}
|
||||
let generation = self.connection_generation.get();
|
||||
let Some(transport) = self.transport.borrow().clone() else {
|
||||
return Err(js_error("not connected"));
|
||||
};
|
||||
let request = CommunicationValue::try_from_bytes_with_type_map_and_limits(
|
||||
&frame,
|
||||
&transport.type_map(),
|
||||
transport.decode_limits(),
|
||||
)
|
||||
.map_err(|e| js_error(format!("parse request: {}", e)))?;
|
||||
let request_id = request
|
||||
.id()
|
||||
.ok_or_else(|| js_error("request frame must contain an id"))?;
|
||||
if request_id == 0 {
|
||||
return Err(js_error("request frame must have a non-zero id"));
|
||||
}
|
||||
if client_pipe::is_expired_request(&self.expired_requests, request_id) {
|
||||
return Err(js_error(format!(
|
||||
"request id {request_id} recently timed out; use a new request id"
|
||||
)));
|
||||
}
|
||||
|
||||
let (sender, receiver) = oneshot::channel();
|
||||
let token = Rc::new(());
|
||||
{
|
||||
let mut pending = self.pending_requests.borrow_mut();
|
||||
if pending.contains_key(&request_id) {
|
||||
return Err(js_error(format!(
|
||||
"request id {request_id} is already pending"
|
||||
)));
|
||||
}
|
||||
pending.insert(
|
||||
request_id,
|
||||
PendingRequest {
|
||||
generation,
|
||||
token: token.clone(),
|
||||
response_type,
|
||||
sender,
|
||||
},
|
||||
);
|
||||
}
|
||||
|
||||
let timeout_ms = timeout_ms.unwrap_or(DEFAULT_REQUEST_TIMEOUT_MS);
|
||||
let response = async {
|
||||
transport.send_frame(&frame).await?;
|
||||
match receiver.await {
|
||||
Ok(result) => result,
|
||||
Err(_) => Err(js_error("request cancelled")),
|
||||
}
|
||||
}
|
||||
.fuse();
|
||||
let timeout = wait_for_timeout(timeout_ms).fuse();
|
||||
pin_mut!(response, timeout);
|
||||
select! {
|
||||
result = response => {
|
||||
if result.is_err() {
|
||||
client_pipe::remove_pending_request(&self.pending_requests, request_id, &token);
|
||||
}
|
||||
result
|
||||
},
|
||||
result = timeout => {
|
||||
client_pipe::expire_pending_request(
|
||||
&self.pending_requests,
|
||||
&self.expired_requests,
|
||||
request_id,
|
||||
&token,
|
||||
);
|
||||
result?;
|
||||
Err(js_error(format!(
|
||||
"request {request_id} timed out after {timeout_ms}ms"
|
||||
)))
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
#[wasm_bindgen]
|
||||
pub fn subscribe(&self, message_type: String, callback: js_sys::Function) -> u32 {
|
||||
let id = self.next_subscription_id.get();
|
||||
self.next_subscription_id.set(id.wrapping_add(1).max(1));
|
||||
self.subscriptions
|
||||
.borrow_mut()
|
||||
.insert(id, (message_type, callback));
|
||||
id
|
||||
}
|
||||
|
||||
#[wasm_bindgen]
|
||||
pub fn unsubscribe(&self, id: u32) -> bool {
|
||||
self.subscriptions.borrow_mut().remove(&id).is_some()
|
||||
}
|
||||
}
|
||||
|
|
@ -1,76 +0,0 @@
|
|||
use wasm_bindgen::prelude::*;
|
||||
|
||||
use crate::client::{ConnectionState, WasmClient};
|
||||
use crate::client_pipe;
|
||||
use crate::error::js_error;
|
||||
use crate::pipe::PipeReader;
|
||||
|
||||
#[wasm_bindgen]
|
||||
impl WasmClient {
|
||||
pub fn set_on_pipe_request(&self, callback: Option<js_sys::Function>) {
|
||||
*self.on_pipe_request.borrow_mut() = callback;
|
||||
}
|
||||
|
||||
#[wasm_bindgen]
|
||||
pub async fn create_pipe(
|
||||
&self,
|
||||
description: &str,
|
||||
) -> Result<client_pipe::WasmPipeHandle, JsValue> {
|
||||
if self.state.get() != ConnectionState::Connected {
|
||||
return Err(js_error("not connected"));
|
||||
}
|
||||
let transport = self
|
||||
.transport
|
||||
.borrow()
|
||||
.clone()
|
||||
.ok_or_else(|| js_error("not connected"))?;
|
||||
|
||||
let pipe_id = client_pipe::random_pipe_id()?;
|
||||
client_pipe::wasm_create_pipe(
|
||||
&transport,
|
||||
description,
|
||||
pipe_id,
|
||||
&self.pending_pipe_creations,
|
||||
&self.expired_pipe_creations,
|
||||
self.connection_generation.get(),
|
||||
&self.connection_generation,
|
||||
)
|
||||
.await
|
||||
}
|
||||
|
||||
#[wasm_bindgen]
|
||||
pub async fn accept_pipe(&self, pipe_id: u32) -> Result<PipeReader, JsValue> {
|
||||
if self.state.get() != ConnectionState::Connected {
|
||||
return Err(js_error("not connected"));
|
||||
}
|
||||
let transport = self
|
||||
.transport
|
||||
.borrow()
|
||||
.clone()
|
||||
.ok_or_else(|| js_error("not connected"))?;
|
||||
|
||||
let generation = self.connection_generation.get();
|
||||
client_pipe::wasm_accept_pipe(
|
||||
&transport,
|
||||
pipe_id,
|
||||
&self.pending_pipes,
|
||||
generation,
|
||||
&self.connection_generation,
|
||||
)
|
||||
.await
|
||||
}
|
||||
|
||||
#[wasm_bindgen]
|
||||
pub async fn deny_pipe(&self, pipe_id: u32) -> Result<(), JsValue> {
|
||||
if self.state.get() != ConnectionState::Connected {
|
||||
return Err(js_error("not connected"));
|
||||
}
|
||||
let transport = self
|
||||
.transport
|
||||
.borrow()
|
||||
.clone()
|
||||
.ok_or_else(|| js_error("not connected"))?;
|
||||
|
||||
client_pipe::wasm_deny_pipe(&transport, pipe_id).await
|
||||
}
|
||||
}
|
||||
|
|
@ -1,335 +0,0 @@
|
|||
use wasm_bindgen::prelude::*;
|
||||
|
||||
use mtp_codec::{CommunicationType, CommunicationValue, DataType, DataValue};
|
||||
|
||||
use crate::client::{ConnectionState, WasmClient};
|
||||
use crate::client_pipe;
|
||||
use crate::error::js_error;
|
||||
use crate::pipe::PipeReader;
|
||||
use crate::transport::WasmTransport;
|
||||
|
||||
use super::MAX_SAFE_JS_INTEGER;
|
||||
use super::dispatch::{
|
||||
PendingPing, PingTimer, frame_id, frame_property, frame_type, reject_pending_requests,
|
||||
route_incoming_frame, set_shared_state, stop_ping_timer,
|
||||
};
|
||||
|
||||
#[wasm_bindgen]
|
||||
impl WasmClient {
|
||||
pub fn start_protocol_pings(&self, interval_ms: u32, client_id: u64) -> Result<(), JsValue> {
|
||||
self.stop_protocol_pings();
|
||||
if self.state.get() != ConnectionState::Connected {
|
||||
return Err(js_error("not connected"));
|
||||
}
|
||||
let Some(transport) = self.transport.borrow().clone() else {
|
||||
return Err(js_error("not connected"));
|
||||
};
|
||||
let generation = self.connection_generation.get();
|
||||
let current_generation = self.connection_generation.clone();
|
||||
let interval_ms = i32::try_from(interval_ms.max(1_000))
|
||||
.map_err(|_| js_error("ping interval is too large"))?;
|
||||
let on_error = self.on_error.clone();
|
||||
let pending_pings = self.pending_pings.clone();
|
||||
let closure = Closure::wrap(Box::new(move || {
|
||||
if current_generation.get() != generation {
|
||||
return;
|
||||
}
|
||||
let transport = transport.clone();
|
||||
let on_error = on_error.clone();
|
||||
let pending_pings = pending_pings.clone();
|
||||
let current_generation = current_generation.clone();
|
||||
wasm_bindgen_futures::spawn_local(async move {
|
||||
if current_generation.get() != generation {
|
||||
return;
|
||||
}
|
||||
let sent_at = js_sys::Date::now();
|
||||
pending_pings.borrow_mut().retain(|_, pending| {
|
||||
pending.generation == generation
|
||||
&& sent_at - pending.sent_at < interval_ms as f64 * 3.0
|
||||
});
|
||||
let timestamp = if sent_at.is_finite()
|
||||
&& sent_at >= 0.0
|
||||
&& sent_at <= MAX_SAFE_JS_INTEGER
|
||||
&& sent_at.fract() == 0.0
|
||||
{
|
||||
sent_at as u64
|
||||
} else {
|
||||
let _ = on_error.call1(&JsValue::NULL, &js_error("invalid clock value"));
|
||||
return;
|
||||
};
|
||||
let type_map = transport.type_map();
|
||||
let frame =
|
||||
CommunicationValue::new_with_type_map(CommunicationType::Ping, &type_map)
|
||||
.add_typed_default(
|
||||
DataType::Description,
|
||||
DataValue::Str("protocol ping".into()),
|
||||
)
|
||||
.add_typed_default(
|
||||
DataType::Timestamp,
|
||||
DataValue::UnsignedNumber(timestamp as u128),
|
||||
)
|
||||
.with_sender(client_id);
|
||||
let Some(ping_id) = frame.id() else {
|
||||
let _ = on_error.call1(&JsValue::NULL, &js_error("ping frame has no id"));
|
||||
return;
|
||||
};
|
||||
let frame = frame
|
||||
.to_bytes()
|
||||
.map_err(|e| js_error(format!("encode ping failed: {}", e)));
|
||||
match frame {
|
||||
Ok(frame) => {
|
||||
if current_generation.get() != generation {
|
||||
return;
|
||||
}
|
||||
pending_pings.borrow_mut().insert(
|
||||
ping_id,
|
||||
PendingPing {
|
||||
generation,
|
||||
sent_at,
|
||||
},
|
||||
);
|
||||
if let Err(error) = transport.send_frame(&frame).await {
|
||||
if pending_pings
|
||||
.borrow()
|
||||
.get(&ping_id)
|
||||
.is_some_and(|ping| ping.generation == generation)
|
||||
{
|
||||
pending_pings.borrow_mut().remove(&ping_id);
|
||||
}
|
||||
let _ = on_error.call1(&JsValue::NULL, &error);
|
||||
}
|
||||
}
|
||||
Err(error) => {
|
||||
let _ = on_error.call1(&JsValue::NULL, &error);
|
||||
}
|
||||
}
|
||||
});
|
||||
}) as Box<dyn FnMut()>);
|
||||
|
||||
let set_interval =
|
||||
js_sys::Reflect::get(&js_sys::global(), &JsValue::from_str("setInterval"))?
|
||||
.dyn_into::<js_sys::Function>()?;
|
||||
let id = set_interval
|
||||
.call2(
|
||||
&JsValue::NULL,
|
||||
closure.as_ref().unchecked_ref(),
|
||||
&JsValue::from_f64(interval_ms as f64),
|
||||
)?
|
||||
.as_f64()
|
||||
.filter(|value| {
|
||||
value.is_finite()
|
||||
&& value.fract() == 0.0
|
||||
&& (i32::MIN as f64..=i32::MAX as f64).contains(value)
|
||||
})
|
||||
.and_then(|value| i32::try_from(value as i64).ok())
|
||||
.ok_or_else(|| js_error("setInterval did not return a valid id"))?;
|
||||
*self.ping_timer.borrow_mut() = Some(PingTimer { id, closure });
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[wasm_bindgen]
|
||||
pub fn stop_protocol_pings(&self) {
|
||||
self.pending_pings.borrow_mut().clear();
|
||||
self.ping_ms.set(None);
|
||||
let Some(timer) = self.ping_timer.borrow_mut().take() else {
|
||||
return;
|
||||
};
|
||||
if let Ok(clear_interval) =
|
||||
js_sys::Reflect::get(&js_sys::global(), &JsValue::from_str("clearInterval"))
|
||||
.and_then(|value| value.dyn_into::<js_sys::Function>())
|
||||
{
|
||||
let _ = clear_interval.call1(&JsValue::NULL, &JsValue::from_f64(timer.id as f64));
|
||||
}
|
||||
drop(timer.closure);
|
||||
}
|
||||
|
||||
pub(super) fn start_receive_loop(
|
||||
&self,
|
||||
transport: WasmTransport,
|
||||
generation: u32,
|
||||
client_id: u64,
|
||||
) -> bool {
|
||||
if self.connection_generation.get() != generation {
|
||||
transport.close();
|
||||
return false;
|
||||
}
|
||||
let loop_transport = transport.clone();
|
||||
self.attempt_transport.borrow_mut().take();
|
||||
*self.transport.borrow_mut() = Some(transport);
|
||||
self.set_state(ConnectionState::Connected);
|
||||
|
||||
let connection_generation = self.connection_generation.clone();
|
||||
let error_generation = connection_generation.clone();
|
||||
let state = self.state.clone();
|
||||
let pending_state_callbacks = self.pending_state_callbacks.clone();
|
||||
let state_callback = self.state_callback.as_ref().clone();
|
||||
let on_msg = self.on_message.clone();
|
||||
let on_err = self.on_error.clone();
|
||||
let subscriptions = self.subscriptions.clone();
|
||||
let pending_requests = self.pending_requests.clone();
|
||||
let loop_pending_requests = pending_requests.clone();
|
||||
let expired_requests = self.expired_requests.clone();
|
||||
let loop_expired_requests = expired_requests.clone();
|
||||
let ping_timer = self.ping_timer.clone();
|
||||
let pending_pings = self.pending_pings.clone();
|
||||
let loop_pending_pings = pending_pings.clone();
|
||||
let ping_ms = self.ping_ms.clone();
|
||||
let loop_ping_ms = ping_ms.clone();
|
||||
let pending_pipe_creations = self.pending_pipe_creations.clone();
|
||||
let expired_pipe_creations = self.expired_pipe_creations.clone();
|
||||
let pending_pipes = self.pending_pipes.clone();
|
||||
let loop_pending_pipes = pending_pipes.clone();
|
||||
let expected_pending_pipes = pending_pipes.clone();
|
||||
let on_pipe_request = self.on_pipe_request.clone();
|
||||
let loop_pipe_creations = pending_pipe_creations.clone();
|
||||
let loop_expired_pipe_creations = expired_pipe_creations.clone();
|
||||
let loop_generation = generation;
|
||||
let frame_generation = connection_generation.clone();
|
||||
let transport_for_cleanup = self.transport.clone();
|
||||
let connection_client_id = self.connection_client_id.clone();
|
||||
wasm_bindgen_futures::spawn_local(async move {
|
||||
loop_transport
|
||||
.receive_loop_with_pipes(
|
||||
move |frame: JsValue| {
|
||||
if frame_generation.get() != loop_generation {
|
||||
return;
|
||||
}
|
||||
let message_type = frame_type(&frame);
|
||||
if let Some(ref msg_type) = message_type {
|
||||
if msg_type == "PipeRequest" {
|
||||
let Some(pipe_id) = frame_id(&frame).filter(|id| *id != 0) else {
|
||||
return;
|
||||
};
|
||||
let description = frame_property(&frame, "data")
|
||||
.and_then(|data| {
|
||||
let desc = js_sys::Reflect::get(
|
||||
&data,
|
||||
&JsValue::from_str("Description"),
|
||||
)
|
||||
.ok()?;
|
||||
desc.as_string()
|
||||
})
|
||||
.unwrap_or_default();
|
||||
|
||||
let cb = on_pipe_request.borrow();
|
||||
if let Some(ref callback) = *cb {
|
||||
let obj = js_sys::Object::new();
|
||||
let _ = js_sys::Reflect::set(
|
||||
&obj,
|
||||
&"pipeId".into(),
|
||||
&JsValue::from_f64(pipe_id as f64),
|
||||
);
|
||||
let _ = js_sys::Reflect::set(
|
||||
&obj,
|
||||
&"description".into(),
|
||||
&JsValue::from_str(&description),
|
||||
);
|
||||
let _ = callback.call1(&JsValue::NULL, &obj.into());
|
||||
}
|
||||
return;
|
||||
}
|
||||
|
||||
if msg_type == "PipeResponse" {
|
||||
let Some(pipe_id) = frame_id(&frame).filter(|id| *id != 0) else {
|
||||
return;
|
||||
};
|
||||
let accepted = frame_property(&frame, "data")
|
||||
.and_then(|data| {
|
||||
let acc = js_sys::Reflect::get(
|
||||
&data,
|
||||
&JsValue::from_str("Accepted"),
|
||||
)
|
||||
.ok()?;
|
||||
acc.as_bool()
|
||||
})
|
||||
.unwrap_or(false);
|
||||
|
||||
let pending = {
|
||||
let mut pending = loop_pipe_creations.borrow_mut();
|
||||
if pending
|
||||
.get(&pipe_id)
|
||||
.is_some_and(|entry| entry.generation == loop_generation)
|
||||
{
|
||||
pending.remove(&pipe_id)
|
||||
} else {
|
||||
None
|
||||
}
|
||||
};
|
||||
if let Some(entry) = pending {
|
||||
let _ = entry.sender.send(Ok(accepted));
|
||||
} else {
|
||||
let _ = client_pipe::consume_expired_pipe_creation(
|
||||
&loop_expired_pipe_creations,
|
||||
pipe_id,
|
||||
);
|
||||
}
|
||||
return;
|
||||
}
|
||||
}
|
||||
|
||||
route_incoming_frame(
|
||||
&frame,
|
||||
loop_generation,
|
||||
&on_msg,
|
||||
&subscriptions,
|
||||
&loop_pending_requests,
|
||||
&loop_expired_requests,
|
||||
&loop_pending_pings,
|
||||
&loop_ping_ms,
|
||||
);
|
||||
},
|
||||
move |error| {
|
||||
if error_generation.get() == generation {
|
||||
let _ = on_err.call1(&JsValue::NULL, &error);
|
||||
}
|
||||
},
|
||||
move |pipe_reader: PipeReader| {
|
||||
let pipe_id = pipe_reader.pipe_id();
|
||||
let mut pending = loop_pending_pipes.borrow_mut();
|
||||
if pending
|
||||
.get(&pipe_id)
|
||||
.is_some_and(|entry| entry.generation == loop_generation)
|
||||
&& let Some(entry) = pending.remove(&pipe_id)
|
||||
{
|
||||
let _ = entry.sender.send(Ok(pipe_reader));
|
||||
}
|
||||
},
|
||||
move |pipe_id| {
|
||||
expected_pending_pipes
|
||||
.borrow()
|
||||
.get(&pipe_id)
|
||||
.is_some_and(|entry| entry.generation == loop_generation)
|
||||
},
|
||||
)
|
||||
.await;
|
||||
if connection_generation.get() != generation {
|
||||
return;
|
||||
}
|
||||
if let Some(current_transport) = transport_for_cleanup.borrow_mut().take() {
|
||||
current_transport.close();
|
||||
}
|
||||
set_shared_state(
|
||||
&state,
|
||||
&pending_state_callbacks,
|
||||
&state_callback,
|
||||
ConnectionState::Disconnected,
|
||||
);
|
||||
stop_ping_timer(&ping_timer);
|
||||
pending_pings.borrow_mut().clear();
|
||||
ping_ms.set(None);
|
||||
reject_pending_requests(&pending_requests, "disconnected");
|
||||
expired_requests.borrow_mut().clear();
|
||||
client_pipe::reject_pending_pipe_creations(&pending_pipe_creations, "disconnected");
|
||||
expired_pipe_creations.borrow_mut().clear();
|
||||
client_pipe::reject_pending_pipes(&pending_pipes, "disconnected");
|
||||
connection_client_id.set(0);
|
||||
});
|
||||
self.connection_client_id.set(client_id);
|
||||
true
|
||||
}
|
||||
|
||||
pub(super) fn reject_pending_requests(&self, message: &str) {
|
||||
reject_pending_requests(&self.pending_requests, message);
|
||||
self.expired_requests.borrow_mut().clear();
|
||||
}
|
||||
}
|
||||
|
|
@ -3,10 +3,8 @@ use std::collections::HashMap;
|
|||
use std::rc::Rc;
|
||||
|
||||
use futures_channel::oneshot;
|
||||
use futures_util::{FutureExt, pin_mut, select};
|
||||
use tracing::debug;
|
||||
use wasm_bindgen::prelude::*;
|
||||
use wasm_bindgen_futures::JsFuture;
|
||||
|
||||
use mtp_codec::{CommunicationType, CommunicationValue, DataType, DataValue};
|
||||
|
||||
|
|
@ -23,17 +21,12 @@ pub(crate) struct PendingRequest {
|
|||
|
||||
pub(crate) struct PendingPipeCreation {
|
||||
pub(crate) generation: u32,
|
||||
pub(crate) token: Rc<()>,
|
||||
pub(crate) sender: oneshot::Sender<Result<bool, JsValue>>,
|
||||
}
|
||||
pub(crate) type PendingPipeCreations = Rc<RefCell<HashMap<u32, PendingPipeCreation>>>;
|
||||
type PipeResponseReceiver = oneshot::Receiver<Result<bool, JsValue>>;
|
||||
type PipeResponseCell = Rc<RefCell<Option<PipeResponseReceiver>>>;
|
||||
|
||||
const DEFAULT_PIPE_CREATION_TIMEOUT_MS: u32 = 30_000;
|
||||
const EXPIRED_PIPE_CREATION_TOMBSTONE_TTL_MS: f64 = 60_000.0;
|
||||
const MAX_EXPIRED_PIPE_CREATION_TOMBSTONES: usize = 1024;
|
||||
|
||||
pub(crate) struct PendingPipe {
|
||||
pub(crate) generation: u32,
|
||||
pub(crate) sender: oneshot::Sender<Result<PipeReader, JsValue>>,
|
||||
|
|
@ -121,10 +114,6 @@ pub struct WasmPipeHandle {
|
|||
description: String,
|
||||
transport: WasmTransport,
|
||||
response_rx: PipeResponseCell,
|
||||
pending: PendingPipeCreations,
|
||||
expired: Rc<RefCell<HashMap<u32, f64>>>,
|
||||
generation: u32,
|
||||
token: Rc<()>,
|
||||
}
|
||||
|
||||
#[wasm_bindgen]
|
||||
|
|
@ -136,37 +125,9 @@ impl WasmPipeHandle {
|
|||
.take()
|
||||
.ok_or_else(|| js_error("handle already consumed"))?;
|
||||
|
||||
let response = rx.fuse();
|
||||
let timeout = wait_for_timeout(DEFAULT_PIPE_CREATION_TIMEOUT_MS).fuse();
|
||||
pin_mut!(response, timeout);
|
||||
let accepted = select! {
|
||||
result = response => match result {
|
||||
Ok(result) => result,
|
||||
Err(_) => {
|
||||
expire_pending_pipe_creation(
|
||||
&self.pending,
|
||||
&self.expired,
|
||||
self.pipe_id,
|
||||
self.generation,
|
||||
&self.token,
|
||||
);
|
||||
return Err(js_error("pipe handle channel closed"));
|
||||
}
|
||||
},
|
||||
result = timeout => {
|
||||
result?;
|
||||
expire_pending_pipe_creation(
|
||||
&self.pending,
|
||||
&self.expired,
|
||||
self.pipe_id,
|
||||
self.generation,
|
||||
&self.token,
|
||||
);
|
||||
return Err(js_error(format!(
|
||||
"pipe creation timed out after {DEFAULT_PIPE_CREATION_TIMEOUT_MS}ms"
|
||||
)));
|
||||
},
|
||||
};
|
||||
let accepted = rx
|
||||
.await
|
||||
.map_err(|_| js_error("pipe handle channel closed"))?;
|
||||
|
||||
match accepted {
|
||||
Ok(true) => {
|
||||
|
|
@ -192,18 +153,6 @@ impl WasmPipeHandle {
|
|||
}
|
||||
}
|
||||
|
||||
impl Drop for WasmPipeHandle {
|
||||
fn drop(&mut self) {
|
||||
expire_pending_pipe_creation(
|
||||
&self.pending,
|
||||
&self.expired,
|
||||
self.pipe_id,
|
||||
self.generation,
|
||||
&self.token,
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) fn random_pipe_id() -> Result<u32, JsValue> {
|
||||
let mut bytes = [0u8; 4];
|
||||
getrandom_v04::fill(&mut bytes).map_err(|_| js_error("rng failed"))?;
|
||||
|
|
@ -217,146 +166,6 @@ pub(crate) fn reject_pending_pipe_creations(pending: &PendingPipeCreations, mess
|
|||
}
|
||||
}
|
||||
|
||||
async fn wait_for_timeout(timeout_ms: u32) -> Result<(), JsValue> {
|
||||
let promise = js_sys::Promise::new(&mut |resolve, reject| {
|
||||
let result = js_sys::Reflect::get(&js_sys::global(), &JsValue::from_str("setTimeout"))
|
||||
.and_then(|value| value.dyn_into::<js_sys::Function>())
|
||||
.and_then(|set_timeout| {
|
||||
set_timeout.call2(
|
||||
&JsValue::NULL,
|
||||
&resolve,
|
||||
&JsValue::from_f64(timeout_ms as f64),
|
||||
)
|
||||
});
|
||||
if let Err(error) = result {
|
||||
let _ = reject.call1(&JsValue::NULL, &error);
|
||||
}
|
||||
});
|
||||
JsFuture::from(promise).await?;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn expire_pending_pipe_creation(
|
||||
pending: &PendingPipeCreations,
|
||||
expired: &Rc<RefCell<HashMap<u32, f64>>>,
|
||||
pipe_id: u32,
|
||||
generation: u32,
|
||||
token: &Rc<()>,
|
||||
) {
|
||||
let removed = {
|
||||
let mut pending = pending.borrow_mut();
|
||||
if pending
|
||||
.get(&pipe_id)
|
||||
.is_some_and(|entry| entry.generation == generation && Rc::ptr_eq(&entry.token, token))
|
||||
{
|
||||
pending.remove(&pipe_id);
|
||||
true
|
||||
} else {
|
||||
false
|
||||
}
|
||||
};
|
||||
if !removed {
|
||||
return;
|
||||
}
|
||||
let now = js_sys::Date::now();
|
||||
let mut expired = expired.borrow_mut();
|
||||
expired.retain(|_, expires_at| *expires_at > now);
|
||||
if expired.len() >= MAX_EXPIRED_PIPE_CREATION_TOMBSTONES
|
||||
&& let Some(oldest) = expired
|
||||
.iter()
|
||||
.min_by(|(_, left), (_, right)| left.total_cmp(right))
|
||||
.map(|(id, _)| *id)
|
||||
{
|
||||
expired.remove(&oldest);
|
||||
}
|
||||
expired.insert(pipe_id, now + EXPIRED_PIPE_CREATION_TOMBSTONE_TTL_MS);
|
||||
}
|
||||
|
||||
struct PendingPipeCreationGuard {
|
||||
pending: PendingPipeCreations,
|
||||
expired: Rc<RefCell<HashMap<u32, f64>>>,
|
||||
pipe_id: u32,
|
||||
generation: u32,
|
||||
token: Rc<()>,
|
||||
armed: bool,
|
||||
}
|
||||
|
||||
impl PendingPipeCreationGuard {
|
||||
fn new(
|
||||
pending: PendingPipeCreations,
|
||||
expired: Rc<RefCell<HashMap<u32, f64>>>,
|
||||
pipe_id: u32,
|
||||
generation: u32,
|
||||
token: Rc<()>,
|
||||
) -> Self {
|
||||
Self {
|
||||
pending,
|
||||
expired,
|
||||
pipe_id,
|
||||
generation,
|
||||
token,
|
||||
armed: true,
|
||||
}
|
||||
}
|
||||
|
||||
fn disarm(&mut self) {
|
||||
self.armed = false;
|
||||
}
|
||||
}
|
||||
|
||||
impl Drop for PendingPipeCreationGuard {
|
||||
fn drop(&mut self) {
|
||||
if self.armed {
|
||||
expire_pending_pipe_creation(
|
||||
&self.pending,
|
||||
&self.expired,
|
||||
self.pipe_id,
|
||||
self.generation,
|
||||
&self.token,
|
||||
);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
struct PendingPipeGuard {
|
||||
pending: PendingPipes,
|
||||
pipe_id: u32,
|
||||
generation: u32,
|
||||
}
|
||||
|
||||
impl PendingPipeGuard {
|
||||
fn new(pending: PendingPipes, pipe_id: u32, generation: u32) -> Self {
|
||||
Self {
|
||||
pending,
|
||||
pipe_id,
|
||||
generation,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl Drop for PendingPipeGuard {
|
||||
fn drop(&mut self) {
|
||||
remove_pending_pipe(&self.pending, self.pipe_id, self.generation);
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) fn consume_expired_pipe_creation(
|
||||
expired: &Rc<RefCell<HashMap<u32, f64>>>,
|
||||
pipe_id: u32,
|
||||
) -> bool {
|
||||
let now = js_sys::Date::now();
|
||||
let mut expired = expired.borrow_mut();
|
||||
expired.retain(|_, expires_at| *expires_at > now);
|
||||
expired.remove(&pipe_id).is_some()
|
||||
}
|
||||
|
||||
fn is_expired_pipe_creation(expired: &Rc<RefCell<HashMap<u32, f64>>>, pipe_id: u32) -> bool {
|
||||
let now = js_sys::Date::now();
|
||||
let mut expired = expired.borrow_mut();
|
||||
expired.retain(|_, expires_at| *expires_at > now);
|
||||
expired.contains_key(&pipe_id)
|
||||
}
|
||||
|
||||
pub(crate) fn reject_pending_pipes(pending: &PendingPipes, message: &str) {
|
||||
let pending = std::mem::take(&mut *pending.borrow_mut());
|
||||
for (_, entry) in pending {
|
||||
|
|
@ -369,26 +178,19 @@ pub(crate) async fn wasm_create_pipe(
|
|||
description: &str,
|
||||
pipe_id: u32,
|
||||
pending_pipe_creations: &PendingPipeCreations,
|
||||
expired_pipe_creations: &Rc<RefCell<HashMap<u32, f64>>>,
|
||||
generation: u32,
|
||||
current_generation: &Rc<std::cell::Cell<u32>>,
|
||||
) -> Result<WasmPipeHandle, JsValue> {
|
||||
let (tx, rx) = oneshot::channel();
|
||||
let token = Rc::new(());
|
||||
let mut pipe_id = pipe_id;
|
||||
for _ in 0..128 {
|
||||
let occupied = pipe_id == 0
|
||||
|| pending_pipe_creations.borrow().contains_key(&pipe_id)
|
||||
|| is_expired_pipe_creation(expired_pipe_creations, pipe_id);
|
||||
let occupied = pipe_id == 0 || pending_pipe_creations.borrow().contains_key(&pipe_id);
|
||||
if !occupied {
|
||||
break;
|
||||
}
|
||||
pipe_id = random_pipe_id()?;
|
||||
}
|
||||
if pipe_id == 0
|
||||
|| pending_pipe_creations.borrow().contains_key(&pipe_id)
|
||||
|| is_expired_pipe_creation(expired_pipe_creations, pipe_id)
|
||||
{
|
||||
if pipe_id == 0 || pending_pipe_creations.borrow().contains_key(&pipe_id) {
|
||||
return Err(js_error("could not allocate a unique pipe id"));
|
||||
}
|
||||
let type_map = transport.type_map();
|
||||
|
|
@ -405,17 +207,9 @@ pub(crate) async fn wasm_create_pipe(
|
|||
pipe_id,
|
||||
PendingPipeCreation {
|
||||
generation,
|
||||
token: token.clone(),
|
||||
sender: tx,
|
||||
},
|
||||
);
|
||||
let mut creation_guard = PendingPipeCreationGuard::new(
|
||||
pending_pipe_creations.clone(),
|
||||
expired_pipe_creations.clone(),
|
||||
pipe_id,
|
||||
generation,
|
||||
token.clone(),
|
||||
);
|
||||
debug!(
|
||||
target = "mtp.wasm",
|
||||
pipe_id,
|
||||
|
|
@ -424,22 +218,31 @@ pub(crate) async fn wasm_create_pipe(
|
|||
"sending pipe request"
|
||||
);
|
||||
if let Err(error) = transport.send_frame(&request_bytes).await {
|
||||
let mut pending = pending_pipe_creations.borrow_mut();
|
||||
if pending
|
||||
.get(&pipe_id)
|
||||
.is_some_and(|entry| entry.generation == generation)
|
||||
{
|
||||
pending.remove(&pipe_id);
|
||||
}
|
||||
return Err(error);
|
||||
}
|
||||
if current_generation.get() != generation {
|
||||
let mut pending = pending_pipe_creations.borrow_mut();
|
||||
if pending
|
||||
.get(&pipe_id)
|
||||
.is_some_and(|entry| entry.generation == generation)
|
||||
{
|
||||
pending.remove(&pipe_id);
|
||||
}
|
||||
return Err(js_error("connection attempt superseded"));
|
||||
}
|
||||
|
||||
creation_guard.disarm();
|
||||
Ok(WasmPipeHandle {
|
||||
pipe_id,
|
||||
description: description.to_string(),
|
||||
transport: transport.clone(),
|
||||
response_rx: Rc::new(RefCell::new(Some(rx))),
|
||||
pending: pending_pipe_creations.clone(),
|
||||
expired: expired_pipe_creations.clone(),
|
||||
generation,
|
||||
token,
|
||||
})
|
||||
}
|
||||
|
||||
|
|
@ -479,7 +282,6 @@ pub(crate) async fn wasm_accept_pipe(
|
|||
},
|
||||
);
|
||||
}
|
||||
let _acceptance_guard = PendingPipeGuard::new(pending_pipes.clone(), pipe_id, generation);
|
||||
|
||||
debug!(
|
||||
target = "mtp.wasm",
|
||||
|
|
@ -489,35 +291,6 @@ pub(crate) async fn wasm_accept_pipe(
|
|||
"sending pipe response"
|
||||
);
|
||||
if let Err(error) = transport.send_frame(&resp_bytes).await {
|
||||
return Err(error);
|
||||
}
|
||||
if current_generation.get() != generation {
|
||||
return Err(js_error("connection attempt superseded"));
|
||||
}
|
||||
|
||||
let response = rx.fuse();
|
||||
let timeout = wait_for_timeout(DEFAULT_PIPE_CREATION_TIMEOUT_MS).fuse();
|
||||
pin_mut!(response, timeout);
|
||||
let result = select! {
|
||||
result = response => match result {
|
||||
Ok(result) => result,
|
||||
Err(_) => {
|
||||
remove_pending_pipe(pending_pipes, pipe_id, generation);
|
||||
return Err(js_error("pipe closed before stream arrived"));
|
||||
}
|
||||
},
|
||||
result = timeout => {
|
||||
result?;
|
||||
remove_pending_pipe(pending_pipes, pipe_id, generation);
|
||||
return Err(js_error(format!(
|
||||
"pipe acceptance timed out after {DEFAULT_PIPE_CREATION_TIMEOUT_MS}ms"
|
||||
)));
|
||||
},
|
||||
};
|
||||
result
|
||||
}
|
||||
|
||||
fn remove_pending_pipe(pending_pipes: &PendingPipes, pipe_id: u32, generation: u32) {
|
||||
let mut pending = pending_pipes.borrow_mut();
|
||||
if pending
|
||||
.get(&pipe_id)
|
||||
|
|
@ -525,6 +298,21 @@ fn remove_pending_pipe(pending_pipes: &PendingPipes, pipe_id: u32, generation: u
|
|||
{
|
||||
pending.remove(&pipe_id);
|
||||
}
|
||||
return Err(error);
|
||||
}
|
||||
if current_generation.get() != generation {
|
||||
let mut pending = pending_pipes.borrow_mut();
|
||||
if pending
|
||||
.get(&pipe_id)
|
||||
.is_some_and(|entry| entry.generation == generation)
|
||||
{
|
||||
pending.remove(&pipe_id);
|
||||
}
|
||||
return Err(js_error("connection attempt superseded"));
|
||||
}
|
||||
|
||||
rx.await
|
||||
.map_err(|_| js_error("pipe closed before stream arrived"))?
|
||||
}
|
||||
|
||||
pub(crate) async fn wasm_deny_pipe(transport: &WasmTransport, pipe_id: u32) -> Result<(), JsValue> {
|
||||
|
|
|
|||
|
|
@ -1,6 +1,5 @@
|
|||
use wasm_bindgen::prelude::*;
|
||||
|
||||
#[derive(Clone)]
|
||||
#[wasm_bindgen]
|
||||
pub struct ConnectionConfig {
|
||||
pub(crate) url: String,
|
||||
|
|
|
|||
|
|
@ -2,8 +2,8 @@ use wasm_bindgen::prelude::*;
|
|||
use zeroize::Zeroizing;
|
||||
|
||||
use mtp_codec::{
|
||||
DataValue, DecodeLimits, EncodeLimits, MtpProtectionPurpose, PROTOCOL_VERSION,
|
||||
ProtectionPolicy, ProtectionPurpose, SealedRelayBuilder, SignaturePolicy, TypeMap,
|
||||
DataValue, MtpProtectionPurpose, PROTOCOL_VERSION, ProtectionPolicy, ProtectionPurpose,
|
||||
SealedRelayBuilder, SignaturePolicy, TypeMap,
|
||||
};
|
||||
use mtp_crypto::{
|
||||
AeadDecrypt, AeadEncrypt, DualSigner, Ed25519Signer, HybridKem, KemPrivateKey, KemPublicKey,
|
||||
|
|
@ -12,18 +12,10 @@ use mtp_crypto::{
|
|||
};
|
||||
|
||||
use crate::error::{from_protection_error, js_error};
|
||||
use crate::relay::{decode_error, decode_frame, relay_error, structured_error};
|
||||
use crate::relay::{decode_frame, relay_error, structured_error};
|
||||
|
||||
fn decode_data_value(value: &[u8]) -> Result<DataValue, JsValue> {
|
||||
DataValue::try_from_bytes_with_limits(value, DecodeLimits::default()).map_err(|error| {
|
||||
let value = decode_error(error, "DataValue decoding failed");
|
||||
let _ = js_sys::Reflect::set(
|
||||
&value,
|
||||
&JsValue::from_str("code"),
|
||||
&JsValue::from_str("invalid-data-value"),
|
||||
);
|
||||
value
|
||||
})
|
||||
DataValue::from_bytes(value).ok_or_else(|| js_error("invalid DataValue"))
|
||||
}
|
||||
|
||||
fn decode_public_key_bundle(
|
||||
|
|
@ -82,19 +74,10 @@ pub struct WasmKeyring {
|
|||
|
||||
#[wasm_bindgen]
|
||||
impl WasmKeyring {
|
||||
/// Serialise the keyring to bytes and report malformed caller-owned
|
||||
/// material as a JavaScript exception.
|
||||
/// Serialise the keyring to bytes.
|
||||
#[wasm_bindgen]
|
||||
pub fn to_bytes(&self) -> Result<Vec<u8>, JsValue> {
|
||||
self.try_to_bytes()
|
||||
}
|
||||
|
||||
#[wasm_bindgen]
|
||||
pub fn try_to_bytes(&self) -> Result<Vec<u8>, JsValue> {
|
||||
self.inner
|
||||
.try_to_bytes()
|
||||
.map(|bytes| bytes.to_vec())
|
||||
.map_err(|error| js_error(format!("Keyring serialization failed: {error}")))
|
||||
pub fn to_bytes(&self) -> Vec<u8> {
|
||||
self.inner.to_bytes().to_vec()
|
||||
}
|
||||
|
||||
/// Deserialise a keyring from bytes.
|
||||
|
|
@ -135,17 +118,8 @@ impl WasmKeyring {
|
|||
|
||||
/// Generate a full keyring with KEM, ML-DSA, and Ed25519 keys.
|
||||
#[wasm_bindgen]
|
||||
pub fn keyring_generate() -> Result<Vec<u8>, JsValue> {
|
||||
keyring_generate_checked()
|
||||
}
|
||||
|
||||
/// Generate a full keyring and report serialization failures to JavaScript.
|
||||
#[wasm_bindgen]
|
||||
pub fn keyring_generate_checked() -> Result<Vec<u8>, JsValue> {
|
||||
Keyring::generate()
|
||||
.try_to_bytes()
|
||||
.map(|bytes| bytes.to_vec())
|
||||
.map_err(|error| js_error(format!("generated keyring serialization failed: {error}")))
|
||||
pub fn keyring_generate() -> Vec<u8> {
|
||||
Keyring::generate().to_bytes().to_vec()
|
||||
}
|
||||
|
||||
/// Build a [`Keyring`] containing only an Ed25519 keypair (no KEM, no ML-DSA).
|
||||
|
|
@ -168,10 +142,7 @@ pub fn keyring_from_ed25519(secret_key: &[u8], public_key: &[u8]) -> Result<Vec<
|
|||
SignaturePublicKey::new(public_key.to_vec()),
|
||||
SignaturePrivateKey::new(secret_key.to_vec()),
|
||||
);
|
||||
keyring
|
||||
.try_to_bytes()
|
||||
.map(|bytes| bytes.to_vec())
|
||||
.map_err(|error| js_error(format!("Keyring serialization failed: {error}")))
|
||||
Ok(keyring.to_bytes().to_vec())
|
||||
}
|
||||
|
||||
// ===========================================================================
|
||||
|
|
@ -201,15 +172,8 @@ impl WasmPublicKeyBundle {
|
|||
}
|
||||
|
||||
#[wasm_bindgen]
|
||||
pub fn to_bytes(&self) -> Result<Vec<u8>, JsValue> {
|
||||
self.try_to_bytes()
|
||||
}
|
||||
|
||||
#[wasm_bindgen]
|
||||
pub fn try_to_bytes(&self) -> Result<Vec<u8>, JsValue> {
|
||||
self.inner
|
||||
.try_as_bytes()
|
||||
.map_err(|error| js_error(format!("public key bundle serialization failed: {error}")))
|
||||
pub fn to_bytes(&self) -> Vec<u8> {
|
||||
self.inner.as_bytes()
|
||||
}
|
||||
|
||||
#[wasm_bindgen]
|
||||
|
|
@ -464,14 +428,6 @@ pub fn wasm_sha256_double(data: &[u8]) -> Vec<u8> {
|
|||
// KDF
|
||||
// ===========================================================================
|
||||
|
||||
/// Length, in bytes, of symmetric keys produced by the MTP key-derivation
|
||||
/// bindings. SDKs should query this instead of duplicating the crypto
|
||||
/// primitive's output size.
|
||||
#[wasm_bindgen]
|
||||
pub fn mtp_symmetric_key_length() -> u32 {
|
||||
32
|
||||
}
|
||||
|
||||
/// HKDF-expand: derive `len` bytes from `ikm` with `salt` and `info`.
|
||||
#[wasm_bindgen]
|
||||
pub fn wasm_hkdf_expand(
|
||||
|
|
@ -496,21 +452,6 @@ pub fn wasm_derive_encryption_key(
|
|||
.map_err(|e| js_error(format!("derive_encryption_key failed: {}", e)))
|
||||
}
|
||||
|
||||
/// Derive a 32-byte key from a passphrase using explicit Argon2id parameters.
|
||||
/// The salt and parameters are part of the caller's protected-data format.
|
||||
#[wasm_bindgen]
|
||||
pub fn wasm_argon2id(
|
||||
passphrase: &[u8],
|
||||
salt: &[u8],
|
||||
memory_kib: u32,
|
||||
iterations: u32,
|
||||
lanes: u32,
|
||||
) -> Result<Vec<u8>, JsValue> {
|
||||
mtp_crypto::derive_password_key(passphrase, salt, memory_kib, iterations, lanes)
|
||||
.map(|key| key.to_vec())
|
||||
.map_err(|e| js_error(format!("argon2id password derivation failed: {e}")))
|
||||
}
|
||||
|
||||
/// Signature suites accepted by high-level protected-value APIs.
|
||||
pub const PROTECTION_SIGNATURE_SUITE_ED25519: u8 = 0x01;
|
||||
pub const PROTECTION_SIGNATURE_SUITE_DUAL: u8 = 0x03;
|
||||
|
|
@ -623,11 +564,10 @@ pub fn verify_data_value_with_policy(
|
|||
let value = decode_data_value(value)?;
|
||||
let bundle = decode_public_key_bundle(public_key_bundle, None)?;
|
||||
let result = if signature_suite == 0 {
|
||||
value.verify_with_policy(
|
||||
value.verify(
|
||||
expected_signer_id,
|
||||
&bundle,
|
||||
ProtectionPurpose::from(expected_purpose),
|
||||
ProtectionPolicy::any_supported(),
|
||||
)
|
||||
} else {
|
||||
value.verify_with_policy(
|
||||
|
|
@ -710,11 +650,7 @@ pub fn decrypt_data_value_with_keyrings(
|
|||
let keyrings = keyrings_from_js(&keyrings)?;
|
||||
let references: Vec<&Keyring> = keyrings.iter().collect();
|
||||
value
|
||||
.decrypt_with_keyrings_and_limits(
|
||||
&references,
|
||||
ProtectionPurpose::from(expected_purpose),
|
||||
DecodeLimits::default(),
|
||||
)
|
||||
.decrypt_with_keyrings(&references, ProtectionPurpose::from(expected_purpose))
|
||||
.map_err(from_protection_error)?
|
||||
.to_bytes()
|
||||
.map_err(|e| js_error(format!("decryption failed: {e}")))
|
||||
|
|
@ -810,13 +746,6 @@ pub fn mtp_protection_signature_suite_dual() -> u8 {
|
|||
PROTECTION_SIGNATURE_SUITE_DUAL
|
||||
}
|
||||
|
||||
/// Explicit compatibility policy value accepting any signature suite
|
||||
/// supported by this WASM build. New callers should prefer a fixed suite.
|
||||
#[wasm_bindgen]
|
||||
pub fn mtp_protection_signature_suite_any_supported() -> u8 {
|
||||
0
|
||||
}
|
||||
|
||||
/// Forward a sealed relay frame to another clear next hop without opening or
|
||||
/// re-encoding its authenticated encrypted payload.
|
||||
#[wasm_bindgen]
|
||||
|
|
@ -847,27 +776,12 @@ fn build_encrypted_relay_frame_impl(
|
|||
signer: &dyn SignatureScheme,
|
||||
metadata_recipient_public_key_bundles: JsValue,
|
||||
content_recipient_public_key_bundles: JsValue,
|
||||
limits: JsValue,
|
||||
) -> Result<Vec<u8>, JsValue> {
|
||||
let tm = TypeMap::new(PROTOCOL_VERSION);
|
||||
let encode_limits = if limits.is_null() || limits.is_undefined() {
|
||||
EncodeLimits::default()
|
||||
} else {
|
||||
crate::client::encode_limits_from_js(&limits)?
|
||||
};
|
||||
let relay_options =
|
||||
crate::relay::relay_open_options(ProtectionPolicy::any_supported(), &limits)?;
|
||||
let application_content =
|
||||
crate::frame::js_to_data_value_with_limits(&data, &tm, encode_limits)?;
|
||||
let application_content = crate::frame::js_to_data_value(&data, &tm)?;
|
||||
let application_metadata = encoded_metadata
|
||||
.as_deref()
|
||||
.map(|bytes| {
|
||||
DataValue::try_from_bytes_with_limits(
|
||||
bytes,
|
||||
DecodeLimits::for_transport_message_size(encode_limits.max_output_size as u64),
|
||||
)
|
||||
.map_err(|error| crate::relay::decode_error(error, "metadata decoding failed"))
|
||||
})
|
||||
.map(decode_data_value)
|
||||
.transpose()?;
|
||||
let content_recipients = public_key_bundles_from_js(&content_recipient_public_key_bundles)?;
|
||||
let metadata_recipients = public_key_bundles_from_js(&metadata_recipient_public_key_bundles)?;
|
||||
|
|
@ -884,8 +798,6 @@ fn build_encrypted_relay_frame_impl(
|
|||
.created_at(created_at)
|
||||
.metadata_recipients(metadata_recipients)
|
||||
.content_recipients(content_recipients)
|
||||
.encode_limits(encode_limits)
|
||||
.protected_limits(relay_options.protected_limits)
|
||||
.type_map(&tm);
|
||||
let builder = match application_metadata {
|
||||
Some(metadata) => builder.metadata(metadata),
|
||||
|
|
@ -895,7 +807,7 @@ fn build_encrypted_relay_frame_impl(
|
|||
builder
|
||||
.build()
|
||||
.map_err(relay_error)?
|
||||
.to_bytes_with_limits(encode_limits)
|
||||
.to_bytes()
|
||||
.map_err(|e| js_error(format!("relay frame encoding failed: {e}")))
|
||||
}
|
||||
|
||||
|
|
@ -932,45 +844,6 @@ pub fn build_encrypted_relay_frame_with_keyring(
|
|||
&signer,
|
||||
metadata_recipient_public_key_bundles,
|
||||
content_recipient_public_key_bundles,
|
||||
JsValue::UNDEFINED,
|
||||
)
|
||||
}
|
||||
|
||||
/// Build a sealed relay frame with explicit encoder and semantic field
|
||||
/// limits. The same limits are applied by the native relay builder.
|
||||
#[wasm_bindgen]
|
||||
#[allow(clippy::too_many_arguments)]
|
||||
pub fn build_encrypted_relay_frame_with_keyring_with_limits(
|
||||
message_type: &str,
|
||||
data: JsValue,
|
||||
signer_id: u64,
|
||||
final_recipient_id: u64,
|
||||
next_hop_id: u64,
|
||||
message_id: &str,
|
||||
created_at: u64,
|
||||
encoded_metadata: Option<Vec<u8>>,
|
||||
keyring_bytes: &[u8],
|
||||
signature_suite: u8,
|
||||
metadata_recipient_public_key_bundles: JsValue,
|
||||
content_recipient_public_key_bundles: JsValue,
|
||||
limits: JsValue,
|
||||
) -> Result<Vec<u8>, JsValue> {
|
||||
let keyring = Keyring::from_bytes(keyring_bytes)
|
||||
.map_err(|e| js_error(format!("keyring initialization failed: {e}")))?;
|
||||
let signer = relay_signer_from_keyring(&keyring, signature_suite)?;
|
||||
build_encrypted_relay_frame_impl(
|
||||
message_type,
|
||||
data,
|
||||
signer_id,
|
||||
final_recipient_id,
|
||||
next_hop_id,
|
||||
message_id,
|
||||
created_at,
|
||||
encoded_metadata,
|
||||
&signer,
|
||||
metadata_recipient_public_key_bundles,
|
||||
content_recipient_public_key_bundles,
|
||||
limits,
|
||||
)
|
||||
}
|
||||
|
||||
|
|
@ -1018,7 +891,7 @@ mod tests {
|
|||
},
|
||||
};
|
||||
|
||||
let bytes = bundle.try_to_bytes().expect("bundle serialization");
|
||||
let bytes = bundle.to_bytes();
|
||||
let restored = WasmPublicKeyBundle::from_bytes_unvalidated(&bytes)
|
||||
.expect("from_bytes_unvalidated failed");
|
||||
assert_eq!(restored.sig_cl_public_key(), pk);
|
||||
|
|
@ -1225,7 +1098,7 @@ mod tests {
|
|||
let value = DataValue::Str("signed through wasm".into())
|
||||
.to_bytes()
|
||||
.expect("value encoding failed");
|
||||
let keyring_bytes = keyring.try_to_bytes().expect("keyring serialization");
|
||||
let keyring_bytes = keyring.to_bytes();
|
||||
let signed = sign_data_value_with_keyring(
|
||||
&value,
|
||||
0xfeed_beef,
|
||||
|
|
@ -1238,7 +1111,7 @@ mod tests {
|
|||
|
||||
verify_data_value_with_policy(
|
||||
&signed,
|
||||
&bundle.try_as_bytes().expect("bundle serialization"),
|
||||
&bundle.as_bytes(),
|
||||
0xfeed_beef,
|
||||
7,
|
||||
PROTECTION_SIGNATURE_SUITE_ED25519,
|
||||
|
|
@ -1248,7 +1121,7 @@ mod tests {
|
|||
assert!(
|
||||
verify_data_value_with_policy(
|
||||
&signed,
|
||||
&wrong_bundle.try_as_bytes().expect("bundle serialization"),
|
||||
&wrong_bundle.as_bytes(),
|
||||
0xfeed_beef,
|
||||
7,
|
||||
PROTECTION_SIGNATURE_SUITE_ED25519,
|
||||
|
|
@ -1264,29 +1137,21 @@ mod tests {
|
|||
let value = DataValue::Array(vec![DataValue::BoolTrue, DataValue::UnsignedNumber(42)])
|
||||
.to_bytes()
|
||||
.expect("value encoding failed");
|
||||
let recipient_bytes = recipient.try_as_bytes().expect("recipient serialization");
|
||||
let encrypted =
|
||||
encrypt_data_value(&value, &recipient_bytes, 9).expect("encrypt_data_value failed");
|
||||
let keyring_bytes = keyring.try_to_bytes().expect("keyring serialization");
|
||||
let decrypted =
|
||||
decrypt_data_value(&encrypted, &keyring_bytes, 9).expect("decrypt_data_value failed");
|
||||
let encrypted = encrypt_data_value(&value, &recipient.as_bytes(), 9)
|
||||
.expect("encrypt_data_value failed");
|
||||
let decrypted = decrypt_data_value(&encrypted, &keyring.to_bytes(), 9)
|
||||
.expect("decrypt_data_value failed");
|
||||
|
||||
assert_eq!(decrypted, value);
|
||||
|
||||
let second_keyring = Keyring::generate();
|
||||
let second_recipient = second_keyring.public_key_bundle();
|
||||
let recipients = js_sys::Array::new();
|
||||
let second_recipient_bytes = second_recipient
|
||||
.try_as_bytes()
|
||||
.expect("second recipient serialization");
|
||||
recipients.push(&js_sys::Uint8Array::from(&recipient_bytes[..]));
|
||||
recipients.push(&js_sys::Uint8Array::from(&second_recipient_bytes[..]));
|
||||
recipients.push(&js_sys::Uint8Array::from(&recipient.as_bytes()[..]));
|
||||
recipients.push(&js_sys::Uint8Array::from(&second_recipient.as_bytes()[..]));
|
||||
let multi = encrypt_data_value_for_recipients(&value, recipients.into(), 9)
|
||||
.expect("multi-recipient encryption failed");
|
||||
let second_keyring_bytes = second_keyring
|
||||
.try_to_bytes()
|
||||
.expect("second keyring serialization");
|
||||
let opened_by_second = decrypt_data_value(&multi, &second_keyring_bytes, 9)
|
||||
let opened_by_second = decrypt_data_value(&multi, &second_keyring.to_bytes(), 9)
|
||||
.expect("second recipient could not decrypt");
|
||||
assert_eq!(opened_by_second, value);
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1,13 +1,9 @@
|
|||
use wasm_bindgen::{JsCast, prelude::*};
|
||||
|
||||
use mtp_codec::{
|
||||
CommunicationType, CommunicationValue, DataType, DataValue, DecodeLimits, EncodeLimits,
|
||||
PROTOCOL_VERSION,
|
||||
};
|
||||
use mtp_codec::{CommunicationType, CommunicationValue, DataType, DataValue, PROTOCOL_VERSION};
|
||||
use mtp_type_map::TypeMap;
|
||||
|
||||
use crate::error::js_error;
|
||||
use crate::relay::decode_error;
|
||||
|
||||
#[wasm_bindgen(typescript_custom_section)]
|
||||
const PARSED_FRAME_TS: &'static str = r#"
|
||||
|
|
@ -141,49 +137,7 @@ pub(crate) fn data_value_to_js(value: &DataValue, tm: &TypeMap) -> Result<JsValu
|
|||
}
|
||||
|
||||
const MAX_SAFE_INT: f64 = 9007199254740991.0; // 2^53 - 1
|
||||
|
||||
struct JsDataValueEncodeContext {
|
||||
limits: EncodeLimits,
|
||||
values: usize,
|
||||
}
|
||||
|
||||
impl JsDataValueEncodeContext {
|
||||
fn visit(&mut self, depth: usize) -> Result<(), JsValue> {
|
||||
if depth > self.limits.max_depth {
|
||||
return Err(js_error("MTP DataValue nesting-depth limit exceeded"));
|
||||
}
|
||||
self.values = self
|
||||
.values
|
||||
.checked_add(1)
|
||||
.ok_or_else(|| js_error("MTP DataValue value-count limit exceeded"))?;
|
||||
if self.values > self.limits.max_values {
|
||||
return Err(js_error("MTP DataValue value-count limit exceeded"));
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
pub(crate) fn js_to_data_value(value: &JsValue, tm: &TypeMap) -> Result<DataValue, JsValue> {
|
||||
js_to_data_value_with_limits(value, tm, EncodeLimits::default())
|
||||
}
|
||||
|
||||
pub(crate) fn js_to_data_value_with_limits(
|
||||
value: &JsValue,
|
||||
tm: &TypeMap,
|
||||
limits: EncodeLimits,
|
||||
) -> Result<DataValue, JsValue> {
|
||||
let mut context = JsDataValueEncodeContext { limits, values: 0 };
|
||||
js_to_data_value_with_context(value, tm, &mut context, 0)
|
||||
}
|
||||
|
||||
fn js_to_data_value_with_context(
|
||||
value: &JsValue,
|
||||
tm: &TypeMap,
|
||||
context: &mut JsDataValueEncodeContext,
|
||||
depth: usize,
|
||||
) -> Result<DataValue, JsValue> {
|
||||
context.visit(depth)?;
|
||||
if value.is_null() || value.is_undefined() {
|
||||
return Ok(DataValue::Null);
|
||||
}
|
||||
|
|
@ -198,17 +152,9 @@ fn js_to_data_value_with_context(
|
|||
}
|
||||
if js_sys::Array::is_array(value) {
|
||||
let array = js_sys::Array::from(value);
|
||||
if array.length() as usize > context.limits.max_values {
|
||||
return Err(js_error("MTP DataValue value-count limit exceeded"));
|
||||
}
|
||||
let mut values = Vec::with_capacity(array.length() as usize);
|
||||
for item in array.iter() {
|
||||
values.push(js_to_data_value_with_context(
|
||||
&item,
|
||||
tm,
|
||||
context,
|
||||
depth + 1,
|
||||
)?);
|
||||
values.push(js_to_data_value(&item, tm)?);
|
||||
}
|
||||
return Ok(DataValue::Array(values));
|
||||
}
|
||||
|
|
@ -252,9 +198,6 @@ fn js_to_data_value_with_context(
|
|||
if value.is_object() {
|
||||
let object = js_sys::Object::from(value.clone());
|
||||
let keys = js_sys::Object::keys(&object);
|
||||
if keys.length() as usize > context.limits.max_values {
|
||||
return Err(js_error("MTP DataValue value-count limit exceeded"));
|
||||
}
|
||||
let mut entries = Vec::with_capacity(keys.length() as usize);
|
||||
for key in keys.iter() {
|
||||
let key = key
|
||||
|
|
@ -269,10 +212,7 @@ fn js_to_data_value_with_context(
|
|||
tm.version
|
||||
))
|
||||
})?;
|
||||
entries.push((
|
||||
id,
|
||||
js_to_data_value_with_context(&value, tm, context, depth + 1)?,
|
||||
));
|
||||
entries.push((id, js_to_data_value(&value, tm)?));
|
||||
}
|
||||
return Ok(DataValue::Container(entries));
|
||||
}
|
||||
|
|
@ -342,16 +282,15 @@ fn apply_frame_options(
|
|||
}
|
||||
|
||||
pub(crate) fn parse_frame_value(frame: &[u8]) -> Result<JsValue, JsValue> {
|
||||
parse_frame_value_with_limits(frame, &TypeMap::latest(), DecodeLimits::default())
|
||||
parse_frame_value_with_type_map(frame, &TypeMap::latest())
|
||||
}
|
||||
|
||||
pub(crate) fn parse_frame_value_with_limits(
|
||||
pub(crate) fn parse_frame_value_with_type_map(
|
||||
frame: &[u8],
|
||||
type_map: &TypeMap,
|
||||
limits: DecodeLimits,
|
||||
) -> Result<JsValue, JsValue> {
|
||||
let comm = CommunicationValue::try_from_bytes_with_type_map_and_limits(frame, type_map, limits)
|
||||
.map_err(|error| decode_error(error, "parse failed"))?;
|
||||
let comm = CommunicationValue::from_bytes_with(frame, type_map)
|
||||
.map_err(|e| js_error(format!("parse failed: {}", e)))?;
|
||||
let tm = type_map;
|
||||
let obj = js_sys::Object::new();
|
||||
|
||||
|
|
@ -416,8 +355,8 @@ pub fn build_ping_frame(
|
|||
/// Parse an auth response frame into a JS object.
|
||||
#[wasm_bindgen(unchecked_return_type = "AuthResponse")]
|
||||
pub fn parse_auth_response(response: &[u8]) -> Result<JsValue, JsValue> {
|
||||
let comm = CommunicationValue::try_from_bytes_with_limits(response, DecodeLimits::default())
|
||||
.map_err(|error| decode_error(error, "parse failed"))?;
|
||||
let comm = CommunicationValue::from_bytes(response)
|
||||
.map_err(|e| js_error(format!("parse failed: {}", e)))?;
|
||||
|
||||
let connected = matches!(
|
||||
comm.get_data(DataType::Connected),
|
||||
|
|
@ -479,8 +418,8 @@ pub fn parse_auth_response(response: &[u8]) -> Result<JsValue, JsValue> {
|
|||
/// Parse any MTP frame into the human-readable CommunicationValue display form.
|
||||
#[wasm_bindgen]
|
||||
pub fn format_frame(frame: &[u8]) -> Result<String, JsValue> {
|
||||
let comm = CommunicationValue::try_from_bytes_with_limits(frame, DecodeLimits::default())
|
||||
.map_err(|error| decode_error(error, "parse failed"))?;
|
||||
let comm = CommunicationValue::from_bytes(frame)
|
||||
.map_err(|e| js_error(format!("parse failed: {}", e)))?;
|
||||
Ok(comm.to_string())
|
||||
}
|
||||
|
||||
|
|
@ -490,80 +429,31 @@ pub fn parse_frame(frame: &[u8]) -> Result<JsValue, JsValue> {
|
|||
parse_frame_value(frame)
|
||||
}
|
||||
|
||||
/// Parse a frame with the caller's bounded receive policy. The compatibility
|
||||
/// `parse_frame` entry point retains the default policy for existing callers.
|
||||
#[wasm_bindgen(unchecked_return_type = "ParsedFrame")]
|
||||
pub fn parse_frame_with_limits(frame: &[u8], limits: JsValue) -> Result<JsValue, JsValue> {
|
||||
let limits = crate::client::decode_limits_from_js(&limits)?;
|
||||
parse_frame_value_with_limits(frame, &TypeMap::latest(), limits)
|
||||
}
|
||||
|
||||
/// Parse a standalone serialized `DataValue` into the same structured form
|
||||
/// used for frame payloads. Protected values remain opaque until the caller
|
||||
/// explicitly opens and verifies them.
|
||||
#[wasm_bindgen(unchecked_return_type = "ParsedDataValue")]
|
||||
pub fn parse_data_value(value: &[u8]) -> Result<JsValue, JsValue> {
|
||||
parse_data_value_with_decode_limits(value, DecodeLimits::default())
|
||||
}
|
||||
|
||||
fn parse_data_value_with_decode_limits(
|
||||
value: &[u8],
|
||||
limits: DecodeLimits,
|
||||
) -> Result<JsValue, JsValue> {
|
||||
let value = DataValue::try_from_bytes_with_limits(value, limits)
|
||||
.map_err(|error| decode_error(error, "decode data value failed"))?;
|
||||
let value = DataValue::from_bytes(value).ok_or_else(|| js_error("invalid DataValue"))?;
|
||||
let tm = TypeMap::new(PROTOCOL_VERSION);
|
||||
data_value_to_js(&value, &tm)
|
||||
}
|
||||
|
||||
/// Parse a standalone serialized `DataValue` with the caller's bounded
|
||||
/// receive policy. The compatibility `parse_data_value` entry point retains
|
||||
/// the default policy for existing callers.
|
||||
#[wasm_bindgen(unchecked_return_type = "ParsedDataValue")]
|
||||
pub fn parse_data_value_with_limits(value: &[u8], limits: JsValue) -> Result<JsValue, JsValue> {
|
||||
let limits = crate::client::decode_limits_from_js(&limits)?;
|
||||
parse_data_value_with_decode_limits(value, limits)
|
||||
}
|
||||
|
||||
/// Encode one standalone `DataValue` using the negotiated/current type map.
|
||||
#[wasm_bindgen]
|
||||
pub fn encode_data_value(value: JsValue) -> Result<Vec<u8>, JsValue> {
|
||||
encode_data_value_with_encode_limits(value, EncodeLimits::default())
|
||||
}
|
||||
|
||||
fn encode_data_value_with_encode_limits(
|
||||
value: JsValue,
|
||||
limits: EncodeLimits,
|
||||
) -> Result<Vec<u8>, JsValue> {
|
||||
let tm = TypeMap::new(PROTOCOL_VERSION);
|
||||
js_to_data_value_with_limits(&value, &tm, limits)?
|
||||
.to_bytes_with_limits(limits)
|
||||
js_to_data_value(&value, &tm)?
|
||||
.to_bytes()
|
||||
.map_err(|e| js_error(format!("encode data value failed: {e}")))
|
||||
}
|
||||
|
||||
/// Encode one standalone `DataValue` using explicit recursion and output
|
||||
/// limits. The compatibility entry point above keeps the historical default.
|
||||
#[wasm_bindgen]
|
||||
pub fn encode_data_value_with_limits(value: JsValue, limits: JsValue) -> Result<Vec<u8>, JsValue> {
|
||||
let limits = crate::client::encode_limits_from_js(&limits)?;
|
||||
encode_data_value_with_encode_limits(value, limits)
|
||||
}
|
||||
|
||||
/// Build a typed MTP frame using generated communication/data type names.
|
||||
#[wasm_bindgen]
|
||||
pub fn build_frame(
|
||||
message_type: &str,
|
||||
data: JsValue,
|
||||
options: JsValue,
|
||||
) -> Result<Vec<u8>, JsValue> {
|
||||
build_frame_with_encode_limits(message_type, data, options, EncodeLimits::default())
|
||||
}
|
||||
|
||||
fn build_frame_with_encode_limits(
|
||||
message_type: &str,
|
||||
data: JsValue,
|
||||
options: JsValue,
|
||||
limits: EncodeLimits,
|
||||
) -> Result<Vec<u8>, JsValue> {
|
||||
let comm_type = CommunicationType::from_name(message_type)
|
||||
.ok_or_else(|| js_error(format!("unknown communication type: {message_type}")))?;
|
||||
|
|
@ -588,7 +478,7 @@ fn build_frame_with_encode_limits(
|
|||
))
|
||||
})?;
|
||||
msg = msg
|
||||
.add_data(id, js_to_data_value_with_limits(&value, &tm, limits)?)
|
||||
.add_data(id, js_to_data_value(&value, &tm)?)
|
||||
.map_err(|e| js_error(format!("add data failed: {e}")))?;
|
||||
}
|
||||
} else if !data.is_null() && !data.is_undefined() {
|
||||
|
|
@ -597,24 +487,10 @@ fn build_frame_with_encode_limits(
|
|||
));
|
||||
}
|
||||
|
||||
msg.to_bytes_with_limits(limits)
|
||||
msg.to_bytes()
|
||||
.map_err(|e| js_error(format!("encode failed: {}", e)))
|
||||
}
|
||||
|
||||
/// Build a typed frame with explicit recursion and complete-frame output
|
||||
/// limits. High-level SDK sends use this entry point with the transport's
|
||||
/// admitted message size.
|
||||
#[wasm_bindgen]
|
||||
pub fn build_frame_with_limits(
|
||||
message_type: &str,
|
||||
data: JsValue,
|
||||
options: JsValue,
|
||||
limits: JsValue,
|
||||
) -> Result<Vec<u8>, JsValue> {
|
||||
let limits = crate::client::encode_limits_from_js(&limits)?;
|
||||
build_frame_with_encode_limits(message_type, data, options, limits)
|
||||
}
|
||||
|
||||
/// Build a typed MTP frame around a complete serialized `DataValue` payload.
|
||||
///
|
||||
/// Unlike [`build_frame`], this does not interpret the payload as a clear data
|
||||
|
|
@ -625,50 +501,19 @@ pub fn build_frame_with_payload(
|
|||
message_type: &str,
|
||||
serialized_payload: &[u8],
|
||||
options: JsValue,
|
||||
) -> Result<Vec<u8>, JsValue> {
|
||||
build_frame_with_payload_with_encode_limits(
|
||||
message_type,
|
||||
serialized_payload,
|
||||
options,
|
||||
EncodeLimits::default(),
|
||||
)
|
||||
}
|
||||
|
||||
fn build_frame_with_payload_with_encode_limits(
|
||||
message_type: &str,
|
||||
serialized_payload: &[u8],
|
||||
options: JsValue,
|
||||
limits: EncodeLimits,
|
||||
) -> Result<Vec<u8>, JsValue> {
|
||||
let comm_type = CommunicationType::from_name(message_type)
|
||||
.ok_or_else(|| js_error(format!("unknown communication type: {message_type}")))?;
|
||||
let payload = DataValue::try_from_bytes_with_limits(
|
||||
serialized_payload,
|
||||
DecodeLimits::for_transport_message_size(limits.max_output_size as u64),
|
||||
)
|
||||
.map_err(|error| decode_error(error, "invalid serialized DataValue payload"))?;
|
||||
let payload = DataValue::from_bytes(serialized_payload)
|
||||
.ok_or_else(|| js_error("invalid serialized DataValue payload"))?;
|
||||
let message =
|
||||
apply_frame_options(CommunicationValue::new(comm_type), &options)?.with_payload(payload);
|
||||
|
||||
message
|
||||
.to_bytes_with_limits(limits)
|
||||
.to_bytes()
|
||||
.map_err(|e| js_error(format!("encode failed: {e}")))
|
||||
}
|
||||
|
||||
/// Build a typed frame around a serialized payload with explicit output
|
||||
/// limits. The payload is also parsed with a policy derived from that limit so
|
||||
/// an oversized/deep input cannot bypass the bounded builder.
|
||||
#[wasm_bindgen]
|
||||
pub fn build_frame_with_payload_with_limits(
|
||||
message_type: &str,
|
||||
serialized_payload: &[u8],
|
||||
options: JsValue,
|
||||
limits: JsValue,
|
||||
) -> Result<Vec<u8>, JsValue> {
|
||||
let limits = crate::client::encode_limits_from_js(&limits)?;
|
||||
build_frame_with_payload_with_encode_limits(message_type, serialized_payload, options, limits)
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
#[cfg(target_arch = "wasm32")]
|
||||
mod tests {
|
||||
|
|
|
|||
102
wasm/src/pipe.rs
102
wasm/src/pipe.rs
|
|
@ -1,6 +1,9 @@
|
|||
use wasm_bindgen::JsCast;
|
||||
use wasm_bindgen::prelude::*;
|
||||
use wasm_bindgen_futures::JsFuture;
|
||||
|
||||
use crate::transport::{BrowserRecvStream, BrowserSendStream, log_stream_error_code};
|
||||
use crate::error::js_error;
|
||||
use crate::transport::release_writer_lock;
|
||||
|
||||
#[wasm_bindgen(typescript_custom_section)]
|
||||
const PIPE_TS: &str = r#"
|
||||
|
|
@ -20,41 +23,54 @@ export interface PipeReader {
|
|||
|
||||
#[wasm_bindgen]
|
||||
pub struct PipeWriter {
|
||||
stream: BrowserSendStream,
|
||||
writer: JsValue,
|
||||
pipe_id: u32,
|
||||
}
|
||||
|
||||
impl PipeWriter {
|
||||
pub(crate) fn new(stream: BrowserSendStream, pipe_id: u32) -> Self {
|
||||
Self { stream, pipe_id }
|
||||
}
|
||||
}
|
||||
|
||||
impl Drop for PipeWriter {
|
||||
fn drop(&mut self) {
|
||||
self.stream.release();
|
||||
pub fn new(writer: JsValue, pipe_id: u32) -> Self {
|
||||
Self { writer, pipe_id }
|
||||
}
|
||||
}
|
||||
|
||||
#[wasm_bindgen]
|
||||
impl PipeWriter {
|
||||
pub async fn write(&mut self, data: &[u8]) -> Result<(), JsValue> {
|
||||
self.stream.write_all(data).await
|
||||
let chunk = js_sys::Uint8Array::from(data);
|
||||
let write_fn = js_sys::Reflect::get(&self.writer, &JsValue::from_str("write"))
|
||||
.map_err(|_| js_error("missing write"))?
|
||||
.dyn_into::<js_sys::Function>()
|
||||
.map_err(|_| js_error("write not a function"))?;
|
||||
let write_promise = write_fn
|
||||
.call1(&self.writer, &chunk)
|
||||
.map_err(|e| js_error(format!("write failed: {:?}", e)))?;
|
||||
JsFuture::from(write_promise.unchecked_into::<js_sys::Promise>()).await?;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
pub async fn close(mut self) -> Result<(), JsValue> {
|
||||
let result = self.stream.finish().await;
|
||||
if let Err(error) = &result {
|
||||
log_stream_error_code(error, "pipe writer close");
|
||||
pub async fn close(self) -> Result<(), JsValue> {
|
||||
let close_fn = js_sys::Reflect::get(&self.writer, &JsValue::from_str("close"))
|
||||
.map_err(|_| js_error("missing close"))?
|
||||
.dyn_into::<js_sys::Function>()
|
||||
.map_err(|_| js_error("close not a function"))?;
|
||||
let close_promise = close_fn
|
||||
.call0(&self.writer)
|
||||
.map_err(|e| js_error(format!("close failed: {:?}", e)))?;
|
||||
if let Err(e) = JsFuture::from(close_promise.unchecked_into::<js_sys::Promise>()).await {
|
||||
crate::transport::log_stream_error_code(&e, "pipe writer close");
|
||||
}
|
||||
self.stream.release();
|
||||
result
|
||||
release_writer_lock(&self.writer);
|
||||
Ok(())
|
||||
}
|
||||
|
||||
pub fn abort(&mut self) -> Result<(), JsValue> {
|
||||
let result = self.stream.reset(0);
|
||||
self.stream.release();
|
||||
result
|
||||
let abort_fn = js_sys::Reflect::get(&self.writer, &JsValue::from_str("abort"))
|
||||
.map_err(|_| js_error("missing abort"))?
|
||||
.dyn_into::<js_sys::Function>()
|
||||
.map_err(|_| js_error("abort not a function"))?;
|
||||
let _ = abort_fn.call0(&self.writer);
|
||||
release_writer_lock(&self.writer);
|
||||
Ok(())
|
||||
}
|
||||
|
||||
pub fn pipe_id(&self) -> u32 {
|
||||
|
|
@ -64,26 +80,19 @@ impl PipeWriter {
|
|||
|
||||
#[wasm_bindgen]
|
||||
pub struct PipeReader {
|
||||
stream: BrowserRecvStream,
|
||||
reader: JsValue,
|
||||
description: String,
|
||||
pipe_id: u32,
|
||||
pending: Vec<u8>,
|
||||
finished: bool,
|
||||
}
|
||||
|
||||
impl PipeReader {
|
||||
pub(crate) fn new(
|
||||
stream: BrowserRecvStream,
|
||||
pipe_id: u32,
|
||||
description: String,
|
||||
pending: Vec<u8>,
|
||||
) -> Self {
|
||||
pub fn new(reader: JsValue, pipe_id: u32, description: String, pending: Vec<u8>) -> Self {
|
||||
Self {
|
||||
stream,
|
||||
reader,
|
||||
pipe_id,
|
||||
description,
|
||||
pending,
|
||||
finished: false,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
|
@ -96,18 +105,27 @@ impl PipeReader {
|
|||
return Ok(js_sys::Uint8Array::from(&data[..]).into());
|
||||
}
|
||||
|
||||
if self.finished {
|
||||
let read_fn = js_sys::Reflect::get(&self.reader, &JsValue::from_str("read"))
|
||||
.map_err(|_| js_error("missing read"))?
|
||||
.dyn_into::<js_sys::Function>()
|
||||
.map_err(|_| js_error("read not a function"))?;
|
||||
let promise = read_fn
|
||||
.call0(&self.reader)
|
||||
.map_err(|_| js_error("read call failed"))?
|
||||
.unchecked_into::<js_sys::Promise>();
|
||||
let result = JsFuture::from(promise).await?;
|
||||
|
||||
let done = js_sys::Reflect::get(&result, &JsValue::from_str("done"))
|
||||
.ok()
|
||||
.and_then(|v| v.as_bool())
|
||||
.unwrap_or(true);
|
||||
if done {
|
||||
return Ok(JsValue::NULL);
|
||||
}
|
||||
|
||||
match self.stream.read_chunk().await? {
|
||||
Some(value) => Ok(js_sys::Uint8Array::from(&value[..]).into()),
|
||||
None => {
|
||||
self.stream.release();
|
||||
self.finished = true;
|
||||
Ok(JsValue::NULL)
|
||||
}
|
||||
}
|
||||
let value = js_sys::Reflect::get(&result, &JsValue::from_str("value"))
|
||||
.map_err(|_| js_error("missing value"))?;
|
||||
Ok(js_sys::Uint8Array::new(&value).into())
|
||||
}
|
||||
|
||||
pub fn pipe_id(&self) -> u32 {
|
||||
|
|
@ -118,9 +136,3 @@ impl PipeReader {
|
|||
self.description.clone()
|
||||
}
|
||||
}
|
||||
|
||||
impl Drop for PipeReader {
|
||||
fn drop(&mut self) {
|
||||
self.stream.release();
|
||||
}
|
||||
}
|
||||
|
|
|
|||
Some files were not shown because too many files have changed in this diff Show more
Loading…
Reference in a new issue