Compare commits

..
Author SHA1 Message Date
Alex Emmet
f6a8b464e0
Merge remote-tracking branch 'refs/remotes/origin/master'
All checks were successful
CI / checks (push) Successful in 5m48s
2026-08-27 22:43:40 +02:00
Alex Emmet
9cea795d1c
ChaCha20 2026-08-27 22:42:32 +02:00
4493ed32cf
chore(deps): update chacha20 from 0.10.1 to 0.10.2
All checks were successful
CI / checks (push) Successful in 5m53s
2026-08-27 22:33:20 +02:00
c7c7afe578
fix(example): web-client type
Some checks failed
CI / checks (push) Failing after 3m16s
2026-08-27 22:16:16 +02:00
8d94bc5498
feat(qol): add direnv
Some checks failed
CI / checks (push) Failing after 3m14s
2026-08-27 21:55:36 +02:00
22d13742ae
feat(qol): remove dup
Some checks failed
CI / checks (push) Failing after 4m44s
2026-08-27 21:10:37 +02:00
df75fd2830
feat(qol): remove dup
Some checks failed
CI / checks (push) Failing after 4m19s
2026-08-27 20:28:26 +02:00
6348d7884f
Fix thing
Some checks failed
CI / checks (push) Failing after 3m59s
2026-08-27 20:17:11 +02:00
7ef6ec9e88
Merge remote-tracking branch 'refs/remotes/origin/master'
Some checks failed
CI / checks (push) Failing after 2s
2026-08-27 19:30:03 +02:00
bd5547ae6f
feat(ts-sdk): add schemas 2026-08-27 19:29:53 +02:00
Alex Emmet
c30315af94
Clean
Some checks failed
CI / checks (push) Failing after 2s
2026-08-27 16:55:52 +02:00
e83cd132a2
feat(wasm, native, h3): make wasm, native and h3 use unified interface
Some checks failed
CI / checks (push) Failing after 2s
2026-08-27 15:31:55 +02:00
Alex Emmet
101b8322a1
[Fix] Policy overwrites
Some checks failed
CI / checks (push) Failing after 2s
2026-08-20 21:59:57 +02:00
Alex Emmet
24167c4aa0
Merge remote-tracking branch 'refs/remotes/origin/master'
Some checks failed
CI / checks (push) Failing after 2s
2026-08-20 21:42:43 +02:00
Alex Emmet
420831cd09
[Add] Docs & patches 2026-08-20 21:42:23 +02:00
687d56f9f1
Merge remote-tracking branch 'refs/remotes/origin/master'
Some checks failed
CI / checks (push) Failing after 2s
2026-08-20 20:55:52 +02:00
4c10b56a6c
Update workflows 2026-08-20 20:55:43 +02:00
Alex Emmet
bd660b2afb
[Debug]
Some checks failed
CI / checks (push) Failing after 2m17s
2026-08-20 20:37:25 +02:00
Alex Emmet
2b0bdc3257
[Fix] Connections
Some checks failed
CI / checks (push) Failing after 3m14s
2026-08-20 20:11:08 +02:00
a5c8d4f0c8
Update ci.yml
Some checks failed
CI / checks (push) Failing after 2m36s
2026-08-19 13:17:05 +02:00
a6c4e56835
[Upd] Docs
Some checks failed
CI / checks (push) Failing after 2m33s
2026-08-19 12:37:22 +02:00
d11eb04d12
[Fix] Clean
Some checks failed
CI / checks (push) Failing after 2m37s
2026-08-19 11:46:40 +02:00
Alex Emmet
b331b9f6a3
[Fix] Clean
Some checks failed
CI / checks (push) Failing after 16m34s
2026-08-18 21:53:02 +02:00
Alex Emmet
3395b91ad1
Merge remote-tracking branch 'refs/remotes/origin/master'
Some checks failed
CI / checks (push) Failing after 2m20s
2026-08-18 20:59:06 +02:00
Alex Emmet
a7e804c603
[Fix] Harden MTP codec, transport, and SDK security 2026-08-18 20:58:01 +02:00
113a54048f Merge pull request 'Update Rust crate tokio-stream to v0.1.19' (#21) from renovate/tokio-stream-0.x-lockfile into master
All checks were successful
CI / checks (push) Successful in 7m47s
2026-08-14 17:00:34 +03:00
bd187106a1 Update Rust crate tokio-stream to v0.1.19
All checks were successful
renovate/stability-days Updates have met minimum release age requirement
CI / checks (pull_request) Successful in 11m23s
2026-08-14 16:00:47 +03:00
Alex Emmet
188caf56cc [Fix] Clean
All checks were successful
CI / checks (push) Successful in 4m36s
2026-08-14 14:39:09 +02:00
105 changed files with 14198 additions and 7474 deletions

1
.envrc Normal file
View file

@ -0,0 +1 @@
use flake

View file

@ -7,16 +7,12 @@ 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
@ -33,7 +29,6 @@ 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

View file

@ -14,16 +14,10 @@ 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:
@ -32,9 +26,6 @@ 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
View file

@ -6,3 +6,4 @@ dist/
*.tgz
wasm/pkg/
web_client/
.direnv

164
Cargo.lock generated
View file

@ -90,9 +90,9 @@ dependencies = [
[[package]]
name = "async-trait"
version = "0.1.91"
version = "0.1.92"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "ae36dc4177970ef04fde5178d3e2429882def40e57a451f919c098f72baa6cec"
checksum = "82f6aeea286b8eb4dd3431a1be1b59d290ace00f5bfd8e2a159bc2a05e2c1667"
dependencies = [
"proc-macro2",
"quote",
@ -221,9 +221,9 @@ checksum = "37b2a672a2cb129a2e41c10b1224bb368f9f37a2b16b612598138befd7b37eb5"
[[package]]
name = "cc"
version = "1.4.2"
version = "1.4.3"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "5d262e149917187838d5b42777c8253bcb64500067342904e7d429499a6f277e"
checksum = "509591b7bcd67f4ef775afad7662703b4935daaa6ec0e5605cfb1090b32a2b6d"
dependencies = [
"find-msvc-tools",
"jobserver",
@ -256,9 +256,9 @@ dependencies = [
[[package]]
name = "chacha20"
version = "0.10.1"
version = "0.10.2"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "d524456ba66e72eb8b115ff89e01e497f8e6d11d78b70b1aa13c0fbd97540a81"
checksum = "65c35e4b699c7e15ccbe7ee35c005e4fc0a278d22238a2857e6ce2dadeda1b06"
dependencies = [
"cfg-if",
"cpufeatures 0.3.0",
@ -596,9 +596,9 @@ checksum = "64cd1e32ddd350061ae6edb1b082d7c54915b5c672c389143b9a63403a109f24"
[[package]]
name = "find-msvc-tools"
version = "0.1.10"
version = "0.1.11"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "26b73573e6edcd2af0cdf47bd6cb58f0b3839491263c314eaad1ccf24430e1de"
checksum = "d45db016d36b838f563236e9193d0ee6ce38f3f68b6c94e914b4929c96bbb890"
[[package]]
name = "fnv"
@ -629,9 +629,9 @@ checksum = "42703706b716c37f96a77aea830392ad231f44c9e9a67872fa5548707e11b11c"
[[package]]
name = "futures"
version = "0.3.33"
version = "0.3.34"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "a88cf1f829d945f548cf8fec32c61b1f202b6d93b45848602fc02af4b12ad218"
checksum = "9a31d2a3fbaaeb2af2368bbdd904aa8e812d3c04a1ee10d3171f52d556e5d0a3"
dependencies = [
"futures-channel",
"futures-core",
@ -660,9 +660,9 @@ checksum = "92d699e522242e69e3003b94ecc1f960f3a5e015aa7c5d7486e65ad01dd94f5e"
[[package]]
name = "futures-executor"
version = "0.3.33"
version = "0.3.34"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "6754879cc9f2c66f88c6e5c35344bb0bdb0708b0352b1201815667c7eabc7458"
checksum = "031b47cf1a3c6cc8bc2fc76cd437f521619387907d469316e7c0bc278f1f5432"
dependencies = [
"futures-core",
"futures-task",
@ -764,9 +764,9 @@ dependencies = [
[[package]]
name = "h2"
version = "0.4.15"
version = "0.4.16"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "6cb093c84e8bd9b188d4c4a8cb6579fc016968d14c99882163cd3ff402a4f155"
checksum = "a9f37a958b41b3b19ee2707c06439c0e9e547e847223eb791ecb0cb821c65e27"
dependencies = [
"atomic-waker",
"bytes",
@ -895,9 +895,9 @@ dependencies = [
[[package]]
name = "http-body-util"
version = "0.1.4"
version = "0.1.5"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "e9f41fd6a08e4d4ec69df65976da761afd5ad5e58a9d4acb46bd1c953a9e3ff2"
checksum = "23169fe34a5fbcdd3f3862e78fb9b6fccd5f02a6dc6f732547005d45631ce71c"
dependencies = [
"bytes",
"futures-core",
@ -966,9 +966,9 @@ dependencies = [
[[package]]
name = "icu_collections"
version = "2.2.0"
version = "2.3.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "2984d1cd16c883d7935b9e07e44071dca8d917fd52ecc02c04d5fa0b5a3f191c"
checksum = "fa68d21081c4a05d5a901a1c62add574c77048b6a1c67be3b50ce0b60d4ca513"
dependencies = [
"displaydoc",
"potential_utf",
@ -980,9 +980,9 @@ dependencies = [
[[package]]
name = "icu_locale_core"
version = "2.2.0"
version = "2.3.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "92219b62b3e2b4d88ac5119f8904c10f8f61bf7e95b640d25ba3075e6cac2c29"
checksum = "d56e28588da92eee5c3201a6eff33fabdd49b62269c8938d4ff050ce4d900deb"
dependencies = [
"displaydoc",
"litemap",
@ -993,9 +993,9 @@ dependencies = [
[[package]]
name = "icu_normalizer"
version = "2.2.0"
version = "2.3.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "c56e5ee99d6e3d33bd91c5d85458b6005a22140021cc324cea84dd0e72cff3b4"
checksum = "12f9cf5f235641ed274641dd81c3f28d870e276763d0797aeeab72317b1c646f"
dependencies = [
"icu_collections",
"icu_normalizer_data",
@ -1007,16 +1007,17 @@ dependencies = [
[[package]]
name = "icu_normalizer_data"
version = "2.2.0"
version = "2.3.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "da3be0ae77ea334f4da67c12f149704f19f81d1adf7c51cf482943e84a2bad38"
checksum = "1563da1ed3e0b3bf3d74c9b85917ac9c56464d2f57242270c09c9e752f8021a0"
[[package]]
name = "icu_properties"
version = "2.2.0"
version = "2.3.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "bee3b67d0ea5c2cca5003417989af8996f8604e34fb9ddf96208a033901e70de"
checksum = "7e7ca276ad3145661a65914e6daf131ca5120cd3dcee8f8f3214b8875184a148"
dependencies = [
"displaydoc",
"icu_collections",
"icu_locale_core",
"icu_properties_data",
@ -1027,15 +1028,15 @@ dependencies = [
[[package]]
name = "icu_properties_data"
version = "2.2.0"
version = "2.3.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "8e2bbb201e0c04f7b4b3e14382af113e17ba4f63e2c9d2ee626b720cbce54a14"
checksum = "e590f038c1464a96894fd6d10127e90a8be4509f56ff7ecef851b15cee0b7caa"
[[package]]
name = "icu_provider"
version = "2.2.0"
version = "2.3.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "139c4cf31c8b5f33d7e199446eff9c1e02decfc2f0eec2c8d71f65befa45b421"
checksum = "92a7ed671a6aad807a8651a2e1782a6598fda9ce5185dd8158549e95a91c6428"
dependencies = [
"displaydoc",
"icu_locale_core",
@ -1153,9 +1154,9 @@ dependencies = [
[[package]]
name = "js-sys"
version = "0.3.103"
version = "0.3.104"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "53b44bfcdb3f8d5837a46dae1ca9660a837176eee74a28b229bc626816589102"
checksum = "0e0c1080212aad755ea003d18543e8768dd432c48819efd73a7bf1e39b7a5a3a"
dependencies = [
"cfg-if",
"futures-util",
@ -1173,9 +1174,9 @@ dependencies = [
[[package]]
name = "keccak"
version = "0.2.0"
version = "0.2.1"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "9e24a010dd405bd7ed803e5253182815b41bf2e6a80cc3bfc066658e03a198aa"
checksum = "ffd9697dc4a9a62e2da93389f34400b77a28f0287711263cabb203b3ccb9c0e4"
dependencies = [
"cfg-if",
"cpufeatures 0.3.0",
@ -1201,9 +1202,9 @@ checksum = "b6d2cec3eae94f9f509c767b45932f1ada8350c4bdb85af2fcab4a3c14807981"
[[package]]
name = "litemap"
version = "0.8.2"
version = "0.8.3"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "92daf443525c4cce67b150400bc2316076100ce0b3686209eb8cf3c31612e6f0"
checksum = "47d9d19d1d6efa0109d2f65ff4c85cddd50bd572e5a00127ab10987290bcefae"
[[package]]
name = "lock_api"
@ -1234,9 +1235,9 @@ checksum = "cf8baf1c55e62ffcace7a9f06f4bd9cd3f0c4beb022d3b367256b91b87513d98"
[[package]]
name = "minicov"
version = "0.3.8"
version = "0.3.9"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "4869b6a491569605d66d3952bcdf03df789e5b536e5f0cf7758a7f08a55ae24d"
checksum = "c3aa3aa12b448ac225b3102217d1ac5cc717908f02722926524b0599c933c7a0"
dependencies = [
"cc",
"walkdir",
@ -1325,9 +1326,6 @@ dependencies = [
"mtp-transport",
"mtp-type-map",
"mtp-webserver",
"rand",
"rcgen",
"tokio",
]
[[package]]
@ -1362,7 +1360,6 @@ name = "mtp-common"
version = "0.3.0"
dependencies = [
"quinn",
"rustls",
"thiserror 2.0.20",
"wtransport",
]
@ -1372,6 +1369,7 @@ name = "mtp-crypto"
version = "0.3.0"
dependencies = [
"aes-gcm",
"argon2",
"base64 0.22.1",
"chacha20poly1305",
"ed25519-dalek",
@ -1395,7 +1393,6 @@ dependencies = [
name = "mtp-files"
version = "0.3.0"
dependencies = [
"argon2",
"mtp-crypto",
"rand",
"thiserror 2.0.20",
@ -1411,6 +1408,7 @@ dependencies = [
"mtp-crypto",
"mtp-transport",
"rand",
"thiserror 2.0.20",
"tokio",
"tracing",
"wtransport",
@ -1485,7 +1483,6 @@ dependencies = [
"mtp-host",
"mtp-transport",
"quinn",
"rand",
"rcgen",
"rustls",
"thiserror 2.0.20",
@ -1665,9 +1662,9 @@ dependencies = [
[[package]]
name = "pkg-config"
version = "0.3.33"
version = "0.3.34"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "19f132c84eca552bf34cab8ec81f1c1dcc229b811638f9d283dceabe58c5569e"
checksum = "f6b464fbc74e149a392436b17d523f769e057cb6877f6a5c4618bc6f11800548"
[[package]]
name = "poly1305"
@ -1700,9 +1697,9 @@ checksum = "05c8b63e8d9609db387f0324918f81d68fe27748f084ef092fb35954d0539a85"
[[package]]
name = "potential_utf"
version = "0.1.5"
version = "0.1.6"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "0103b1cef7ec0cf76490e969665504990193874ea05c85ff9bab8b911d0a0564"
checksum = "d83eb9bc6d8e5cf568e7a1101d60ee05e81ed50ea106026f3d18deeb046d7661"
dependencies = [
"zerovec",
]
@ -1745,9 +1742,9 @@ dependencies = [
[[package]]
name = "quinn-proto"
version = "0.11.16"
version = "0.11.17"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "2f4bfc015262b9df63c8845072ce59068853ff5872180c2ce2f13038b970e560"
checksum = "04759210543be93709136e28212294a659ef5001836ff4eab4d663e4529bba83"
dependencies = [
"aws-lc-rs",
"bytes",
@ -1803,7 +1800,7 @@ version = "0.10.2"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "c7f5fa3a058cd35567ef9bfa5e75732bee0f9e4c55fa90477bef2dfcdbc4be80"
dependencies = [
"chacha20 0.10.1",
"chacha20 0.10.2",
"getrandom 0.4.3",
"rand_core 0.10.1",
]
@ -2120,7 +2117,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "09057cb2149ad4cbd2da1e26b351f9a4c354219421229c69c3063e6f61947c4a"
dependencies = [
"digest 0.11.3",
"keccak 0.2.0",
"keccak 0.2.1",
"sponge-cursor",
]
@ -2345,9 +2342,9 @@ dependencies = [
[[package]]
name = "tinystr"
version = "0.8.3"
version = "0.8.4"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "c8323304221c2a851516f22236c5722a72eaa19749016521d6dff0824447d96d"
checksum = "b1e27c91459209c2986af3dcf603a5a74a4368754ce37414f59acc971167f643"
dependencies = [
"displaydoc",
"zerovec",
@ -2408,9 +2405,9 @@ dependencies = [
[[package]]
name = "tokio-stream"
version = "0.1.18"
version = "0.1.19"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "32da49809aab5c3bc678af03902d4ccddea2a87d028d86392a4b1560c6906c70"
checksum = "a3d06f0b082ba57c26b79407372e57cf2a1e28124f78e9479fe80322cf53420b"
dependencies = [
"futures-core",
"pin-project-lite",
@ -2419,13 +2416,14 @@ dependencies = [
[[package]]
name = "tokio-util"
version = "0.7.18"
version = "0.7.19"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "9ae9cec805b01e8fc3fd2fe289f89149a9b66dd16786abd8b19cfa7b48cb0098"
checksum = "494815d09bf52b5548659851081238f0ca39ff638363907596da739561c62c52"
dependencies = [
"bytes",
"futures-core",
"futures-sink",
"libc",
"pin-project-lite",
"tokio",
]
@ -2570,9 +2568,9 @@ checksum = "ccf3ec651a847eb01de73ccad15eb7d99f80485de043efb2f370cd654f4ea44b"
[[package]]
name = "wasm-bindgen"
version = "0.2.126"
version = "0.2.127"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "4b067c0c11094aef6b7a801c1e34a26affafdf3d051dba08456b868789aaf9a4"
checksum = "1b70935747edd64d89de3efa29d73789b806c15798f8e7dca4d8ac356b50ce70"
dependencies = [
"cfg-if",
"once_cell",
@ -2583,9 +2581,9 @@ dependencies = [
[[package]]
name = "wasm-bindgen-futures"
version = "0.4.76"
version = "0.4.77"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "c62df1340f32221cb9c54d6a27b030e3dba64361d4a95bed55f9aacb44da291d"
checksum = "6b7777d5cc23d0e91404e53ce2d5e8ec7acae3026b16233dba62cd3246457950"
dependencies = [
"js-sys",
"wasm-bindgen",
@ -2593,9 +2591,9 @@ dependencies = [
[[package]]
name = "wasm-bindgen-macro"
version = "0.2.126"
version = "0.2.127"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "167ce5e579f6bcf889c4f7175a8a5a585de84e8ff93976ce393efa5f2837aab1"
checksum = "77775f8f3f7217702089053b94958f8f54061a3f663417df76e19cbdcca29bc1"
dependencies = [
"quote",
"wasm-bindgen-macro-support",
@ -2603,9 +2601,9 @@ dependencies = [
[[package]]
name = "wasm-bindgen-macro-support"
version = "0.2.126"
version = "0.2.127"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "f3997c7839262f4ef12cf90b818d6340c18e80f263f1a94bf157d0ec4420380e"
checksum = "e11d33f857dc2fb11b8bc75aee111aa9cbeb12cd9f25efd3d4c2a3dd4e235284"
dependencies = [
"bumpalo",
"proc-macro2",
@ -2616,18 +2614,18 @@ dependencies = [
[[package]]
name = "wasm-bindgen-shared"
version = "0.2.126"
version = "0.2.127"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "dc1b4cb0cc549fcf58d7dfc081778139b3d283a081644e833e84682ad71cea24"
checksum = "7ef64dbcc55df09c7e5a46182d181c2cfa3e925f3da937ea764728b4bbb9dcbf"
dependencies = [
"unicode-ident",
]
[[package]]
name = "wasm-bindgen-test"
version = "0.3.76"
version = "0.3.77"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "2a0d555ca874445df8d314f94f5c948a4e74e5418f332c89f660a3d8310a96f4"
checksum = "895a2607575412a4eda1df892084a375ea10dfeadc4d7d2ab87b854e4ddc7ba1"
dependencies = [
"async-trait",
"cast",
@ -2647,9 +2645,9 @@ dependencies = [
[[package]]
name = "wasm-bindgen-test-macro"
version = "0.3.76"
version = "0.3.77"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "94eb68555b95bcea5e8cf4abe280b529049479fa995bfc23734af96a6aedc120"
checksum = "4288cb0ebe215033bf949ae1fd046726daa4c32a157f24b9dc6ac387a52aa759"
dependencies = [
"proc-macro2",
"quote",
@ -2658,9 +2656,9 @@ dependencies = [
[[package]]
name = "wasm-bindgen-test-shared"
version = "0.2.126"
version = "0.2.127"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "c31d56021e873866c968588ed85ccdf56db5c426e44afdb4618c39895104b920"
checksum = "33ff1c1b360982e93b6d8ea9c04836f71dba0817a16f91e229cf3a51bdd9d987"
[[package]]
name = "wasm-tracing"
@ -2791,9 +2789,9 @@ checksum = "589f6da84c646204747d1270a2a5661ea66ed1cced2631d546fdfb155959f9ec"
[[package]]
name = "writeable"
version = "0.6.3"
version = "0.6.4"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "1ffae5123b2d3fc086436f8834ae3ab053a283cfac8fe0a0b8eaae044768a4c4"
checksum = "3ad82d2a33cdc9674dc7465672f271e096168fcdbe0f799d9e6db8c5892679dc"
[[package]]
name = "wtransport"
@ -2938,9 +2936,9 @@ dependencies = [
[[package]]
name = "zerotrie"
version = "0.2.4"
version = "0.2.5"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "0f9152d31db0792fa83f70fb2f83148effb5c1f5b8c7686c3459e361d9bc20bf"
checksum = "4ea269c3bd32f0a32c321907a2ae912ba6f4649bb0fc764a15627e99a7095a3f"
dependencies = [
"displaydoc",
"yoke",
@ -2949,9 +2947,9 @@ dependencies = [
[[package]]
name = "zerovec"
version = "0.11.6"
version = "0.11.7"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "90f911cbc359ab6af17377d242225f4d75119aec87ea711a880987b18cd7b239"
checksum = "94b5c6b5976d66c1d703c4fd17d3f5e43c8cedaacf604961b171adc7130896d8"
dependencies = [
"yoke",
"zerofrom",
@ -2960,13 +2958,13 @@ dependencies = [
[[package]]
name = "zerovec-derive"
version = "0.11.3"
version = "0.11.5"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "625dc425cab0dca6dc3c3319506e6593dcb08a9f387ea3b284dbd52a92c40555"
checksum = "9f212a141d820099d57ffafb9569be9617a6f27d3dc881fbee8fb56642f917a9"
dependencies = [
"proc-macro2",
"quote",
"syn 2.0.119",
"syn 3.0.3",
]
[[package]]

View file

@ -113,10 +113,5 @@ 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"]

View file

@ -47,18 +47,29 @@ Feature summary:
| Feature | Pulls in | Enables |
| --- | --- | --- |
| `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 |
| `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` |
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)
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)
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::transport`, `mtp::common`, `mtp::type_map`, `mtp::crypto`, `mtp::host`, and `mtp::client`.
`mtp::codec`, `mtp::common`, `mtp::type_map`, `mtp::crypto`, `mtp::host`,
`mtp::client`, `mtp::files`, and `mtp::webserver` when their features are enabled.
### Codec

View file

@ -13,6 +13,10 @@ 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 {
@ -134,17 +138,33 @@ 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().await;
let mut pending = self
.pipe_dispatcher
.pending_creations
.lock()
.map_err(|_| mtp_common::PipeError::ConnectionClosed)?;
let pipe_id = loop {
let candidate = rand::random::<u32>();
if candidate != 0 && !pending.contains_key(&candidate) {
if candidate != 0
&& !pending.contains_key(&candidate)
&& !is_expired_creation(&self.pipe_dispatcher, candidate)
{
break candidate;
}
};
pending.insert(pipe_id, tx);
pending.insert(
pipe_id,
PendingCreation {
token: token.clone(),
sender: 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,
@ -154,19 +174,17 @@ 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,
})
}
@ -207,18 +225,19 @@ 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>>(
config.policy.receiver_queue_capacity,
receiver_queue_capacity,
);
let (pipe_req_tx, pipe_req_rx) =
mpsc::channel::<PipeRequest>(config.policy.receiver_queue_capacity);
let (pipe_req_tx, pipe_req_rx) = mpsc::channel::<PipeRequest>(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: Mutex::new(std::collections::HashMap::new()),
pending_creations: std::sync::Mutex::new(std::collections::HashMap::new()),
expired_creations: std::sync::Mutex::new(std::collections::HashMap::new()),
pending_pipes: Mutex::new(std::collections::HashMap::new()),
policy: Arc::new(config.policy),
});
@ -255,8 +274,9 @@ 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>>(
config.policy.receiver_queue_capacity,
receiver_queue_capacity,
);
let dispatcher = Arc::new(PipeDispatcher {
pending_requests: Mutex::new(std::collections::HashMap::new()),

View file

@ -222,6 +222,10 @@ 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)
@ -233,10 +237,7 @@ 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(keys.public_key_bundle().as_bytes()),
);
.add_typed_default(DataType::PublicKeys, DataValue::Bytes(public_key_bytes));
if let Some(desc) = &config.description {
ident = ident.add_typed_default(DataType::Description, DataValue::Str(desc.clone()));
}
@ -415,7 +416,9 @@ 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.as_bytes();
let pk_bytes = pk_bundle
.try_as_bytes()
.map_err(|error| CommunicationError::ParseError(error.to_string()))?;
let mut register = CommunicationValue::new_with_type_map(CommunicationType::Register, &tm)
.add_typed_default(DataType::Version, DataValue::Str(version_str.clone()))
@ -601,7 +604,9 @@ mod tests {
#[cfg(feature = "pipes")]
type_map: mtp_codec::TypeMap::latest(),
#[cfg(feature = "pipes")]
pending_creations: Mutex::new(HashMap::new()),
pending_creations: std::sync::Mutex::new(HashMap::new()),
#[cfg(feature = "pipes")]
expired_creations: std::sync::Mutex::new(HashMap::new()),
#[cfg(feature = "pipes")]
pending_pipes: Mutex::new(HashMap::new()),
#[cfg(feature = "pipes")]
@ -639,7 +644,9 @@ mod tests {
#[cfg(feature = "pipes")]
type_map: mtp_codec::TypeMap::latest(),
#[cfg(feature = "pipes")]
pending_creations: Mutex::new(HashMap::new()),
pending_creations: std::sync::Mutex::new(HashMap::new()),
#[cfg(feature = "pipes")]
expired_creations: std::sync::Mutex::new(HashMap::new()),
#[cfg(feature = "pipes")]
pending_pipes: Mutex::new(HashMap::new()),
#[cfg(feature = "pipes")]

View file

@ -66,8 +66,8 @@ pub(crate) async fn start_ping_session(
return None;
}
let (pong_tx, mut pong_rx) = mpsc::unbounded_channel();
receiver.observe_pongs(pong_tx).await;
let (pong_tx, mut pong_rx) = mpsc::channel(1);
receiver.observe_pongs_bounded(pong_tx).await;
let last_ping = Arc::new(Mutex::new(None));
let ping_state = last_ping.clone();
let interval = config.ping_interval;
@ -75,6 +75,7 @@ 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 {
@ -91,6 +92,7 @@ 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;
@ -121,7 +123,9 @@ 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;
}

View file

@ -5,6 +5,8 @@ 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};
@ -21,6 +23,8 @@ 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")]
@ -33,9 +37,11 @@ impl PipeHandle {
&self.description
}
pub async fn wait(self) -> Result<Option<mtp_transport::PipeWriter>, PipeError> {
match self.response_rx.await {
Ok(Ok(true)) => {
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))) => {
let writer = self
.sender
.open_pipe(self.pipe_id, &self.description)
@ -43,10 +49,27 @@ impl PipeHandle {
.map_err(PipeError::from)?;
Ok(Some(writer))
}
Ok(Ok(false)) => Ok(None),
Ok(Err(e)) => Err(e),
Err(_) => Err(PipeError::StreamClosed),
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)
}
}
}
}
#[cfg(feature = "pipes")]
impl Drop for PipeHandle {
fn drop(&mut self) {
expire_pending_creation(&self.dispatcher, self.pipe_id, &self.token);
}
}
@ -55,9 +78,41 @@ 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 {
@ -69,6 +124,10 @@ 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;
@ -92,7 +151,10 @@ impl PipeRequest {
let timeout = self.dispatcher.policy.read_timeout;
match tokio::time::timeout(timeout, pipe_rx).await {
Ok(Ok(reader)) => Ok(reader),
Ok(Ok(reader)) => {
expected_pipe.disarm();
Ok(reader)
}
Ok(Err(_)) => {
self.dispatcher
.pending_pipes
@ -129,14 +191,54 @@ 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:
Mutex<HashMap<u32, tokio::sync::oneshot::Sender<Result<bool, PipeError>>>>,
pub(crate) pending_creations: StdMutex<HashMap<u32, PendingCreation>>,
#[cfg(feature = "pipes")]
pub(crate) expired_creations: StdMutex<HashMap<u32, Instant>>,
#[cfg(feature = "pipes")]
pub(crate) pending_pipes:
Mutex<HashMap<u32, tokio::sync::oneshot::Sender<mtp_transport::PipeReader>>>,
@ -144,6 +246,90 @@ 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>>,
@ -200,15 +386,14 @@ pub(crate) async fn expire_pending_request(
let mut expired = dispatcher.expired_requests.lock().await;
let now = Instant::now();
expired.retain(|_, expires_at| *expires_at > now);
if expired.len() >= MAX_EXPIRED_REQUEST_TOMBSTONES {
if let Some(oldest) = expired
if expired.len() >= MAX_EXPIRED_REQUEST_TOMBSTONES
&& let Some(oldest) = expired
.iter()
.min_by_key(|(_, expires_at)| **expires_at)
.map(|(id, _)| *id)
{
expired.remove(&oldest);
}
}
expired.insert(request_id, now + EXPIRED_REQUEST_TOMBSTONE_TTL);
}
}
@ -267,6 +452,7 @@ 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;
@ -284,9 +470,15 @@ pub(crate) async fn run_dispatcher(
continue;
};
let accepted = msg.get_bool(DataType::Accepted).unwrap_or(false);
let mut pending = dispatcher.pending_creations.lock().await;
if let Some(tx) = pending.remove(&pipe_id) {
let _ = tx.send(Ok(accepted));
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);
}
continue;
}
@ -304,6 +496,10 @@ 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;
}
@ -326,6 +522,10 @@ 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
View file

@ -187,9 +187,9 @@ dependencies = [
[[package]]
name = "chacha20"
version = "0.10.1"
version = "0.10.2"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "d524456ba66e72eb8b115ff89e01e497f8e6d11d78b70b1aa13c0fbd97540a81"
checksum = "65c35e4b699c7e15ccbe7ee35c005e4fc0a278d22238a2857e6ce2dadeda1b06"
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.1",
"chacha20 0.10.2",
"getrandom 0.4.3",
"rand_core 0.10.1",
]

View file

@ -2,7 +2,7 @@ use byteorder::{BigEndian, ReadBytesExt, WriteBytesExt};
use std::fmt;
use std::io::Cursor;
use crate::data_value::{DataKind, DataValue, DecodeLimits};
use crate::data_value::{DataKind, DataValue, DecodeError, DecodeLimits, EncodeLimits};
use crate::rand_u32;
use mtp_common::CodecError;
use mtp_type_map::{
@ -260,23 +260,50 @@ impl CommunicationValue {
#[must_use]
pub fn reply_to(&self, comm_type: CommunicationType) -> Self {
let mut response = Self::new(comm_type);
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);
response.sender = self.receiver;
response.receiver = self.sender;
response
}
pub fn merge(&mut self, other: &Self) {
if self.mapping_error.is_none() {
self.mapping_error.clone_from(&other.mapping_error);
/// 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 Some(other_entries) = other.payload.container_entries() else {
self.mapping_error
.get_or_insert(CodecError::InvalidEncoding);
return;
};
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 {
let _ = self.insert_data(*id, value.clone());
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);
}
}
@ -315,9 +342,22 @@ 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)?;
@ -344,44 +384,71 @@ impl CommunicationValue {
body.write_u64::<BigEndian>(receiver)
.map_err(|_| CodecError::InvalidEncoding)?;
}
body.extend_from_slice(&self.payload.to_bytes()?);
body.extend_from_slice(&payload);
let length = u32::try_from(body.len()).map_err(|_| CodecError::TooManyEntries)?;
let mut out = Vec::with_capacity(4 + body.len());
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);
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(|_| CodecError::InvalidEncoding)? as usize;
.map_err(|_| DecodeError::MalformedEncoding)? as usize;
let end = 4usize
.checked_add(length)
.ok_or(CodecError::InvalidEncoding)?;
.ok_or(DecodeError::MalformedEncoding)?;
if end != bytes.len() {
return Err(CodecError::InvalidEncoding);
return Err(DecodeError::MalformedEncoding);
}
let comm_type = CommunicationTypeId(
cursor
.read_u16::<BigEndian>()
.map_err(|_| CodecError::InvalidEncoding)?,
.map_err(|_| DecodeError::MalformedEncoding)?,
);
let flags = cursor.read_u8().map_err(|_| CodecError::InvalidEncoding)?;
let flags = cursor
.read_u8()
.map_err(|_| DecodeError::MalformedEncoding)?;
if flags & !FLAG_KNOWN != 0 {
return Err(CodecError::InvalidEncoding);
return Err(DecodeError::MalformedEncoding);
}
let id = if flags & FLAG_HAS_ID != 0 {
Some(
cursor
.read_u32::<BigEndian>()
.map_err(|_| CodecError::InvalidEncoding)?,
.map_err(|_| DecodeError::MalformedEncoding)?,
)
} else {
None
@ -390,7 +457,7 @@ impl CommunicationValue {
Some(
cursor
.read_u64::<BigEndian>()
.map_err(|_| CodecError::InvalidEncoding)?,
.map_err(|_| DecodeError::MalformedEncoding)?,
)
} else {
None
@ -399,14 +466,14 @@ impl CommunicationValue {
Some(
cursor
.read_u64::<BigEndian>()
.map_err(|_| CodecError::InvalidEncoding)?,
.map_err(|_| DecodeError::MalformedEncoding)?,
)
} else {
None
};
let payload = DataValue::read_from_with_limits(&mut cursor, limits)?;
let payload = DataValue::read_from_with_diagnostics(&mut cursor, limits)?;
if cursor.position() as usize != end {
return Err(CodecError::InvalidEncoding);
return Err(DecodeError::MalformedEncoding);
}
Ok(Self {
id,
@ -420,13 +487,36 @@ impl CommunicationValue {
}
pub fn from_bytes_with(bytes: &[u8], type_map: &TypeMap) -> Result<Self, CodecError> {
let mut value = Self::from_bytes(bytes)?;
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)?;
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());
}
@ -441,7 +531,8 @@ impl CommunicationValue {
.comm_id_enum(comm)
.ok_or_else(|| CodecError::UnknownCommunicationType(comm_name.to_string()))?,
);
let payload = migrate_data_value(&self.payload, source, target)?;
let mut context = MigrationContext::new(limits);
let payload = migrate_data_value(&self.payload, source, target, &mut context)?;
Ok(Self {
id: self.id,
comm_type,
@ -454,15 +545,63 @@ 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) => {
let mut migrated = Vec::with_capacity(entries.len());
context.enter()?;
let count = u16::try_from(entries.len()).map_err(|_| CodecError::TooManyEntries)?;
let mut migrated = Vec::with_capacity(usize::from(count));
for (old_id, value) in entries {
let name = source
.data_type_name(old_id.0)
@ -474,16 +613,20 @@ 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)?));
migrated.push((new_id, migrate_data_value(value, source, target, context)?));
}
context.leave();
Ok(DataValue::Container(migrated))
}
DataValue::Array(values) => Ok(DataValue::Array(
values
.iter()
.map(|value| migrate_data_value(value, source, target))
.collect::<Result<Vec<_>, _>>()?,
)),
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))
}
#[cfg(feature = "crypto")]
DataValue::Signed(_) | DataValue::Encrypted(_) => Err(CodecError::InvalidEncoding),
scalar => Ok(scalar.clone()),
@ -641,6 +784,39 @@ 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![
@ -740,9 +916,19 @@ mod tests {
let mut signer_public_keys = recipient.public_key_bundle();
signer_public_keys.sig_cl_public_key = signer_public_key;
signed.verify(SENDER_ID, &signer_public_keys, ProtectionPurpose::from(1))?;
signed.verify_with_policy(
SENDER_ID,
&signer_public_keys,
ProtectionPurpose::from(1),
crate::ProtectionPolicy::any_supported(),
)?;
assert_eq!(
signed.into_verified(SENDER_ID, &signer_public_keys, ProtectionPurpose::from(1))?,
signed.into_verified_with_policy(
SENDER_ID,
&signer_public_keys,
ProtectionPurpose::from(1),
crate::ProtectionPolicy::any_supported(),
)?,
clear_payload
);
Ok(())

File diff suppressed because it is too large Load diff

View file

@ -11,20 +11,33 @@ pub use data_value::{
ApplicationProtectionPurpose, EncryptedValue, MtpProtectionPurpose, ProtectionError,
ProtectionPolicy, ProtectionPurpose, ProtectionPurposeError, SignaturePolicy, SignedValue,
};
pub use data_value::{DataKind, DataValue, DecodeLimits};
pub use data_value::{
DEFAULT_TRANSPORT_ALLOCATION_FACTOR, DataKind, DataValue, DecodeError, DecodeLimits,
EncodeLimits,
};
pub use mtp_common::{CodecError, TimeError, unix_time_millis};
#[cfg(feature = "crypto")]
#[allow(deprecated)]
pub use protected::{
CURRENT_PROTECTED_VERSION, InMemoryReplayGuard, ProtectedError, ProtectedMessageBuilder,
ReplayError, ReplayGuard, VerifiedProtectedMessage, open_protected, open_protected_with,
open_protected_with_keys, protected_claimed_signer_id,
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,
};
#[cfg(feature = "crypto")]
#[allow(deprecated)]
pub use relay::{
CURRENT_RELAY_VERSION, RelayError, SealedRelayBuilder, VerifiedRelayContent,
CURRENT_RELAY_VERSION, RelayError, RelayOpenOptions, SealedRelayBuilder, VerifiedRelayContent,
VerifiedRelayMetadata, forward_relay_frame, open_relay_content,
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,
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,
};
pub use mtp_type_map::{

File diff suppressed because it is too large Load diff

View file

@ -2,6 +2,7 @@ use mtp_common::CodecError;
use mtp_type_map::{PROTOCOL_VERSION, TypeMap, Version};
use crate::CommunicationValue;
use crate::EncodeLimits;
pub use mtp_type_map::Registry;
@ -42,7 +43,41 @@ impl VersionedCodec {
/// Encode a value using the codec's negotiated framing rules.
pub fn encode(&self, value: &CommunicationValue) -> Result<Vec<u8>, CodecError> {
value.to_bytes()
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)
}
/// Decode a frame and retain the negotiated type map for typed access.
@ -58,3 +93,34 @@ 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
View file

@ -139,9 +139,9 @@ checksum = "f079e83a288787bcd14a6aea84cee5c87a67c5a3e660c30f557a3d24761b3527"
[[package]]
name = "chacha20"
version = "0.10.1"
version = "0.10.2"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "d524456ba66e72eb8b115ff89e01e497f8e6d11d78b70b1aa13c0fbd97540a81"
checksum = "65c35e4b699c7e15ccbe7ee35c005e4fc0a278d22238a2857e6ce2dadeda1b06"
dependencies = [
"cfg-if",
"cpufeatures",

View file

@ -15,7 +15,6 @@ 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",

View file

@ -41,6 +41,10 @@ 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}")]
@ -160,6 +164,9 @@ 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),
@ -178,6 +185,38 @@ 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 {
@ -208,6 +247,7 @@ 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"))]

243
create-web-release.mjs Normal file
View file

@ -0,0 +1,243 @@
#!/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
View file

@ -181,9 +181,9 @@ dependencies = [
[[package]]
name = "chacha20"
version = "0.10.1"
version = "0.10.2"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "d524456ba66e72eb8b115ff89e01e497f8e6d11d78b70b1aa13c0fbd97540a81"
checksum = "65c35e4b699c7e15ccbe7ee35c005e4fc0a278d22238a2857e6ce2dadeda1b06"
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.1",
"chacha20 0.10.2",
"getrandom 0.4.3",
"rand_core 0.10.1",
]

View file

@ -23,6 +23,7 @@ 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 }
@ -43,3 +44,4 @@ hkdf = ["dep:hkdf", "dep:sha2"]
sha2 = ["dep:sha2"]
tls = ["dep:rcgen", "dep:time"]
parallel = ["dep:tokio"]
password-kdf = ["dep:argon2"]

View file

@ -6,6 +6,8 @@ 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")]

View file

@ -40,6 +40,114 @@ 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> {
@ -72,50 +180,7 @@ impl MultiEncryptedMessage {
/// Parse the canonical envelope body.
pub fn from_bytes(bytes: &[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() + 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(),
})
Ok(MultiEncryptedMessageRef::from_bytes(bytes)?.to_owned())
}
}
@ -197,20 +262,51 @@ pub fn decrypt_multi_for(
purpose: u8,
keyring: &Keyring,
) -> Result<Vec<u8>, CryptoError> {
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()
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()
})
{
return Err(CryptoError::MalformedEnvelope);
}
let payload_aad = payload_aad(message)?;
for entry in &message.recipients {
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 shared_secret =
match HybridKem::decapsulate(&keyring.kem_secret_key, &entry.kem_ciphertext) {
Ok(secret) => secret,
@ -219,30 +315,51 @@ pub fn decrypt_multi_for(
let wrap_key = Zeroizing::new(derive_encryption_key(
&shared_secret,
KEY_WRAP_DOMAIN,
&[message.encryption_type.to_byte(), purpose],
&[encryption_type.to_byte(), purpose],
)?);
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,
) {
let aad = wrap_aad(encryption_type, purpose, &entry.kem_ciphertext);
let cek = match open_with_key(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(
message.encryption_type,
cek,
&message.ciphertext,
&payload_aad,
);
return open_with_key(encryption_type, cek, 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::*;
@ -325,4 +442,30 @@ 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(())
}
}

View file

@ -36,3 +36,29 @@ 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)
}

View file

@ -340,9 +340,9 @@ impl Keyring {
Ok(out)
}
pub fn to_bytes(&self) -> Zeroizing<Vec<u8>> {
#[deprecated(note = "use try_to_bytes for the primary fallible serializer")]
pub fn to_bytes(&self) -> Result<Zeroizing<Vec<u8>>, crate::error::CryptoError> {
self.try_to_bytes()
.expect("key material length exceeds wire limit")
}
pub fn from_bytes(bytes: &[u8]) -> Result<Self, crate::error::CryptoError> {
@ -383,16 +383,26 @@ impl Keyring {
})
}
pub fn to_hex(&self) -> String {
bytes_to_hex(&self.to_bytes())
#[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 from_hex(s: &str) -> Result<Self, crate::error::CryptoError> {
Self::from_bytes(&hex_to_bytes(s)?)
}
pub fn to_base64(&self) -> String {
bytes_to_base64(&self.to_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_to_bytes()?))
}
pub fn from_base64(s: &str) -> Result<Self, crate::error::CryptoError> {
@ -507,9 +517,9 @@ impl PublicKeyBundle {
Ok(out)
}
pub fn as_bytes(&self) -> Vec<u8> {
#[deprecated(note = "use try_as_bytes for the primary fallible serializer")]
pub fn as_bytes(&self) -> Result<Vec<u8>, crate::error::CryptoError> {
self.try_as_bytes()
.expect("public key bundle field exceeds wire limit")
}
/// Parse a complete suite-compatible public bundle.
@ -582,8 +592,13 @@ impl PublicKeyBundle {
Self::from_bytes(bytes)
}
pub fn to_base64(&self) -> String {
bytes_to_base64(&self.as_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 from_base64(s: &str) -> Result<Self, crate::error::CryptoError> {
@ -602,12 +617,6 @@ 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")
@ -629,7 +638,7 @@ mod tests {
let cl = SignaturePublicKey::new(vec![3u8; 32]);
let bundle = PublicKeyBundle::new(kem, pq, cl);
let bytes = bundle.as_bytes();
let bytes = bundle.try_as_bytes()?;
let recovered = PublicKeyBundle::from_bytes_unvalidated(&bytes)?;
assert_eq!(
@ -672,9 +681,9 @@ mod tests {
SignaturePqPublicKey::new(vec![0xCDu8; 96]),
SignaturePublicKey::new(vec![0xEFu8; 32]),
);
let bytes: Vec<u8> = Vec::from(&bundle);
let bytes = bundle.try_as_bytes()?;
let recovered = PublicKeyBundle::from_bytes_unvalidated(bytes.as_slice())?;
assert_eq!(bundle.as_bytes(), recovered.as_bytes());
assert_eq!(bundle.try_as_bytes()?, recovered.try_as_bytes()?);
Ok(())
}
@ -688,8 +697,8 @@ mod tests {
SignaturePublicKey::new(vec![5u8; 32]),
SignaturePrivateKey::new(vec![6u8; 32]),
);
let bytes = keyring.to_bytes();
let recovered = Keyring::from_bytes(&bytes)?;
let bytes = keyring.try_to_bytes()?;
let recovered = Keyring::from_bytes(bytes.as_slice())?;
assert_eq!(
keyring.kem_public_key.as_bytes(),
recovered.kem_public_key.as_bytes()
@ -722,7 +731,7 @@ mod tests {
}
#[test]
fn canonical_key_parsers_reject_trailing_bytes() {
fn canonical_key_parsers_reject_trailing_bytes() -> Result<(), Box<dyn std::error::Error>> {
let keyring = Keyring::new(
KemPublicKey::new(vec![1u8; 16]),
KemPrivateKey::new(vec![2u8; 16]),
@ -731,25 +740,40 @@ mod tests {
SignaturePublicKey::new(vec![5u8; 16]),
SignaturePrivateKey::new(vec![6u8; 16]),
);
let mut keyring_bytes = keyring.to_bytes().to_vec();
let mut keyring_bytes = keyring.try_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.as_bytes();
let mut bundle_bytes = bundle.try_as_bytes()?;
bundle_bytes.push(0xBB);
assert!(PublicKeyBundle::from_bytes(&bundle_bytes).is_err());
Ok(())
}
#[test]
fn validated_bundle_rejects_partial_suite_keys() {
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>> {
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.as_bytes()).is_err());
assert!(PublicKeyBundle::from_bytes_validated(&bundle.try_as_bytes()?).is_err());
Ok(())
}
#[test]
@ -762,9 +786,9 @@ mod tests {
SignaturePublicKey::new(vec![4u8; 16]),
SignaturePrivateKey::new(vec![5u8; 16]),
);
let bytes = keyring.to_bytes();
let bytes = keyring.try_to_bytes()?;
let recovered = Keyring::try_from(bytes.as_slice())?;
assert_eq!(keyring.to_bytes(), recovered.to_bytes());
assert_eq!(keyring.try_to_bytes()?, recovered.try_to_bytes()?);
Ok(())
}
@ -788,9 +812,9 @@ mod tests {
SignaturePublicKey::new(vec![5u8; 16]),
SignaturePrivateKey::new(vec![6u8; 16]),
);
let hex = keyring.to_hex();
let hex = keyring.try_to_hex()?;
let recovered = Keyring::from_hex(&hex)?;
assert_eq!(keyring.to_bytes(), recovered.to_bytes());
assert_eq!(keyring.try_to_bytes()?, recovered.try_to_bytes()?);
Ok(())
}
@ -804,9 +828,9 @@ mod tests {
SignaturePublicKey::new(vec![5u8; 16]),
SignaturePrivateKey::new(vec![6u8; 16]),
);
let b64 = keyring.to_base64();
let b64 = keyring.try_to_base64()?;
let recovered = Keyring::from_base64(&b64)?;
assert_eq!(keyring.to_bytes(), recovered.to_bytes());
assert_eq!(keyring.try_to_bytes()?, recovered.try_to_bytes()?);
Ok(())
}
@ -817,9 +841,9 @@ mod tests {
SignaturePqPublicKey::new(vec![2u8; 64]),
SignaturePublicKey::new(vec![3u8; 32]),
);
let b64 = bundle.to_base64();
let b64 = bundle.try_to_base64()?;
let recovered = PublicKeyBundle::from_base64_unvalidated(&b64)?;
assert_eq!(bundle.as_bytes(), recovered.as_bytes());
assert_eq!(bundle.try_as_bytes()?, recovered.try_as_bytes()?);
Ok(())
}

View file

@ -60,6 +60,8 @@ 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};
@ -83,7 +85,9 @@ pub use helper::{ENCRYPT_DOMAIN, KEY_WRAP_DOMAIN};
#[cfg(all(feature = "mlkem-tls", feature = "hkdf"))]
pub use helper::{
MAX_RECIPIENTS, MultiEncryptedMessage, RecipientEntry, decrypt_multi_for, encrypt_multi_for,
MAX_RECIPIENTS, MultiEncryptedMessage, MultiEncryptedMessageRef, RecipientEntry,
decrypt_multi_for, decrypt_multi_for_parts, decrypt_multi_for_parts_with_limit,
encrypt_multi_for,
};
/* ================================ TESTS ================================ */
@ -311,7 +315,9 @@ mod tests {
#[test]
fn keyring_serialize_roundtrip() {
let kr = Keyring::generate();
let bytes = kr.to_bytes();
let bytes = kr
.try_to_bytes()
.expect("keyring serialization should succeed");
let loaded = Keyring::from_bytes(&bytes).expect("keyring roundtrip should succeed");
assert_eq!(
kr.kem_public_key.as_bytes(),
@ -332,7 +338,9 @@ mod tests {
fn public_key_bundle_serialize_roundtrip() {
let kr = Keyring::generate();
let bundle = kr.public_key_bundle();
let bytes = bundle.as_bytes();
let bytes = bundle
.try_as_bytes()
.expect("bundle serialization should succeed");
let loaded = PublicKeyBundle::from_bytes(&bytes).expect("bundle roundtrip should succeed");
assert_eq!(
bundle.kem_public_key.as_bytes(),

View file

@ -8,9 +8,23 @@ ignore = []
[bans]
# Flag multiple versions of the same crate so duplicate trees are visible.
multiple-versions = "warn"
multiple-versions = "deny"
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 = [

View file

@ -1,19 +1,23 @@
# MTP Connections
Native clients and hosts share the same connection shape after the opening handshake. The client creates the connection; the host receives it from `accept()`.
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()`.
| 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` | Receives application frames | Receives application frames | Receives application frames |
| `receiver` | Underlying receiver; use `receive()` for application frames | Underlying receiver; use `receive()` for application frames | Underlying receiver; use `receive()` for 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` |
| `request_path` | / | / | WebTransport CONNECT path (e.g. `/mtp`) |
| `path` | — | Native hosts use `/` | WebTransport CONNECT path (e.g. `/mtp`) |
| `remote_addr` | Server `SocketAddr` when available | Peer `SocketAddr` | Peer `SocketAddr` |
`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.
`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.
Server-side MTP connections expose `remote_addr`, the peer address observed by
QUIC. HTTP route handlers receive the peer address as `HttpRequest::remote_addr`.

View file

@ -4,23 +4,28 @@ 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):
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.
```rust
use mtp::codec::registry::Registry;
use mtp::codec::{Version, registry::Registry};
let registry = Registry::builtin(); // loads all TypeMaps from config
let registry = Registry::builtin(); // loads all TypeMaps from the build config
// Check if a version is supported
assert!(registry.supports(&Version(1, 0)));
assert!(registry.supports(&Version(3, 0)));
// Find highest mutual version for a client
let client_versions = &[Version(0, 0), Version(1, 0)];
let client_versions = &[Version(2, 0), Version(3, 0)];
let negotiated = registry.negotiate(client_versions);
assert_eq!(negotiated, Some(Version(1, 0)));
assert_eq!(negotiated, Some(Version(3, 0)));
// Look up a version's TypeMap
let tm = registry.get(&Version(2, 0)).unwrap();
let tm = registry.get(&Version(3, 0)).unwrap();
```
The `Registry::builtin()` constructor uses the `TypeMap::vX_Y()` methods generated from the config.
@ -54,9 +59,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.receiver for raw CommunicationValue I/O
// conn.sender / conn.receive() for application CommunicationValue I/O
let msg = conn.receiver.receive().await?;
let msg = conn.receive().await?;
}
```
@ -94,29 +99,31 @@ The client's `PROTOCOL_VERSION` constant is set by `protocol_version` in `type-m
## Version Negotiation Flow
```
Client (v2.0) Host (v0.0, v1.0, v2.0)
Client (v3.0) Host (v3.0)
| |
| QUIC connect |
|----------------------->|
| |
| CommValue{ Ident. } |
| Version -> "2.0" |
| Version -> "3.0" |
| Id -> 8765 |
| (unsigned hello; auth |
| challenge follows) |
|----------------------->|
| | registry.negotiate(&[Version(2,0)])
| | -> Some(Version(2,0))
| | registry.negotiate(&[Version(3,0)])
| | -> Some(Version(3,0))
| |
| Response | selected v2.0 TypeMap
| Response | selected v3.0 TypeMap
|<-----------------------|
| Status, version |
| |
| subsequent messages |
| use v2.0 TypeMap |
| use v3.0 TypeMap |
```
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.
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.
## Protocol Ping and Pong

View file

@ -12,6 +12,8 @@ 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. |
@ -43,4 +45,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. 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; a handshake that exceeds the configured limit returns `AcceptError::AuthenticationTimedOut`. The authentication flow and its signed fields are defined in [Security](SECURITY.md).

View file

@ -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();
conn.sender.close().await;
```
## 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.to_bytes();
let keyring_bytes = keyring.try_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.to_bytes()` -> `Vec<u8>`
- Serialise: `keyring.try_to_bytes()` -> `Result<Zeroizing<Vec<u8>>, CryptoError>`
- 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::transport::Policy`):
Two send modes (configured via `mtp::client::Policy`):
- `PersistentStream` (default): reuses one QUIC unidirectional stream
- `SingleStreamPerMessage`: opens a new stream per message
@ -248,12 +248,16 @@ Inbound frames are queued internally. The `receive()` method returns the next av
### Close
```rust
conn.sender.close();
conn.sender.close().await;
// or
conn.receiver.close();
```
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.
`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.
### Pipes
@ -309,7 +313,7 @@ For `public_signer`, call `verify` and `into_verified` before calling `decrypt`;
The `Policy` struct controls transport behaviour:
```rust
use mtp::transport::{Policy, SendMode};
use mtp::client::{Policy, SendMode};
let policy = Policy {
send_mode: SendMode::PersistentStream,

View file

@ -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, request 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, `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()` 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.
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.
### Authentication
@ -164,12 +164,18 @@ On success, the connection has `AuthState::Authenticated`, the assigned `client_
## Errors
`MTPWebServer::new` returns `CommunicationError` for certificate parsing, certificate loading, bind failures, and rejected authentication policy.
`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.
`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)

View file

@ -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 `type-maps.yaml` by `Registry::builtin()`.
`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.
### Registry
```rust
use mtp::codec::registry::Registry;
use mtp::codec::Version;
let registry = host.registry();
assert!(registry.supports(&Version(2, 0)));
assert!(registry.supports(&Version(3, 0)));
let negotiated = registry.negotiate(&[Version(1, 0), Version(2, 0)]);
// -> Some(Version(2, 0)) if both versions are registered
let negotiated = registry.negotiate(&[Version(2, 0), Version(3, 0)]);
// -> Some(Version(3, 0)) for this repository's builtin map
```
## Authentication Flow
@ -105,13 +105,15 @@ After a successful handshake, `MTPConnection` exposes `AuthState::Authenticated`
## Handling Messages
Use `conn.sender` and `conn.receiver` for bidirectional message exchange:
Use `conn.sender` and `conn.receive()` for bidirectional message exchange. The
connection dispatcher owns the underlying receiver, especially when `pipes` is
enabled:
```rust
while let Some(conn) = host.accept().await? {
tokio::spawn(async move {
loop {
match conn.receiver.receive().await {
match conn.receive().await {
Ok(msg) => {
let response = process_message(&msg, &conn);
conn.sender.send(&response).await.ok();
@ -212,7 +214,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.to_bytes();
let bytes = host_keyring.try_to_bytes()?;
std::fs::write("host_keys.bin", bytes)?;
```

View file

@ -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()`; 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().await`; its `drain_timeout` controls graceful TCP HTTP completion and the QUIC drain period before remaining connection tasks are terminated.

View file

@ -167,6 +167,9 @@ 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" {
@ -182,7 +185,7 @@ while let Ok(request) = conn.receive_pipe().await {
let mut reader = accept_pipe_session(
reader.into_inner(), &params, &own_keyring, &client_public_bundle,
).await?;
let mut hasher = sha2::Sha256::new();
let mut hasher = Sha256::new();
while let Some(chunk) = reader.read_record().await? {
hasher.update(&chunk);
process_chunk(&chunk).await?;

View file

@ -55,8 +55,28 @@ Verified SDK results expose the authenticated `protectedVersion` and
`finalRecipientId` alongside the application content.
Native applications use the same schema through `ProtectedMessageBuilder` and
`open_protected`; language bindings delegate envelope construction and opening
to this codec boundary.
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.
## Authentication Flow
@ -77,10 +97,24 @@ 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 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.
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.

View file

@ -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. `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. The lower-level `mtp_transport::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,6 +160,14 @@ 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
@ -174,9 +182,11 @@ 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. It
never derives receive policy from the recipient keyring. The sender's
signature suite remains a separate choice. Signature policy must be applied
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
independently to relay metadata, relay content, and pipe session establishment.
### Key history and rotation
@ -202,6 +212,7 @@ 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`.
@ -261,16 +272,43 @@ 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,
and 64 encrypted recipients. Decrypted values are parsed with the same limits.
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`.
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.
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.
## Security Limitations

View file

@ -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. 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. `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.
`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.

View file

@ -1,6 +1,14 @@
# Type Map
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).
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.
## Binary Frame Format
@ -85,6 +93,15 @@ 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.
@ -120,7 +137,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::SomeType).unwrap();
let id = tm.data_id_enum(DataType::ExampleText).unwrap();
```
For native builds with the `registry` feature, the enums are a **union across
@ -134,30 +151,31 @@ compiled by the Vite plugin.
Encoding/decoding uses a `TypeMap` to resolve type names to wire IDs:
```rust
use mtp::codec::{encode, decode, DataValue};
use mtp::codec::{CommunicationType, CommunicationValue, DataType, DataValue};
use mtp::type_map::TypeMap;
let tm = TypeMap::v2_0();
let value = DataValue::Str("hello".into());
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 bytes = encode(&value, &tm).unwrap();
let decoded = decode(&bytes, &tm).unwrap();
let bytes = value.to_bytes().unwrap();
let decoded = CommunicationValue::from_bytes_with(&bytes, &tm).unwrap();
```
```rust
let tm_v3 = TypeMap::v3_0();
assert!(tm_v3.data_id_enum(DataType::SomeType).is_some());
assert!(tm_v3.data_id_enum(DataType::ExampleText).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 self-delimiting codec begins at protocol version `3.0`; older versions are not codec 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 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.
### 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::SomeType → host encodes with v3.0 TypeMap → wire ID 32
v3.0 host receives an unsupported pre-v3.0 peer → version negotiation error
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
```
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.
@ -175,17 +193,33 @@ 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::new(registry);
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);
// Encode with a specific version
let bytes = codec.encode(&value, Version(3, 0)).unwrap();
// The value must retain the negotiated map used to construct it.
let bytes = codec.encode(&value).unwrap();
// Decode with a specific version
let decoded = codec.decode(&bytes, 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();
```
`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.

View file

@ -104,6 +104,9 @@ 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. |
@ -471,6 +474,60 @@ 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
@ -602,8 +659,19 @@ The SDK logger receives parsed events:
```typescript
type MTPLogEvent =
| { hint: "info" | "warning"; type: string; data: unknown }
| { hint: "error"; type: string | "error"; error: string };
| {
hint: "info" | "warning";
type: string;
data: unknown;
direction?: "send" | "recv";
}
| {
hint: "error";
type: string | "error";
error: string;
data?: unknown;
direction?: "send" | "recv";
};
```
Incoming non-error frames and sent frames are logged as `info`. Error frames and transport errors are logged as `error`.

11
example/Cargo.lock generated
View file

@ -225,9 +225,9 @@ dependencies = [
[[package]]
name = "chacha20"
version = "0.10.1"
version = "0.10.2"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "d524456ba66e72eb8b115ff89e01e497f8e6d11d78b70b1aa13c0fbd97540a81"
checksum = "65c35e4b699c7e15ccbe7ee35c005e4fc0a278d22238a2857e6ce2dadeda1b06"
dependencies = [
"cfg-if",
"cpufeatures 0.3.0",
@ -1305,7 +1305,6 @@ name = "mtp-common"
version = "0.3.0"
dependencies = [
"quinn",
"rustls",
"thiserror 2.0.20",
"wtransport",
]
@ -1314,6 +1313,7 @@ dependencies = [
name = "mtp-crypto"
version = "0.3.0"
dependencies = [
"argon2",
"base64 0.22.1",
"chacha20poly1305",
"ed25519-dalek",
@ -1337,7 +1337,6 @@ dependencies = [
name = "mtp-files"
version = "0.3.0"
dependencies = [
"argon2",
"mtp-crypto",
"rand",
"thiserror 2.0.20",
@ -1353,6 +1352,7 @@ dependencies = [
"mtp-crypto",
"mtp-transport",
"rand",
"thiserror 2.0.20",
"tokio",
"tracing",
"wtransport",
@ -1404,7 +1404,6 @@ dependencies = [
"mtp-host",
"mtp-transport",
"quinn",
"rand",
"rustls",
"thiserror 2.0.20",
"tokio",
@ -1702,7 +1701,7 @@ version = "0.10.2"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "c7f5fa3a058cd35567ef9bfa5e75732bee0f9e4c55fa90477bef2dfcdbc4be80"
dependencies = [
"chacha20 0.10.1",
"chacha20 0.10.2",
"getrandom 0.4.3",
"rand_core 0.10.1",
]

View file

@ -62,7 +62,7 @@ pub fn build_demo_message(
)
.add_typed_default(
DataType::Timestamp,
DataValue::UnsignedNumber(timestamp as u128),
DataValue::UnsignedNumber(timestamp),
)
.add_typed_default(DataType::Data, DataValue::Str("Hello, MTP!".into()))
.add_typed_default(DataType::Flags, DataValue::BoolTrue)

View file

@ -79,6 +79,7 @@ pub struct ClientMetrics {
}
impl ClientMetrics {
#[cfg(test)]
pub fn new() -> Self {
Self {
sessions: Vec::new(),

View file

@ -3,8 +3,9 @@ use std::time::{Duration, Instant};
use mtp::client::MTPConnection;
use mtp::codec::{
CommunicationType, DataType, DataValue, ProtectionPolicy, ProtectedMessageBuilder,
ProtectionPurpose, SealedRelayBuilder, SignaturePolicy, TypeMap, open_relay_content,
open_relay_metadata,
ProtectionPurpose, RelayOpenOptions, SealedRelayBuilder, SignaturePolicy, TypeMap,
open_relay_content_with_limits_without_replay,
open_relay_metadata_without_replay,
};
use mtp::common::unix_time_millis;
use mtp::crypto::{Ed25519Signer, Keyring, PublicKeyBundle};
@ -159,22 +160,22 @@ pub async fn send_sealed_relay(
return Err("relay forwarding changed the sealed-sender boundary".into());
}
let metadata = open_relay_metadata(
let metadata = open_relay_metadata_without_replay(
&forwarded,
&final_recipient_keyring,
signer_id,
&signer_keyring.public_key_bundle(),
RELAY_SIGNATURE_POLICY,
RelayOpenOptions::new(RELAY_SIGNATURE_POLICY),
)?;
let application_metadata = metadata
.metadata()
.ok_or("forwarded relay metadata was missing")?;
let content = open_relay_content(
let content = open_relay_content_with_limits_without_replay(
&metadata,
&final_recipient_keyring,
&signer_keyring.public_key_bundle(),
FINAL_RECIPIENT_ID,
RELAY_SIGNATURE_POLICY,
&[&final_recipient_keyring],
&[signer_keyring.public_key_bundle()],
Some(FINAL_RECIPIENT_ID),
RelayOpenOptions::new(RELAY_SIGNATURE_POLICY),
)?;
if content.message_type != "ProtectedMessage" {
return Err(format!("unexpected relay message type: {}", content.message_type).into());

View file

@ -17,14 +17,22 @@ 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.to_bytes(), loaded_keyring.to_bytes());
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.public_key_bundle().as_bytes(),
loaded_bundle.as_bytes()
bundle_bytes,
loaded_bundle_bytes
);
println!(
"\nPrivateKeyRing (base64):\n{}",
keyring.try_to_base64()?
);
println!("\nPrivateKeyRing (base64):\n{}", keyring.to_base64());
println!("\nPublicKeyBundle (base64):\n{}", loaded_bundle.to_base64());
println!(
"\nPublicKeyBundle (base64):\n{}",
loaded_bundle.try_to_base64()?
);
println!("Wrote keyring -> {}", keyring_path.display());
println!("Wrote bundle -> {}", bundle_path.display());

View file

@ -2,8 +2,11 @@ use std::collections::HashMap;
use mtp::codec::{
CommunicationType, CommunicationValue, DataType, DataTypeId, DataValue, InMemoryReplayGuard,
ProtectionPolicy, ProtectionPurpose, SignaturePolicy, TypeMap,
forward_relay_frame, open_protected_with, open_relay_content, open_relay_metadata_with,
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,
};
use mtp::crypto::{Keyring, PublicKeyBundle};
@ -44,7 +47,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 as u128))
.add_data(ts_id, DataValue::UnsignedNumber(now))
.map_err(|e| e.to_string())?
.add_data(data_id, DataValue::Str(data.into()))
.map_err(|e| e.to_string())
@ -65,16 +68,18 @@ fn process_direct_protected(
));
}
let opened = open_protected_with(
let opened = open_protected_with_checked(
msg,
std::slice::from_ref(&host_keyring),
None,
|signer_id| resolve_signer_key(signer_id, registered_clients).map(|key| vec![key]),
ProtectedOpenOptions::new(
Some(DIRECT_DESTINATION_ID),
ProtectionPurpose::from(DIRECT_SIGNATURE_PURPOSE),
ProtectionPurpose::from(DIRECT_ENCRYPTION_PURPOSE),
SIGNATURE_POLICY,
Some(accepted_messages),
),
accepted_messages,
)
.map_err(|e| format!("direct protected message could not be authenticated: {e}"))?;
let signer_id = opened.signer_id;
@ -131,15 +136,15 @@ fn process_sealed_relay(
));
}
let metadata = open_relay_metadata_with(
let metadata = open_relay_metadata_with_checked(
msg,
std::slice::from_ref(&host_keyring),
None,
|signer_id| {
resolve_signer_key(signer_id, registered_clients).map(|key| vec![key])
},
SIGNATURE_POLICY,
Some(accepted_messages),
RelayOpenOptions::new(SIGNATURE_POLICY),
accepted_messages,
)
.map_err(|e| format!("metadata relay could not authenticate metadata: {e}"))?;
println!(
@ -158,13 +163,13 @@ fn process_sealed_relay(
.len()
);
let content_result = open_relay_content(
let content_result = open_relay_content_with_limits_without_replay(
&metadata,
host_keyring,
&resolve_signer_key(metadata.signer_id(), registered_clients)
.ok_or("metadata signer key disappeared")?,
FINAL_RECIPIENT_ID,
SIGNATURE_POLICY,
&[host_keyring],
&[resolve_signer_key(metadata.signer_id(), registered_clients)
.ok_or("metadata signer key disappeared")?],
Some(FINAL_RECIPIENT_ID),
RelayOpenOptions::new(SIGNATURE_POLICY),
);
if content_result.is_ok() {
return Err("metadata relay unexpectedly decrypted final-recipient content".into());
@ -305,12 +310,22 @@ 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(signer_id, pk_bundle, mtp::codec::ProtectionPurpose::from(2))
.verify_with_policy(
signer_id,
pk_bundle,
mtp::codec::ProtectionPurpose::from(2),
SIGNATURE_POLICY,
)
.is_ok()
{
let dv = sig
.clone()
.into_verified(signer_id, pk_bundle, mtp::codec::ProtectionPurpose::from(2))
.into_verified_with_policy(
signer_id,
pk_bundle,
mtp::codec::ProtectionPurpose::from(2),
SIGNATURE_POLICY,
)
.ok();
if let Some(entries) = dv.and_then(|value| value.as_container()) {
println!(" Verified SignedPayload: {:?}", entries);
@ -331,16 +346,22 @@ 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(
.verify_with_policy(
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(signer_id, pk_bundle, mtp::codec::ProtectionPurpose::from(3))
.into_verified_with_policy(
signer_id,
pk_bundle,
mtp::codec::ProtectionPurpose::from(3),
SIGNATURE_POLICY,
)
.ok();
if let Some(entries) = dv.and_then(|value| value.as_container()) {
println!(" Verified SecurePayload: {:?}", entries);
@ -367,7 +388,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 as u128))
.add_data(ts_id, DataValue::UnsignedNumber(now))
.map_err(|e| e.to_string())?
.add_data(
data_id,

View file

@ -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.as_bytes());
let bundle_hex = hex::encode(bundle.try_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?;

View file

@ -38,6 +38,7 @@ 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>> {
@ -100,19 +101,20 @@ async fn main() -> Result<(), Box<dyn std::error::Error>> {
serde_json::to_string_pretty(&*db).ok()
};
if let Some(json) = json {
if let Err(error) = tokio::fs::write("clients.json", json).await {
if let Some(json) = json
&& 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(&host_keyring.to_bytes()) {
match mtp::crypto::Keyring::from_bytes(&decrypt_keyring_bytes) {
Ok(keyring) => keyring,
Err(e) => {
return Err(format!("failed to re-load host keyring for decryption: {e}").into());

View file

@ -108,6 +108,7 @@ pub struct ServerMetrics {
}
impl ServerMetrics {
#[cfg(test)]
pub fn new() -> Self {
Self {
inner: Mutex::new(Inner {
@ -218,6 +219,7 @@ impl ServerMetrics {
}
}
#[cfg(test)]
pub fn snapshot(&self) -> ServerMetricsFile {
let inner = self.inner.lock().unwrap();
self.to_file(&inner)
@ -553,11 +555,11 @@ mod tests {
let metrics = ServerMetrics::new();
for i in 0..3 {
let mut session = metrics.start_session(1000 + i as u64, format!("session {i}"));
let mut session = metrics.start_session(1000 + i, format!("session {i}"));
for _ in 0..(i + 1) * 2 {
session.record_message(Duration::from_millis(1 + i), true);
}
session.record_pipe((i as u64 + 1) * 1000);
session.record_pipe((i + 1) * 1000);
session.finish(format!("exit {i}"));
}

View file

@ -6,8 +6,7 @@ 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"] }
argon2 = "0.5"
mtp-crypto = { version = "0.3.0", path = "../crypto", default-features = false, features = ["chacha20poly1305", "hkdf", "password-kdf"] }
rand = "0.10.2"
thiserror = "2"

View file

@ -119,6 +119,20 @@ 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;
@ -146,7 +160,7 @@ fn write_secret_atomic(path: &Path, bytes: &[u8]) -> io::Result<()> {
let _ = fs::remove_file(&temporary);
return Err(error);
}
Ok(())
sync_parent_directory(path)
}
fn derive_key(
@ -156,21 +170,12 @@ fn derive_key(
iterations: u32,
lanes: u32,
) -> Result<Zeroizing<[u8; 32]>, FileError> {
if salt.len() != SALT_LEN
|| !(8 * 1024..=256 * 1024).contains(&memory_kib)
|| !(1..=10).contains(&iterations)
|| !(1..=8).contains(&lanes)
{
if salt.len() != SALT_LEN {
return Err(FileError::Crypto(CryptoError::KdfError));
}
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)
Ok(Zeroizing::new(mtp_crypto::derive_password_key(
passphrase, salt, memory_kib, iterations, lanes,
)?))
}
fn protected_header_aad(parameters: &[u8]) -> Vec<u8> {
@ -206,7 +211,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.to_bytes();
let plaintext = keyring.try_to_bytes()?;
let encrypted = cipher.encrypt(&plaintext, &protected_header_aad(&parameters))?;
let mut payload = Vec::with_capacity(PROTECTED_PARAMS_LEN + encrypted.len());
payload.extend_from_slice(&parameters);
@ -259,7 +264,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.to_bytes();
let payload = keyring.try_to_bytes()?;
let bytes = Zeroizing::new(encode(KEYRING_MAGIC, RAW_FORMAT_VERSION, &payload));
write_secret_atomic(path.as_ref(), &bytes)?;
Ok(())
@ -286,7 +291,8 @@ pub fn save_public_key_bundle(
bundle: &PublicKeyBundle,
path: impl AsRef<Path>,
) -> Result<(), FileError> {
let bytes = encode(BUNDLE_MAGIC, BUNDLE_FORMAT_VERSION, &bundle.as_bytes());
let bundle_bytes = bundle.try_as_bytes()?;
let bytes = encode(BUNDLE_MAGIC, BUNDLE_FORMAT_VERSION, &bundle_bytes);
fs::write(path, bytes)?;
Ok(())
}
@ -339,7 +345,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.to_bytes(), loaded.to_bytes());
assert_eq!(keyring.try_to_bytes()?, loaded.try_to_bytes()?);
let _ = fs::remove_file(&path);
Ok(())
}
@ -350,7 +356,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.as_bytes(), loaded.as_bytes());
assert_eq!(bundle.try_as_bytes()?, loaded.try_as_bytes()?);
let _ = fs::remove_file(&path);
Ok(())
}
@ -414,7 +420,7 @@ mod tests {
Err(FileError::UnprotectedKeyring)
));
let loaded = load_keyring_raw(&path)?;
assert_eq!(keyring.to_bytes(), loaded.to_bytes());
assert_eq!(keyring.try_to_bytes()?, loaded.try_to_bytes()?);
let _ = fs::remove_file(&path);
Ok(())
}
@ -423,7 +429,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.to_bytes();
let serialized = keyring.try_to_bytes()?;
save_keyring(&keyring, &path, b"passphrase")?;
let stored = fs::read(&path)?;
assert!(

View file

@ -1,31 +1,34 @@
{
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:
nixpkgs.lib.foldl' nixpkgs.lib.recursiveUpdate {} (
map (system: nixpkgs.lib.mapAttrs (_: value: {${system} = value;}) (f system)) systems
eachSystem =
f:
nixpkgs.lib.foldl' nixpkgs.lib.recursiveUpdate { } (
map (system: nixpkgs.lib.mapAttrs (_: value: { ${system} = value; }) (f system)) systems
);
in
eachSystem (
system: let
overlays = [rust-overlay.overlays.default];
pkgs = import nixpkgs {inherit system overlays;};
system:
let
overlays = [ rust-overlay.overlays.default ];
pkgs = import nixpkgs { inherit system overlays; };
rustToolchain = pkgs.rust-bin.stable.latest.default.override {
extensions = [
@ -33,12 +36,12 @@
"clippy"
"rustfmt"
];
targets = ["wasm32-unknown-unknown"];
targets = [ "wasm32-unknown-unknown" ];
};
clippyCheck = pkgs.writeShellApplication {
name = "mtp-clippy";
runtimeInputs = [rustToolchain];
runtimeInputs = [ rustToolchain ];
text = ''
export MTP_TYPE_MAPS="''${MTP_TYPE_MAPS:-$PWD/example/type-maps.yaml}"
cargo clippy --workspace --exclude mtp-wasm --all-targets --all-features -- -D warnings -W unreachable-pub
@ -47,7 +50,7 @@
macheteCheck = pkgs.writeShellApplication {
name = "mtp-machete";
runtimeInputs = [pkgs.cargo-machete];
runtimeInputs = [ pkgs.cargo-machete ];
text = ''
cargo machete "$@"
'';
@ -55,7 +58,15 @@
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}"
@ -67,7 +78,6 @@
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
@ -80,13 +90,17 @@
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";

View file

@ -9,6 +9,7 @@ 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"

View file

@ -5,10 +5,14 @@ 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;
@ -82,6 +86,159 @@ 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,
@ -94,6 +251,8 @@ 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,
@ -113,6 +272,10 @@ 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 {
@ -127,6 +290,8 @@ 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,
@ -153,6 +318,13 @@ 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,
}
}
@ -173,7 +345,9 @@ 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);
@ -183,6 +357,7 @@ 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
}
@ -210,4 +385,141 @@ 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(&registration)
.expect("registration attempt decision")
);
assert!(
!limiter
.allow(&registration)
.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
);
}
}

View file

@ -11,7 +11,10 @@ use tokio::sync::{Mutex, mpsc};
#[cfg(feature = "crypto")]
use crate::error::random_client_id;
#[cfg(feature = "pipes")]
use crate::pipe::{PipeDispatcher, PipeReceiver, PipeRequest, PipeSender, run_dispatcher};
use crate::pipe::{
PendingCreationGuard, PipeDispatcher, PipeReceiver, PipeRequest, PipeSender,
is_expired_creation, run_dispatcher,
};
#[cfg(feature = "pipes")]
use mtp_transport::Policy;
@ -70,7 +73,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, P>>>,
pub(crate) pipe_req_rx: Mutex<mpsc::Receiver<PipeRequest<S, R, P>>>,
#[cfg(feature = "pipes")]
pub(crate) pipe_dispatcher: Arc<PipeDispatcher<P>>,
#[cfg(not(feature = "pipes"))]
@ -145,10 +148,12 @@ where
remote_addr: Option<SocketAddr>,
) -> Self {
let policy = Arc::new(Policy::default());
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 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 dispatcher = Arc::new(PipeDispatcher {
pending_creations: Mutex::new(std::collections::HashMap::new()),
pending_creations: std::sync::Mutex::new(std::collections::HashMap::new()),
expired_creations: std::sync::Mutex::new(std::collections::HashMap::new()),
pending_pipes: Mutex::new(std::collections::HashMap::new()),
policy,
type_map: codec.type_map().clone(),
@ -198,10 +203,12 @@ where
remote_addr: Option<SocketAddr>,
policy: Arc<Policy>,
) -> Self {
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 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 dispatcher = Arc::new(PipeDispatcher {
pending_creations: Mutex::new(std::collections::HashMap::new()),
pending_creations: std::sync::Mutex::new(std::collections::HashMap::new()),
expired_creations: std::sync::Mutex::new(std::collections::HashMap::new()),
pending_pipes: Mutex::new(std::collections::HashMap::new()),
policy,
type_map: codec.type_map().clone(),
@ -321,19 +328,37 @@ where
pub async fn create_pipe(
&self,
description: &str,
) -> Result<crate::pipe::PipeHandle<S>, mtp_common::PipeError> {
) -> Result<crate::pipe::PipeHandle<S, P>, mtp_common::PipeError> {
let (response_tx, response_rx) = tokio::sync::oneshot::channel();
let pipe_id = {
let mut pending = self.pipe_dispatcher.pending_creations.lock().await;
let mut pending = self
.pipe_dispatcher
.pending_creations
.lock()
.map_err(|_| mtp_common::PipeError::ConnectionClosed)?;
let pipe_id = loop {
let candidate = rand::random::<u32>();
if candidate != 0 && !pending.contains_key(&candidate) {
if candidate != 0
&& !pending.contains_key(&candidate)
&& !is_expired_creation(&self.pipe_dispatcher, candidate)
{
break candidate;
}
};
pending.insert(pipe_id, response_tx);
pipe_id
let token = Arc::new(());
pending.insert(
pipe_id,
crate::pipe::PendingCreation {
token: token.clone(),
sender: response_tx,
},
);
drop(pending);
(pipe_id, token)
};
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,
@ -342,23 +367,21 @@ 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, P>, CommunicationError> {
pub async fn receive_pipe(&self) -> Result<PipeRequest<S, R, P>, CommunicationError> {
self.pipe_req_rx
.lock()
.await

189
host/src/engine.rs Normal file → Executable file
View 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::HostConfig;
use crate::config::{AuthenticationContext, HostConfig};
use crate::error::AcceptError;
use mtp_codec::{
CommunicationType, CommunicationValue, DataType, DataValue, TypeMap, Version,
@ -127,19 +127,32 @@ 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(
self.accept_until_with_context(
sender,
receiver,
tokio::time::Instant::now() + self.config.auth_timeout,
context,
)
.await
}
#[cfg(not(feature = "crypto"))]
{
let result = self.accept_inner(sender, receiver).await;
let result = self.accept_inner(sender, receiver, &context).await;
if result.is_err() {
sender.close();
}
@ -159,7 +172,22 @@ impl HandshakeEngine {
receiver: &R,
deadline: tokio::time::Instant,
) -> Result<HandshakeResult, AcceptError> {
match tokio::time::timeout_at(deadline, self.accept_inner(sender, receiver)).await {
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
{
Ok(result) => {
if result.is_err() {
sender.close();
@ -186,8 +214,15 @@ 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(),
@ -242,6 +277,11 @@ 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()))?;
@ -257,6 +297,48 @@ 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(
@ -344,7 +426,7 @@ impl HandshakeEngine {
// Authenticated clients include PublicKeys in Identification as an
// intent marker; this avoids acknowledging the opening as a guest
// connection and leaving the client waiting for a Challenge.
if Some(first_msg.get_type()) == CommunicationType::Register.try_to_id(&tm)
if Some(first_msg.get_type()) == CommunicationType::Register.try_to_id(tm)
|| first_msg.get_data(DataType::PublicKeys).is_some()
{
send_rejection_generic(
@ -400,7 +482,7 @@ impl HandshakeEngine {
let tm = codec.type_map();
// Register frames always go through full authentication
if Some(first_msg.get_type()) == CommunicationType::Register.try_to_id(&tm) {
if Some(first_msg.get_type()) == CommunicationType::Register.try_to_id(tm) {
let bundle = match extract_register_bundle(&first_msg) {
Ok(bundle) => bundle,
Err(error) => {
@ -408,7 +490,14 @@ impl HandshakeEngine {
return Err(error);
}
};
let pk_bytes = bundle.as_bytes();
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);
}
};
return self
.complete_auth_handshake(
sender,
@ -425,7 +514,7 @@ impl HandshakeEngine {
}
// Identification: try lookup, fall back to guest
if Some(first_msg.get_type()) == CommunicationType::Identification.try_to_id(&tm) {
if Some(first_msg.get_type()) == CommunicationType::Identification.try_to_id(tm) {
let cid = match first_msg.get_data(DataType::Id) {
Some(DataValue::UnsignedNumber(n)) => u64::try_from(*n).unwrap_or(0),
_ => 0,
@ -453,6 +542,28 @@ 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(),
);
@ -461,6 +572,7 @@ 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) => {
@ -469,6 +581,7 @@ 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)?;
@ -504,7 +617,7 @@ impl HandshakeEngine {
let tm = codec.type_map();
let (flow, response_type) = if Some(first_msg.get_type())
== CommunicationType::Identification.try_to_id(&tm)
== CommunicationType::Identification.try_to_id(tm)
{
let cid = match first_msg.get_data(DataType::Id) {
Some(DataValue::UnsignedNumber(n)) => match u64::try_from(*n) {
@ -525,6 +638,14 @@ 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,
@ -540,12 +661,13 @@ impl HandshakeEngine {
"unknown client id".into(),
));
}
}
};
(
Flow::Login { id: cid, bundle },
CommunicationType::IdentificationResponse,
)
} else if Some(first_msg.get_type()) == CommunicationType::Register.try_to_id(&tm) {
} else if Some(first_msg.get_type()) == CommunicationType::Register.try_to_id(tm) {
let bundle = match extract_register_bundle(&first_msg) {
Ok(bundle) => bundle,
Err(error) => {
@ -553,7 +675,14 @@ impl HandshakeEngine {
return Err(error);
}
};
let pk_bytes = bundle.as_bytes();
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);
}
};
(
Flow::Register { bundle, pk_bytes },
CommunicationType::RegisterResponse,
@ -710,7 +839,7 @@ impl HandshakeEngine {
sender.close();
AcceptError::Receive(e)
})?;
if Some(proof.get_type()) != CommunicationType::ChallengeResponse.try_to_id(&tm) {
if Some(proof.get_type()) != CommunicationType::ChallengeResponse.try_to_id(tm) {
let error = AcceptError::AuthenticationFailed("missing challenge response".into());
reject_error_generic(sender, &error, tm).await;
return Err(error);
@ -787,7 +916,14 @@ impl HandshakeEngine {
Flow::Login { id, bundle } => (id, bundle),
Flow::Register { bundle, .. } => {
let _registration_guard = self.config.registration_lock.lock().await;
let identity = bundle.as_bytes();
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 cached_id = self
.config
.registration_ids
@ -988,6 +1124,12 @@ 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;
}
@ -1021,6 +1163,11 @@ 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
}
@ -1041,8 +1188,8 @@ impl HandshakeSender for mtp_transport::Sender {
) -> impl std::future::Future<Output = Result<(), CommunicationError>> + Send {
mtp_transport::Sender::finish_stream(self)
}
fn set_type_map(&self, type_map: &TypeMap) -> impl std::future::Future<Output = ()> + Send {
async move { self.set_type_map(type_map).await }
async fn set_type_map(&self, type_map: &TypeMap) {
self.set_type_map(type_map).await;
}
fn close(&self) {
let sender = self.clone();
@ -1058,8 +1205,8 @@ impl HandshakeReceiver for mtp_transport::Receiver {
mtp_transport::Receiver::receive(self)
}
fn set_type_map(&self, type_map: &TypeMap) -> impl std::future::Future<Output = ()> + Send {
async move { self.set_type_map(type_map).await }
async fn set_type_map(&self, type_map: &TypeMap) {
self.set_type_map(type_map).await;
}
}
@ -1075,8 +1222,8 @@ impl<C: mtp_transport::TransportConnection> HandshakeSender for mtp_transport::G
) -> impl std::future::Future<Output = Result<(), CommunicationError>> + Send {
mtp_transport::GenericSender::finish_stream(self)
}
fn set_type_map(&self, type_map: &TypeMap) -> impl std::future::Future<Output = ()> + Send {
async move { self.set_type_map(type_map).await }
async fn set_type_map(&self, type_map: &TypeMap) {
self.set_type_map(type_map).await;
}
fn close(&self) {
mtp_transport::GenericSender::close(self);
@ -1093,8 +1240,8 @@ impl<C: mtp_transport::TransportConnection> HandshakeReceiver
mtp_transport::GenericReceiver::receive(self)
}
fn set_type_map(&self, type_map: &TypeMap) -> impl std::future::Future<Output = ()> + Send {
async move { self.set_type_map(type_map).await }
async fn set_type_map(&self, type_map: &TypeMap) {
self.set_type_map(type_map).await;
}
}

View file

@ -11,7 +11,7 @@ use std::time::Instant;
#[cfg(feature = "pipes")]
use tokio::sync::mpsc;
use crate::config::HostConfig;
use crate::config::{AuthenticationContext, HostConfig};
use crate::connection::MTPConnection;
use crate::engine::HandshakeEngine;
use crate::error::AcceptError;
@ -135,7 +135,16 @@ impl HandshakeContext {
receiver: Receiver,
) -> Result<Option<MTPConnection>, AcceptError> {
let engine = HandshakeEngine::new(self.registry.clone(), self.config.clone());
let result = engine.accept(&sender, &receiver).await?;
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?;
#[cfg(feature = "crypto")]
{
Ok(Some(self.connection_from_handshake_result(
@ -171,12 +180,13 @@ impl HandshakeContext {
receiver.respond_to_pings(sender.clone());
}
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 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 dispatcher = Arc::new(PipeDispatcher {
pending_creations: tokio::sync::Mutex::new(std::collections::HashMap::new()),
pending_creations: std::sync::Mutex::new(std::collections::HashMap::new()),
expired_creations: std::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(),
@ -260,12 +270,13 @@ impl HandshakeContext {
receiver.respond_to_pings(sender.clone());
}
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 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 dispatcher = Arc::new(PipeDispatcher {
pending_creations: tokio::sync::Mutex::new(std::collections::HashMap::new()),
pending_creations: std::sync::Mutex::new(std::collections::HashMap::new()),
expired_creations: std::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,

View file

@ -29,8 +29,9 @@ pub use mtp_codec::registry::Registry;
#[cfg(feature = "crypto")]
pub use config::{
AuthenticationPolicy, CompleteRegister, FindRegisteredClient, GetExistingClient,
GuestIdGenerator,
AuthenticationAttempt, AuthenticationAttemptLimiter, AuthenticationContext,
AuthenticationLimitError, AuthenticationPolicy, CompleteRegister, FindRegisteredClient,
GetExistingClient, GuestIdGenerator, InMemoryAuthenticationAttemptLimiter,
};
#[cfg(feature = "crypto")]
pub use error::AuthState;

View file

@ -3,6 +3,7 @@ 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.
@ -26,6 +27,10 @@ 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;
@ -51,6 +56,14 @@ 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> {
@ -86,6 +99,14 @@ 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> {
@ -93,14 +114,20 @@ where
}
}
pub struct PipeHandle<S: PipeSender> {
pub struct PipeHandle<S: PipeSender, P = wtransport::RecvStream> {
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: PipeSender> PipeHandle<S> {
impl<S, P> PipeHandle<S, P>
where
S: PipeSender,
P: tokio::io::AsyncRead + Send + Unpin + 'static,
{
pub fn pipe_id(&self) -> u32 {
self.pipe_id
}
@ -109,31 +136,96 @@ impl<S: PipeSender> PipeHandle<S> {
&self.description
}
pub async fn wait(self) -> Result<Option<PipeWriter<S::Writer>>, PipeError> {
match self.response_rx.await {
Ok(Ok(true)) => self
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
.sender
.open_pipe_stream(self.pipe_id, &self.description)
.await
.map(Some)
.map_err(PipeError::from),
Ok(Ok(false)) => Ok(None),
Ok(Err(error)) => Err(error),
Err(_) => Err(PipeError::StreamClosed),
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)
}
}
}
}
pub struct PipeRequest<S, P> {
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(crate) pipe_id: u32,
pub(crate) description: String,
pub(crate) sender: S,
pub(crate) receiver: R,
pub(crate) dispatcher: Arc<PipeDispatcher<P>>,
}
impl<S, P> PipeRequest<S, 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>
where
S: PipeSender,
R: PipeReceiver<P>,
P: tokio::io::AsyncRead + Send + Unpin + 'static,
{
pub fn id(&self) -> u32 {
@ -145,6 +237,10 @@ 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
@ -168,7 +264,10 @@ where
}
match tokio::time::timeout(self.dispatcher.policy.read_timeout, pipe_rx).await {
Ok(Ok(reader)) => Ok(reader),
Ok(Ok(reader)) => {
expected_pipe.disarm();
Ok(reader)
}
Ok(Err(_)) => {
self.dispatcher
.pending_pipes
@ -203,18 +302,137 @@ where
}
pub(crate) struct PipeDispatcher<P> {
pub(crate) pending_creations:
Mutex<HashMap<u32, tokio::sync::oneshot::Sender<Result<bool, PipeError>>>>,
pub(crate) pending_creations: StdMutex<HashMap<u32, PendingCreation>>,
pub(crate) expired_creations: StdMutex<HashMap<u32, tokio::time::Instant>>,
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, P>>,
pipe_req_tx: mpsc::Sender<PipeRequest<S, R, P>>,
dispatcher: Arc<PipeDispatcher<P>>,
) where
S: PipeSender,
@ -241,6 +459,7 @@ 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;
@ -256,10 +475,17 @@ pub(crate) async fn run_dispatcher<S, R, P>(
}
continue;
};
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)));
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");
}
continue;
}
@ -292,15 +518,19 @@ 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;
}
}
}
}
}

View file

@ -25,7 +25,6 @@ rustls = "0.23"
tracing = "0.1"
thiserror = "2"
async-trait = "0.1"
rand = { version = "0.10.1", optional = true }
[dev-dependencies]
rcgen = "0.14"
@ -33,5 +32,5 @@ hyper = { version = "1", features = ["client", "http2"] }
[features]
default = []
crypto = ["mtp-host/crypto", "dep:rand"]
crypto = ["mtp-host/crypto"]
pipes = ["mtp-host/pipes", "mtp-transport/pipes"]

View file

@ -143,6 +143,11 @@ 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(),

View file

@ -32,6 +32,8 @@ pub struct H3TransportSender {
pub struct H3TransportReceiver {
stream: H3RecvStream,
quinn: quinn::Connection,
read_exact_calls: u64,
}
impl H3TransportConnection {
@ -42,6 +44,11 @@ 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]
@ -50,14 +57,14 @@ impl TransportSendStream for H3TransportSender {
self.stream
.write_all(buf)
.await
.map_err(|_| CommunicationError::StreamError)?;
.map_err(|_| CommunicationError::DeliveryUnknown)?;
// 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::StreamError)
.map_err(|_| CommunicationError::DeliveryUnknown)
}
async fn finish(&mut self) -> Result<(), CommunicationError> {
@ -66,23 +73,53 @@ 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(|_| ())
.map(|_| {
if first_read {
tracing::debug!(
remote = %self.quinn.remote_address(),
bytes = buf.len(),
header = ?buf,
"received first bytes from WebTransport MTP stream"
);
}
})
.map_err(|error| {
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.
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.
*/
return CommunicationError::StreamClosed;
}
error!("[mtp-webserver] receive stream read_exact failed ({} bytes): {error}", buf.len());
tracing::warn!(len = buf.len(), %error, "WebTransport receive stream read_exact failed");
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"
);
CommunicationError::StreamError
})
}
@ -96,6 +133,9 @@ 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
@ -105,6 +145,11 @@ 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 {
@ -162,10 +207,27 @@ impl TransportConnection for H3TransportConnection {
loop {
match self.session.accept_uni().await {
Ok(Some((id, stream))) if id == self.session.session_id() => {
return Ok(H3TransportReceiver { stream });
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,
});
}
Ok(Some(_)) => {
Ok(Some((stream_session_id, _stream))) => {
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),
@ -220,6 +282,7 @@ pub type WebMtpReceiver = GenericReceiver<H3TransportConnection>;
pub type WebMTPConnection =
mtp_host::MTPConnection<WebMtpSender, WebMtpReceiver, H3TransportReceiver>;
#[allow(clippy::too_many_arguments)]
pub(crate) async fn accept_web_connection(
session: Arc<Session>,
path: String,
@ -268,6 +331,7 @@ pub(crate) async fn accept_web_connection(
.await
}
#[allow(clippy::too_many_arguments)]
async fn accept_web_connection_inner(
session: Arc<Session>,
path: String,
@ -281,6 +345,8 @@ 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());
@ -288,12 +354,27 @@ async fn accept_web_connection_inner(
let engine = mtp_host::HandshakeEngine::new(Registry::builtin(), host_config);
#[cfg(feature = "crypto")]
let result = engine
.accept_until(
.accept_until_with_context(
&sender,
&receiver,
deadline.expect("crypto WebTransport handshakes have a deadline"),
mtp_host::AuthenticationContext {
peer_network_identity: Some(remote_addr.to_string()),
connection_id,
},
)
.await?;
.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?;
#[cfg(not(feature = "crypto"))]
let result = engine.accept(&sender, &receiver).await?;

View file

@ -58,19 +58,20 @@
"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 clean && pnpm run build && pnpm pack",
"pack": "pnpm run release:web",
"release:web": "node create-web-release.mjs",
"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:types && pnpm run test:vite && pnpm run test:boundary"
"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"
},
"devDependencies": {
"@types/node": "^26.0.1",
"jscpd": "4.2.5",
"jscpd": "5.0.14",
"typescript": "^7.0.0"
},
"dependencies": {

935
pnpm-lock.yaml generated

File diff suppressed because it is too large Load diff

2760
src/sdk/client.ts Normal file

File diff suppressed because it is too large Load diff

959
src/sdk/codec.ts Normal file
View file

@ -0,0 +1,959 @@
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;
}

26
src/sdk/credentials.ts Normal file
View file

@ -0,0 +1,26 @@
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);
}

File diff suppressed because it is too large Load diff

View file

@ -0,0 +1,33 @@
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) });
}
};

258
src/sdk/protection.ts Normal file
View file

@ -0,0 +1,258 @@
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 Normal file
View file

@ -0,0 +1,204 @@
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 };

234
src/sdk/schema.ts Normal file
View file

@ -0,0 +1,234 @@
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();
};
}
}

View file

@ -92,8 +92,14 @@ export function signatureVerificationPolicyValue(
return bindings.mtp_protection_signature_suite_ed25519();
case "dual":
return bindings.mtp_protection_signature_suite_dual();
case "any-supported":
return 0;
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;
}
}
}

27
src/sdk/timeout.ts Normal file
View file

@ -0,0 +1,27 @@
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);
}
}
}

27
src/sdk/wasm-init.ts Normal file
View file

@ -0,0 +1,27 @@
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();

20
test/wasm-init.mjs Normal file
View file

@ -0,0 +1,20 @@
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);
});

View file

@ -1,12 +1,18 @@
use crate::ConnectionHandle;
use crate::framing::RetryClassifier;
#[cfg(feature = "pipes")]
use crate::pipe::PipeReader;
use mtp_codec::{CommunicationValue, DecodeLimits, TypeMap};
use mtp_codec::{CommunicationValue, DecodeError, DecodeLimits, EncodeLimits, 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, sleep, timeout};
use tokio::time::{Duration, Instant, sleep, timeout, timeout_at};
use tracing::{debug, info, instrument, trace, warn};
use wtransport::Connection;
@ -19,6 +25,63 @@ 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,
@ -135,27 +198,62 @@ 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<Policy>,
policy: Arc<RuntimePolicy>,
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,
@ -174,7 +272,11 @@ impl Sender {
data: &CommunicationValue,
policy: &Policy,
) -> Result<(), CommunicationError> {
let bytes = data.to_bytes().map_err(|_| CommunicationError::Encode)?;
let bytes = data
.to_bytes_with_limits(EncodeLimits::for_transport_message_size(
policy.max_message_size,
))
.map_err(|_| CommunicationError::Encode)?;
if bytes.len() as u64 > policy.max_message_size
|| bytes.len() as u64 >= policy.close_frame_len as u64
{
@ -190,15 +292,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::StreamClosed)
Err(CommunicationError::DeliveryUnknown)
}
Ok(Err(other)) => {
warn!("[Sender] write failed: {other}");
Err(CommunicationError::StreamError)
Err(CommunicationError::DeliveryUnknown)
}
Err(_) => {
warn!("[Sender] write timed out (len={})", bytes.len());
Err(CommunicationError::StreamError)
Err(CommunicationError::DeliveryUnknown)
}
}
}
@ -269,11 +371,11 @@ impl Sender {
return Ok(());
}
let err = res.err().unwrap_or(CommunicationError::StreamError);
if !matches!(
err,
CommunicationError::StreamError | CommunicationError::StreamClosed
) {
let err = match res {
Ok(()) => return Ok(()),
Err(error) => error,
};
if !RetryClassifier::retry_persistent_stream(&err) {
return Err(err);
}
*stream_opt = None;
@ -300,15 +402,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::StreamClosed)
Err(CommunicationError::DeliveryUnknown)
}
Ok(Err(other)) => {
warn!("[Sender] finish failed: {other}");
Err(CommunicationError::StreamError)
Err(CommunicationError::DeliveryUnknown)
}
Err(_) => {
warn!("[Sender] finish timed out");
Err(CommunicationError::StreamError)
Err(CommunicationError::DeliveryUnknown)
}
}
}
@ -356,20 +458,24 @@ impl Sender {
#[instrument(skip(self, data), level = "trace")]
pub async fn send(&self, data: &CommunicationValue) -> Result<(), CommunicationError> {
if self.handle.is_closed() {
let _send_lock = self.send_guard.lock().await;
{
let state = self.state.lock().await;
if *state != SenderState::Open {
return Err(self
.handle
.close_reason()
.unwrap_or(CommunicationError::UseAfterClosed));
.unwrap_or(CommunicationError::StreamClosed));
}
}
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);
}
@ -402,6 +508,7 @@ 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()));
}
@ -413,6 +520,12 @@ 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 {
@ -445,11 +558,12 @@ impl Sender {
pipe_id: u32,
description: &str,
) -> Result<crate::pipe::PipeWriter, CommunicationError> {
if self.handle.is_closed() {
let _send_lock = self.send_guard.lock().await;
if *self.state.lock().await != SenderState::Open {
return Err(self
.handle
.close_reason()
.unwrap_or(CommunicationError::UseAfterClosed));
.unwrap_or(CommunicationError::StreamClosed));
}
if self.connection.quic_connection().close_reason().is_some() {
@ -486,13 +600,26 @@ 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(())) => {}
@ -505,12 +632,15 @@ 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(
@ -528,13 +658,23 @@ impl Sender {
let connection = self.connection.clone();
let handle = self.handle.clone();
let policy = self.policy.clone();
let mut stream_opt = self.stream_guard.lock().await;
let _send_lock = self.send_guard.lock().await;
{
let mut state = self.state.lock().await;
if *state != SenderState::Open {
return;
}
*state = SenderState::Closing;
}
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 {
@ -553,10 +693,13 @@ 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(
@ -605,6 +748,9 @@ 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 {
@ -626,12 +772,26 @@ impl Drop for Receiver {
#[derive(Clone, Default)]
struct PingControl {
pong_sender: Option<Sender>,
pong_observer: Option<mpsc::UnboundedSender<CommunicationValue>>,
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
}
}
}
impl Receiver {
pub fn new(connection: Connection, handle: Arc<ConnectionHandle>, policy: Arc<Policy>) -> Self {
Self::new_with_max_message_size(connection, handle, policy.clone(), policy.max_message_size)
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)
}
#[cfg(feature = "host")]
@ -640,6 +800,7 @@ 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);
@ -649,7 +810,7 @@ impl Receiver {
fn new_with_max_message_size(
connection: Connection,
handle: Arc<ConnectionHandle>,
policy: Arc<Policy>,
policy: Arc<RuntimePolicy>,
initial_max_message_size: u64,
) -> Self {
#[cfg(feature = "pipes")]
@ -674,6 +835,12 @@ 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!(
@ -742,6 +909,9 @@ 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;
@ -761,7 +931,14 @@ impl Receiver {
}
let frame_limit = stream_max_message_size.load(Ordering::Relaxed);
match Self::read_one_frame(&mut s, &stream_policy, frame_limit).await {
match Self::read_one_frame(
&mut s,
&stream_policy,
frame_limit,
&stream_decode_rejections,
)
.await
{
Ok(ReceivedFrame::Message(mut msg)) => {
let negotiated_type_map =
stream_type_map.read().await.clone();
@ -770,17 +947,32 @@ impl Receiver {
#[cfg(feature = "pipes")]
{
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(),
if frame_count == 1 {
let is_pipe_request = msg.is_type(
mtp_codec::CommunicationType::PipeRequest,
);
let _ = msg_tx_stream.send(Err(error.clone())).await;
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;
stream_handle.close(Some(error));
break;
}
};
if let FirstFrameDisposition::Pipe(pipe_id) = disposition {
let description = msg
.get_str(mtp_codec::DataType::Description)
.unwrap_or("")
@ -804,25 +996,19 @@ impl Receiver {
break;
}
}
}
let control = {
let control = stream_ping_control.read().await;
if msg.is_type(mtp_codec::CommunicationType::Ping) {
control
let pong_sender = if msg.is_type(mtp_codec::CommunicationType::Ping) {
stream_ping_control
.read()
.await
.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((Some(sender), _)) = control {
if let Some(sender) = pong_sender {
let mut pong = CommunicationValue::new_with_type_map(
mtp_codec::CommunicationType::Pong,
&negotiated_type_map,
@ -844,8 +1030,18 @@ impl Receiver {
continue;
}
if let Some((_, Some(observer))) = control {
let _ = observer.send(msg);
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);
}
continue;
}
@ -942,6 +1138,9 @@ impl Receiver {
queue_notify,
max_message_size,
type_map,
decode_rejections,
#[cfg(feature = "pipes")]
expected_pipes,
}),
}
}
@ -958,6 +1157,34 @@ 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() {
@ -967,16 +1194,21 @@ impl Receiver {
}
}
/* Route reserved Pong frames to a connection-level observer. */
pub async fn observe_pongs(&self, observer: mpsc::UnboundedSender<CommunicationValue>) {
/* Route only the currently expected reserved Pong through a bounded observer. */
pub async fn observe_pongs_bounded(&self, observer: mpsc::Sender<CommunicationValue>) {
self.inner.ping_control.write().await.pong_observer = Some(observer);
}
#[instrument(skip(stream, policy), level = "trace")]
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")]
async fn read_one_frame(
stream: &mut wtransport::RecvStream,
policy: &Policy,
policy: &RuntimePolicy,
max_message_size: u64,
decode_rejections: &DecodeRejectionCounters,
) -> Result<ReceivedFrame, CommunicationError> {
use wtransport::error::{StreamReadError, StreamReadExactError};
@ -1011,6 +1243,7 @@ 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
@ -1020,29 +1253,29 @@ impl Receiver {
return Err(CommunicationError::MessageTooLarge);
}
// 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))
// 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)
.map_err(|_| CommunicationError::MessageTooLarge)?;
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]),
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]),
)
.await
{
Ok(Ok(())) => {
buf.try_reserve(chunk_len)
.map_err(|_| CommunicationError::MessageTooLarge)?;
buf.extend_from_slice(&chunk[..chunk_len]);
}
Ok(Ok(())) => body_offset += chunk_len,
Ok(Err(StreamReadExactError::FinishedEarly(n))) => {
warn!(
"[Receiver] body read ended early ({}/{body_len} bytes): stream closed by peer",
buf.len() + n
body_offset.saturating_sub(4) + n
);
return Err(CommunicationError::StreamError);
}
@ -1063,14 +1296,21 @@ impl Receiver {
}
}
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(
let message = match CommunicationValue::try_from_bytes_with_limits(
&frame,
DecodeLimits::for_transport_message_size(max_message_size),
)
.map_err(|_| CommunicationError::ParseCommunicationValue)?;
) {
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);
}
};
Ok(ReceivedFrame::Message(message))
}
@ -1078,18 +1318,12 @@ impl Receiver {
#[instrument(skip(self), level = "trace")]
pub async fn receive(&self) -> Result<CommunicationValue, CommunicationError> {
let mut close_rx = self.inner.handle.subscribe_close();
if close_rx.borrow().is_some() {
return Err(self
.inner
.handle
.close_reason()
.unwrap_or(CommunicationError::StreamClosed));
}
#[cfg(feature = "pipes")]
{
let mut rx = self.inner.msg_rx.lock().await;
let result = tokio::select! {
biased;
message = rx.recv() => message,
_ = close_rx.changed() => return Err(close_rx
.borrow()
@ -1113,6 +1347,7 @@ impl Receiver {
{
let mut rx = self.inner.rx.lock().await;
let result = tokio::select! {
biased;
message = rx.recv() => message,
_ = close_rx.changed() => return Err(close_rx
.borrow()
@ -1136,17 +1371,11 @@ impl Receiver {
#[cfg(feature = "pipes")]
#[instrument(skip(self), level = "trace")]
pub async fn receive_event(&self) -> Result<TransportEvent, CommunicationError> {
if self.inner.handle.is_closed() {
return Err(self
.inner
.handle
.close_reason()
.unwrap_or(CommunicationError::StreamClosed));
}
let mut close_rx = self.inner.handle.subscribe_close();
let mut msg_rx = self.inner.msg_rx.lock().await;
let mut pipe_rx = self.inner.pipe_rx.lock().await;
tokio::select! {
biased;
msg = msg_rx.recv() => {
match msg {
Some(Ok(val)) => {
@ -1174,6 +1403,10 @@ impl Receiver {
.unwrap_or(CommunicationError::StreamClosed)),
}
}
_ = close_rx.changed() => Err(close_rx
.borrow()
.clone()
.unwrap_or(CommunicationError::StreamClosed)),
}
}
@ -1274,4 +1507,63 @@ 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,
}
);
}
}

View file

@ -2,22 +2,26 @@ use mtp_common::CommunicationError;
use std::net::SocketAddr;
use std::sync::{
Arc,
atomic::{AtomicBool, Ordering},
atomic::{AtomicBool, AtomicU64, 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,
@ -35,6 +39,11 @@ 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)
}

View file

@ -1,7 +1,21 @@
use crate::{Policy, TransportSendStream};
use mtp_codec::CommunicationValue;
use mtp_codec::{CommunicationValue, EncodeLimits};
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
@ -12,7 +26,11 @@ pub(crate) async fn write_frame<S: TransportSendStream>(
value: &CommunicationValue,
policy: &Policy,
) -> Result<(), CommunicationError> {
let bytes = value.to_bytes().map_err(|_| CommunicationError::Encode)?;
let bytes = value
.to_bytes_with_limits(EncodeLimits::for_transport_message_size(
policy.max_message_size,
))
.map_err(|_| CommunicationError::Encode)?;
if bytes.len() as u64 > policy.max_message_size
|| bytes.len() as u64 >= policy.close_frame_len as u64
{
@ -64,6 +82,10 @@ mod tests {
async fn finish(&mut self) -> Result<(), CommunicationError> {
Ok(())
}
fn reset(&mut self, _code: u32) -> Result<(), CommunicationError> {
Ok(())
}
}
#[tokio::test]

View file

@ -5,21 +5,27 @@
//! wrappers while the framing implementation below is shared by adapters.
use crate::{
Policy, TransportConnection, TransportRecvStream, TransportSendStream, framing::write_frame,
Policy, TransportConnection, TransportRecvStream, TransportSendStream,
connection::{DecodeRejectionCounters, RuntimePolicy, classify_decode_error},
framing::{RetryClassifier, write_frame},
};
use mtp_codec::{CommunicationValue, DecodeLimits, TypeMap};
use mtp_common::CommunicationError;
use mtp_codec::{CommunicationValue, DataType, DecodeLimits, TypeMap};
use mtp_common::{CommunicationError, FirstFrameDisposition, classify_first_frame};
#[cfg(feature = "pipes")]
use std::collections::HashSet;
use std::sync::Arc;
#[cfg(feature = "pipes")]
use std::sync::Mutex as StdMutex;
use std::sync::atomic::{AtomicU64, Ordering};
use tokio::sync::{Mutex, RwLock, Semaphore, mpsc};
use tokio::time::timeout;
use tokio::sync::{Mutex, Notify, RwLock, Semaphore, mpsc};
use tokio::time::{Instant, timeout, timeout_at};
#[cfg(feature = "pipes")]
use crate::pipe::{PipeReader, PipeWriter};
pub struct GenericSender<C: TransportConnection> {
connection: C,
policy: Arc<Policy>,
policy: Arc<RuntimePolicy>,
persistent: Arc<Mutex<Option<C::SendStream>>>,
send_lock: Arc<Mutex<()>>,
type_map: Arc<RwLock<TypeMap>>,
@ -39,6 +45,7 @@ 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,
@ -64,6 +71,15 @@ 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?;
@ -73,9 +89,10 @@ impl<C: TransportConnection> GenericSender<C> {
)
.await
.map_err(|_| CommunicationError::StreamError)??;
timeout(self.policy.write_timeout, stream.finish())
.await
.map_err(|_| CommunicationError::StreamError)?
match timeout(self.policy.write_timeout, stream.finish()).await {
Ok(Ok(())) => Ok(()),
Ok(Err(_)) | Err(_) => Err(CommunicationError::DeliveryUnknown),
}
}
crate::SendMode::PersistentStream => {
let mut stream = self.persistent.lock().await;
@ -84,20 +101,30 @@ impl<C: TransportConnection> GenericSender<C> {
if stream.is_none() {
*stream = Some(self.open().await?);
}
let result = timeout(
let result = match stream.as_mut() {
Some(stream) => timeout(
self.policy.write_timeout,
write_frame(stream.as_mut().unwrap(), value, &self.policy),
write_frame(stream, value, &self.policy),
)
.await
.map_err(|_| CommunicationError::StreamError)
.and_then(|r| r);
.and_then(|result| result),
None => Err(CommunicationError::StreamError),
};
if result.is_ok() {
return result;
return Ok(());
}
let error = match result {
Ok(()) => return Ok(()),
Err(error) => error,
};
if !RetryClassifier::retry_persistent_stream(&error) {
return Err(error);
}
*stream = None;
attempts += 1;
if attempts > self.policy.persistent_stream_max_retries {
return result;
return Err(error);
}
tokio::time::sleep(
self.policy.persistent_stream_retry_backoff * attempts as u32,
@ -114,6 +141,7 @@ 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);
}
@ -179,6 +207,10 @@ 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<()>>,
}
@ -192,6 +224,10 @@ 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(),
}
}
@ -207,6 +243,7 @@ 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);
@ -222,6 +259,14 @@ 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();
@ -237,8 +282,10 @@ 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 {
tokio::time::sleep(std::time::Duration::from_millis(1)).await;
notified.await;
continue;
}
@ -277,11 +324,14 @@ 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;
loop {
'stream: loop {
if policy
.max_frames_per_stream
.is_some_and(|max| frames >= max)
@ -304,9 +354,21 @@ impl<C: TransportConnection> GenericReceiver<C> {
policy.application_close_code,
b"frame header read error",
);
break;
break 'stream;
}
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;
}
}
@ -314,6 +376,7 @@ 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) {
@ -330,29 +393,28 @@ impl<C: TransportConnection> GenericReceiver<C> {
connection.close(policy.application_close_code, b"frame too large");
break;
}
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 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 _ = tx.send(Err(CommunicationError::MessageTooLarge)).await;
connection
.close(policy.application_close_code, b"frame allocation failed");
break;
}
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]),
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]),
)
.await;
if !matches!(&body_read, Ok(Ok(())))
|| body.try_reserve(chunk_len).is_err()
{
if !matches!(&body_read, Ok(Ok(()))) {
if matches!(&body_read, Ok(Err(CommunicationError::StreamClosed))) {
break 'stream;
}
tracing::warn!(
pipe_chunk_len = chunk_len,
?body_read,
@ -361,24 +423,23 @@ impl<C: TransportConnection> GenericReceiver<C> {
let _ = tx.send(Err(CommunicationError::StreamError)).await;
connection
.close(policy.application_close_code, b"frame body read error");
break;
break 'stream;
}
body.extend_from_slice(&chunk[..chunk_len]);
}
if body.len() != target_len {
break;
body_offset += chunk_len;
}
frames += 1;
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(
let mut message = match CommunicationValue::try_from_bytes_with_limits(
&frame,
DecodeLimits::for_transport_message_size(frame_limit),
) {
Ok(message) => message,
Err(_) => {
tracing::warn!("MTP receive stream contained an invalid frame");
Err(error) => {
tracing::warn!(
?error,
class = ?classify_decode_error(&error),
"MTP receive stream rejected by bounded decode"
);
decode_rejections.record(&error);
let _ = tx
.send(Err(CommunicationError::ParseCommunicationValue))
.await;
@ -386,25 +447,44 @@ 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 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(),
);
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) => {
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("")
@ -424,6 +504,7 @@ impl<C: TransportConnection> GenericReceiver<C> {
return;
}
}
}
if message.is_type(mtp_codec::CommunicationType::Ping) {
if let Some(sender) = ping_sender.read().await.clone() {
@ -463,6 +544,10 @@ 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),
}
}
@ -470,6 +555,25 @@ 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
@ -480,13 +584,23 @@ 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> {
self.incoming
let result = self
.incoming
.lock()
.await
.recv()
.await
.unwrap_or(Err(CommunicationError::StreamClosed))
.unwrap_or(Err(CommunicationError::StreamClosed));
if result.is_ok() {
self.queue_notify.notify_one();
}
result
}
#[cfg(feature = "pipes")]
@ -498,7 +612,10 @@ impl<C: TransportConnection> GenericReceiver<C> {
tokio::select! {
msg = incoming.recv() => {
match msg {
Some(Ok(val)) => Ok(crate::TransportEvent::Message(val)),
Some(Ok(val)) => {
self.queue_notify.notify_one();
Ok(crate::TransportEvent::Message(val))
}
Some(Err(e)) => Err(e),
None => Err(self
.connection
@ -508,7 +625,10 @@ impl<C: TransportConnection> GenericReceiver<C> {
}
pipe = pipes.recv() => {
match pipe {
Some(reader) => Ok(crate::TransportEvent::Pipe(reader)),
Some(reader) => {
self.queue_notify.notify_one();
Ok(crate::TransportEvent::Pipe(reader))
}
None => Err(self
.connection
.close_reason()
@ -520,12 +640,17 @@ impl<C: TransportConnection> GenericReceiver<C> {
#[cfg(feature = "pipes")]
pub async fn receive_pipe(&self) -> Result<PipeReader<C::RecvStream>, CommunicationError> {
self.pipes
let result = self
.pipes
.lock()
.await
.recv()
.await
.ok_or(CommunicationError::StreamClosed)
.ok_or(CommunicationError::StreamClosed);
if result.is_ok() {
self.queue_notify.notify_one();
}
result
}
#[cfg(feature = "pipes")]
@ -534,7 +659,10 @@ 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) => Ok(Some(reader)),
Ok(reader) => {
self.queue_notify.notify_one();
Ok(Some(reader))
}
Err(mpsc::error::TryRecvError::Empty) => Ok(None),
Err(mpsc::error::TryRecvError::Disconnected) => {
Err(CommunicationError::StreamClosed)

View file

@ -11,7 +11,10 @@ pub mod encrypted_pipe;
#[cfg(feature = "pipes")]
pub mod pipe;
pub use connection::{Policy, Receiver, SendMode, Sender};
pub use connection::{
DecodeRejectionClass, DecodeRejectionCounters, DecodeRejectionCounts, Policy, Receiver,
SendMode, Sender, classify_decode_error,
};
pub use generic_connection::{GenericReceiver, GenericSender};
#[cfg(feature = "pipes")]

View file

@ -17,6 +17,7 @@ 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.
@ -28,6 +29,9 @@ 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.
@ -47,7 +51,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::StreamError)
.map_err(|_| CommunicationError::DeliveryUnknown)
}
async fn finish(&mut self) -> Result<(), CommunicationError> {
@ -55,14 +59,23 @@ 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> {
wtransport::RecvStream::read_exact(self, buf)
.await
.map_err(|_| CommunicationError::StreamError)
match wtransport::RecvStream::read_exact(self, buf).await {
Ok(()) => Ok(()),
Err(wtransport::error::StreamReadExactError::FinishedEarly(0)) => {
Err(CommunicationError::StreamClosed)
}
Err(_) => Err(CommunicationError::StreamError),
}
}
async fn read_chunk(&mut self, max: usize) -> Result<Option<Vec<u8>>, CommunicationError> {
@ -76,6 +89,11 @@ 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]

View file

@ -4,7 +4,7 @@ use async_trait::async_trait;
use mtp_codec::CommunicationValue;
use mtp_common::CommunicationError;
use mtp_transport::{
GenericReceiver, GenericSender, Policy, TransportConnection, TransportEvent,
GenericReceiver, GenericSender, Policy, SendMode, TransportConnection, TransportEvent,
TransportRecvStream, TransportSendStream,
};
use std::sync::Arc;
@ -53,6 +53,10 @@ impl TransportSendStream for MockSendStream {
.await
.map_err(|_| CommunicationError::StreamError)
}
fn reset(&mut self, _code: u32) -> Result<(), CommunicationError> {
Ok(())
}
}
struct MockRecvStream {
@ -89,6 +93,10 @@ impl TransportRecvStream for MockRecvStream {
Err(_) => Err(CommunicationError::StreamError),
}
}
fn stop(self, _code: u32) -> Result<(), CommunicationError> {
Ok(())
}
}
#[derive(Clone)]
@ -157,6 +165,7 @@ 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?;
@ -174,6 +183,7 @@ 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";
@ -195,6 +205,7 @@ 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();
@ -223,6 +234,7 @@ 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? {
@ -261,22 +273,23 @@ 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());
let policy = Arc::new(Policy::default().with_send_mode(SendMode::SingleStreamPerMessage));
let sender = GenericSender::new(conn_a, policy.clone());
let receiver = GenericReceiver::new(conn_b, policy);
let msg = CommunicationValue::new(mtp_codec::CommunicationType::Pong);
sender.send(&msg).await?;
let _pipe_writer = sender.open_pipe(1, "mixed-pipe").await?;
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 received = receiver.receive().await?;
assert_eq!(
received.get_type(),
mtp_codec::CommunicationType::Pong
.try_to_id(&mtp_codec::TypeMap::latest())
.unwrap()
);
assert!(received.is_type(mtp_codec::CommunicationType::PipeRequest));
receiver.expect_pipe(1)?;
let _pipe_writer = sender.open_pipe(1, "mixed-pipe").await?;
let pipe_reader = receiver.receive_pipe().await?;
assert_eq!(pipe_reader.pipe_id(), 1);
@ -291,6 +304,8 @@ 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?;

View file

@ -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::Pong, 99, &tm);
let resp = numbered_message(CommunicationType::BadRequest, 99, &tm);
host_tx.send(&resp).await?;
// Client receives it
let client_received = client_rx.receive().await?;
assert_numbered_message(&client_received, CommunicationType::Pong, 99, &tm);
assert_numbered_message(&client_received, CommunicationType::BadRequest, 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::Pong, i * 10, &tm);
let msg = numbered_message(CommunicationType::BadRequest, i * 10, &tm);
client_tx.send(&msg).await?;
}
for i in 0..3u128 {
let received = host_rx.receive().await?;
assert_numbered_message(&received, CommunicationType::Pong, i * 10, &tm);
assert_numbered_message(&received, CommunicationType::BadRequest, 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::Pong, 7, &tm);
let resp = numbered_message(CommunicationType::BadRequest, 7, &tm);
host_tx.send(&resp).await?;
let got = client_rx.receive().await?;
assert_numbered_message(&got, CommunicationType::Pong, 7, &tm);
assert_numbered_message(&got, CommunicationType::BadRequest, 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::Pong, 22, &tm);
let msg2 = numbered_message(CommunicationType::BadRequest, 22, &tm);
client_tx.send(&msg2).await?;
let received2 = host_rx.receive().await?;
assert_numbered_message(&received2, CommunicationType::Pong, 22, &tm);
assert_numbered_message(&received2, CommunicationType::BadRequest, 22, &tm);
client_tx.close().await;
host_tx.close().await;
@ -342,13 +342,17 @@ async fn test_max_frames_per_stream_enforced() -> Result<(), Box<dyn std::error:
let tm = TypeMap::latest();
client_tx
.send(&numbered_message(CommunicationType::Ping, 1, &tm))
.await?;
let first = host_rx.receive().await?;
.await
.expect("first frame should be sent");
let first = host_rx
.receive()
.await
.expect("first frame should be received");
assert_numbered_message(&first, CommunicationType::Ping, 1, &tm);
client_tx
let _ = client_tx
.send(&numbered_message(CommunicationType::Ping, 2, &tm))
.await?;
.await;
let second = host_rx.receive().await;
assert!(second.is_err(), "stream should be closed after frame limit");

View file

@ -16,3 +16,7 @@ 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"]

View file

@ -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"] }
mtp-crypto = { version = "0.3.0", path = "../crypto", features = ["wasm", "password-kdf"] }
zeroize = "1.9"
wasm-bindgen-test = "0.3.76"

File diff suppressed because it is too large Load diff

View file

@ -0,0 +1,680 @@
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")
);
}
}

View file

@ -0,0 +1,102 @@
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
}
}

184
wasm/src/client/dispatch.rs Normal file
View file

@ -0,0 +1,184 @@
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();
}
}

299
wasm/src/client/mod.rs Normal file
View file

@ -0,0 +1,299 @@
// 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()
}
}

76
wasm/src/client/pipes.rs Normal file
View file

@ -0,0 +1,76 @@
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
}
}

335
wasm/src/client/receive.rs Normal file
View file

@ -0,0 +1,335 @@
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();
}
}

View file

@ -3,8 +3,10 @@ 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};
@ -21,12 +23,17 @@ 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>>,
@ -114,6 +121,10 @@ pub struct WasmPipeHandle {
description: String,
transport: WasmTransport,
response_rx: PipeResponseCell,
pending: PendingPipeCreations,
expired: Rc<RefCell<HashMap<u32, f64>>>,
generation: u32,
token: Rc<()>,
}
#[wasm_bindgen]
@ -125,9 +136,37 @@ impl WasmPipeHandle {
.take()
.ok_or_else(|| js_error("handle already consumed"))?;
let accepted = rx
.await
.map_err(|_| js_error("pipe handle channel closed"))?;
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"
)));
},
};
match accepted {
Ok(true) => {
@ -153,6 +192,18 @@ 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"))?;
@ -166,6 +217,146 @@ 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 {
@ -178,19 +369,26 @@ 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);
let occupied = pipe_id == 0
|| pending_pipe_creations.borrow().contains_key(&pipe_id)
|| is_expired_pipe_creation(expired_pipe_creations, pipe_id);
if !occupied {
break;
}
pipe_id = random_pipe_id()?;
}
if pipe_id == 0 || pending_pipe_creations.borrow().contains_key(&pipe_id) {
if pipe_id == 0
|| pending_pipe_creations.borrow().contains_key(&pipe_id)
|| is_expired_pipe_creation(expired_pipe_creations, pipe_id)
{
return Err(js_error("could not allocate a unique pipe id"));
}
let type_map = transport.type_map();
@ -207,9 +405,17 @@ 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,
@ -218,31 +424,22 @@ 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,
})
}
@ -282,6 +479,7 @@ pub(crate) async fn wasm_accept_pipe(
},
);
}
let _acceptance_guard = PendingPipeGuard::new(pending_pipes.clone(), pipe_id, generation);
debug!(
target = "mtp.wasm",
@ -291,16 +489,35 @@ pub(crate) async fn wasm_accept_pipe(
"sending pipe response"
);
if let Err(error) = transport.send_frame(&resp_bytes).await {
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(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)
@ -308,11 +525,6 @@ pub(crate) async fn wasm_accept_pipe(
{
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> {

View file

@ -1,5 +1,6 @@
use wasm_bindgen::prelude::*;
#[derive(Clone)]
#[wasm_bindgen]
pub struct ConnectionConfig {
pub(crate) url: String,

View file

@ -2,8 +2,8 @@ use wasm_bindgen::prelude::*;
use zeroize::Zeroizing;
use mtp_codec::{
DataValue, MtpProtectionPurpose, PROTOCOL_VERSION, ProtectionPolicy, ProtectionPurpose,
SealedRelayBuilder, SignaturePolicy, TypeMap,
DataValue, DecodeLimits, EncodeLimits, MtpProtectionPurpose, PROTOCOL_VERSION,
ProtectionPolicy, ProtectionPurpose, SealedRelayBuilder, SignaturePolicy, TypeMap,
};
use mtp_crypto::{
AeadDecrypt, AeadEncrypt, DualSigner, Ed25519Signer, HybridKem, KemPrivateKey, KemPublicKey,
@ -12,10 +12,18 @@ use mtp_crypto::{
};
use crate::error::{from_protection_error, js_error};
use crate::relay::{decode_frame, relay_error, structured_error};
use crate::relay::{decode_error, decode_frame, relay_error, structured_error};
fn decode_data_value(value: &[u8]) -> Result<DataValue, JsValue> {
DataValue::from_bytes(value).ok_or_else(|| js_error("invalid DataValue"))
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
})
}
fn decode_public_key_bundle(
@ -74,10 +82,19 @@ pub struct WasmKeyring {
#[wasm_bindgen]
impl WasmKeyring {
/// Serialise the keyring to bytes.
/// Serialise the keyring to bytes and report malformed caller-owned
/// material as a JavaScript exception.
#[wasm_bindgen]
pub fn to_bytes(&self) -> Vec<u8> {
self.inner.to_bytes().to_vec()
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}")))
}
/// Deserialise a keyring from bytes.
@ -118,8 +135,17 @@ impl WasmKeyring {
/// Generate a full keyring with KEM, ML-DSA, and Ed25519 keys.
#[wasm_bindgen]
pub fn keyring_generate() -> Vec<u8> {
Keyring::generate().to_bytes().to_vec()
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}")))
}
/// Build a [`Keyring`] containing only an Ed25519 keypair (no KEM, no ML-DSA).
@ -142,7 +168,10 @@ 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()),
);
Ok(keyring.to_bytes().to_vec())
keyring
.try_to_bytes()
.map(|bytes| bytes.to_vec())
.map_err(|error| js_error(format!("Keyring serialization failed: {error}")))
}
// ===========================================================================
@ -172,8 +201,15 @@ impl WasmPublicKeyBundle {
}
#[wasm_bindgen]
pub fn to_bytes(&self) -> Vec<u8> {
self.inner.as_bytes()
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}")))
}
#[wasm_bindgen]
@ -428,6 +464,14 @@ 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(
@ -452,6 +496,21 @@ 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;
@ -564,10 +623,11 @@ 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(
value.verify_with_policy(
expected_signer_id,
&bundle,
ProtectionPurpose::from(expected_purpose),
ProtectionPolicy::any_supported(),
)
} else {
value.verify_with_policy(
@ -650,7 +710,11 @@ pub fn decrypt_data_value_with_keyrings(
let keyrings = keyrings_from_js(&keyrings)?;
let references: Vec<&Keyring> = keyrings.iter().collect();
value
.decrypt_with_keyrings(&references, ProtectionPurpose::from(expected_purpose))
.decrypt_with_keyrings_and_limits(
&references,
ProtectionPurpose::from(expected_purpose),
DecodeLimits::default(),
)
.map_err(from_protection_error)?
.to_bytes()
.map_err(|e| js_error(format!("decryption failed: {e}")))
@ -746,6 +810,13 @@ 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]
@ -776,12 +847,27 @@ 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 application_content = crate::frame::js_to_data_value(&data, &tm)?;
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_metadata = encoded_metadata
.as_deref()
.map(decode_data_value)
.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"))
})
.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)?;
@ -798,6 +884,8 @@ 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),
@ -807,7 +895,7 @@ fn build_encrypted_relay_frame_impl(
builder
.build()
.map_err(relay_error)?
.to_bytes()
.to_bytes_with_limits(encode_limits)
.map_err(|e| js_error(format!("relay frame encoding failed: {e}")))
}
@ -844,6 +932,45 @@ 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,
)
}
@ -891,7 +1018,7 @@ mod tests {
},
};
let bytes = bundle.to_bytes();
let bytes = bundle.try_to_bytes().expect("bundle serialization");
let restored = WasmPublicKeyBundle::from_bytes_unvalidated(&bytes)
.expect("from_bytes_unvalidated failed");
assert_eq!(restored.sig_cl_public_key(), pk);
@ -1098,7 +1225,7 @@ mod tests {
let value = DataValue::Str("signed through wasm".into())
.to_bytes()
.expect("value encoding failed");
let keyring_bytes = keyring.to_bytes();
let keyring_bytes = keyring.try_to_bytes().expect("keyring serialization");
let signed = sign_data_value_with_keyring(
&value,
0xfeed_beef,
@ -1111,7 +1238,7 @@ mod tests {
verify_data_value_with_policy(
&signed,
&bundle.as_bytes(),
&bundle.try_as_bytes().expect("bundle serialization"),
0xfeed_beef,
7,
PROTECTION_SIGNATURE_SUITE_ED25519,
@ -1121,7 +1248,7 @@ mod tests {
assert!(
verify_data_value_with_policy(
&signed,
&wrong_bundle.as_bytes(),
&wrong_bundle.try_as_bytes().expect("bundle serialization"),
0xfeed_beef,
7,
PROTECTION_SIGNATURE_SUITE_ED25519,
@ -1137,21 +1264,29 @@ mod tests {
let value = DataValue::Array(vec![DataValue::BoolTrue, DataValue::UnsignedNumber(42)])
.to_bytes()
.expect("value encoding 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");
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");
assert_eq!(decrypted, value);
let second_keyring = Keyring::generate();
let second_recipient = second_keyring.public_key_bundle();
let recipients = js_sys::Array::new();
recipients.push(&js_sys::Uint8Array::from(&recipient.as_bytes()[..]));
recipients.push(&js_sys::Uint8Array::from(&second_recipient.as_bytes()[..]));
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[..]));
let multi = encrypt_data_value_for_recipients(&value, recipients.into(), 9)
.expect("multi-recipient encryption failed");
let opened_by_second = decrypt_data_value(&multi, &second_keyring.to_bytes(), 9)
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)
.expect("second recipient could not decrypt");
assert_eq!(opened_by_second, value);
}

View file

@ -1,9 +1,13 @@
use wasm_bindgen::{JsCast, prelude::*};
use mtp_codec::{CommunicationType, CommunicationValue, DataType, DataValue, PROTOCOL_VERSION};
use mtp_codec::{
CommunicationType, CommunicationValue, DataType, DataValue, DecodeLimits, EncodeLimits,
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#"
@ -137,7 +141,49 @@ 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);
}
@ -152,9 +198,17 @@ pub(crate) fn js_to_data_value(value: &JsValue, tm: &TypeMap) -> Result<DataValu
}
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(&item, tm)?);
values.push(js_to_data_value_with_context(
&item,
tm,
context,
depth + 1,
)?);
}
return Ok(DataValue::Array(values));
}
@ -198,6 +252,9 @@ pub(crate) fn js_to_data_value(value: &JsValue, tm: &TypeMap) -> Result<DataValu
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
@ -212,7 +269,10 @@ pub(crate) fn js_to_data_value(value: &JsValue, tm: &TypeMap) -> Result<DataValu
tm.version
))
})?;
entries.push((id, js_to_data_value(&value, tm)?));
entries.push((
id,
js_to_data_value_with_context(&value, tm, context, depth + 1)?,
));
}
return Ok(DataValue::Container(entries));
}
@ -282,15 +342,16 @@ fn apply_frame_options(
}
pub(crate) fn parse_frame_value(frame: &[u8]) -> Result<JsValue, JsValue> {
parse_frame_value_with_type_map(frame, &TypeMap::latest())
parse_frame_value_with_limits(frame, &TypeMap::latest(), DecodeLimits::default())
}
pub(crate) fn parse_frame_value_with_type_map(
pub(crate) fn parse_frame_value_with_limits(
frame: &[u8],
type_map: &TypeMap,
limits: DecodeLimits,
) -> Result<JsValue, JsValue> {
let comm = CommunicationValue::from_bytes_with(frame, type_map)
.map_err(|e| js_error(format!("parse failed: {}", e)))?;
let comm = CommunicationValue::try_from_bytes_with_type_map_and_limits(frame, type_map, limits)
.map_err(|error| decode_error(error, "parse failed"))?;
let tm = type_map;
let obj = js_sys::Object::new();
@ -355,8 +416,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::from_bytes(response)
.map_err(|e| js_error(format!("parse failed: {}", e)))?;
let comm = CommunicationValue::try_from_bytes_with_limits(response, DecodeLimits::default())
.map_err(|error| decode_error(error, "parse failed"))?;
let connected = matches!(
comm.get_data(DataType::Connected),
@ -418,8 +479,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::from_bytes(frame)
.map_err(|e| js_error(format!("parse failed: {}", e)))?;
let comm = CommunicationValue::try_from_bytes_with_limits(frame, DecodeLimits::default())
.map_err(|error| decode_error(error, "parse failed"))?;
Ok(comm.to_string())
}
@ -429,31 +490,80 @@ 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> {
let value = DataValue::from_bytes(value).ok_or_else(|| js_error("invalid DataValue"))?;
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 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(&value, &tm)?
.to_bytes()
js_to_data_value_with_limits(&value, &tm, limits)?
.to_bytes_with_limits(limits)
.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}")))?;
@ -478,7 +588,7 @@ pub fn build_frame(
))
})?;
msg = msg
.add_data(id, js_to_data_value(&value, &tm)?)
.add_data(id, js_to_data_value_with_limits(&value, &tm, limits)?)
.map_err(|e| js_error(format!("add data failed: {e}")))?;
}
} else if !data.is_null() && !data.is_undefined() {
@ -487,10 +597,24 @@ pub fn build_frame(
));
}
msg.to_bytes()
msg.to_bytes_with_limits(limits)
.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
@ -501,19 +625,50 @@ 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::from_bytes(serialized_payload)
.ok_or_else(|| js_error("invalid serialized DataValue payload"))?;
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 message =
apply_frame_options(CommunicationValue::new(comm_type), &options)?.with_payload(payload);
message
.to_bytes()
.to_bytes_with_limits(limits)
.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 {

Some files were not shown because too many files have changed in this diff Show more