diff --git a/client/src/lib.rs b/client/src/lib.rs index 0df6e7e..24ed84a 100644 --- a/client/src/lib.rs +++ b/client/src/lib.rs @@ -4,6 +4,20 @@ use mtp_codec::{CommunicationValue, DataTypeId, DataValue, PROTOCOL_VERSION, Ver use mtp_common::CommunicationError; use mtp_transport::{Policy, Receiver, Sender}; +#[cfg(feature = "crypto")] +fn unexpected_response_type_error( + context: &str, + expected_type: mtp_codec::CommunicationTypeId, + response: &CommunicationValue, +) -> CommunicationError { + CommunicationError::AuthenticationFailed(format!( + "unexpected response type during {context}: expected {:?}, got {:?}; parsed {}", + expected_type, + response.get_type(), + response + )) +} + pub struct ClientConfig { pub url: String, pub server_cert: Option>, @@ -124,6 +138,15 @@ impl MTPClient { // 2. Receive host response (single message) let response = receiver.receive().await?; + let expected_type = mtp_codec::CommunicationTypeId(16); // IdentificationResponse + if response.get_type() != expected_type { + return Err(unexpected_response_type_error( + "auth_connect", + expected_type, + &response, + )); + } + let connected = response.get_data(DataTypeId(11)); match connected { DataValue::BoolTrue => {} @@ -264,6 +287,15 @@ impl MTPClient { // 2. Receive host response (single message) let response = receiver.receive().await?; + let expected_type = mtp_codec::CommunicationTypeId(18); // RegisterResponse + if response.get_type() != expected_type { + return Err(unexpected_response_type_error( + "auth_register", + expected_type, + &response, + )); + } + let connected = response.get_data(DataTypeId(11)); match connected { DataValue::BoolTrue => {} diff --git a/example-usage/client/src/main.rs b/example-usage/client/src/main.rs index 5f3c9b4..d838685 100644 --- a/example-usage/client/src/main.rs +++ b/example-usage/client/src/main.rs @@ -2,13 +2,29 @@ mod auth; mod messages; use std::fs; +use std::path::Path; use mtp::client::ClientConfig; use mtp::crypto::{KemPublicKey, PublicKeyBundle, SignaturePqPublicKey, SignaturePublicKey}; +fn dev_cert_path() -> String { + std::env::var("MTP_DEV_CERT").unwrap_or_else(|_| { + if Path::new("example-usage/dev-cert/cert.pem").exists() { + "example-usage/dev-cert/cert.pem".to_string() + } else { + "dev-cert/cert.pem".to_string() + } + }) +} + #[tokio::main] async fn main() -> Result<(), Box> { - let cert_pem = fs::read("server.pem").expect("Missing server.pem: run server first"); + let cert_path = dev_cert_path(); + let cert_pem = fs::read(&cert_path).unwrap_or_else(|e| { + panic!( + "Missing TLS certificate at {cert_path}: enter the Nix shell first or run the server to generate it: {e}" + ) + }); let host_public_key = PublicKeyBundle::new( KemPublicKey::new( fs::read("host_enc_kem_pk.bin") diff --git a/example-usage/client/src/messages.rs b/example-usage/client/src/messages.rs index 1c17e5a..e69ef15 100644 --- a/example-usage/client/src/messages.rs +++ b/example-usage/client/src/messages.rs @@ -13,8 +13,8 @@ fn derive_demo_key() -> [u8; 32] { pub fn build_demo_message(client_id: u64, keyring: &Keyring) -> CommunicationValue { let cipher = ChaCha20Poly1305::new(derive_demo_key()); - let signer = Ed25519Signer::new(&keyring.sig_cl_secret_key) - .expect("Ed25519 signer from keyring"); + let signer = + Ed25519Signer::new(&keyring.sig_cl_secret_key).expect("Ed25519 signer from keyring"); let inner_enc = DataValue::Container(vec![ (DataTypeId(1), DataValue::Str("secret inner data".into())), @@ -31,7 +31,10 @@ pub fn build_demo_message(client_id: u64, keyring: &Keyring) -> CommunicationVal dv_sig.sign_container(SigAlgorithm::ED25519, &signer); let inner_sec = DataValue::Container(vec![ - (DataTypeId(1), DataValue::Str("signed+encrypted payload".into())), + ( + DataTypeId(1), + DataValue::Str("signed+encrypted payload".into()), + ), (DataTypeId(2), DataValue::UnsignedNumber(7)), ]); let mut dv_sec = inner_sec; @@ -42,13 +45,22 @@ pub fn build_demo_message(client_id: u64, keyring: &Keyring) -> CommunicationVal .unwrap() .as_secs(); - CommunicationValue::new(CommunicationType::Ping) - .add_typed_default(DataType::Description, DataValue::Str("MTP Data Type Demo".into())) - .add_typed_default(DataType::Timestamp, DataValue::UnsignedNumber(timestamp as u128)) + let msg = CommunicationValue::new(CommunicationType::Ping) + .add_typed_default( + DataType::Description, + DataValue::Str("MTP Data Type Demo".into()), + ) + .add_typed_default( + DataType::Timestamp, + DataValue::UnsignedNumber(timestamp as u128), + ) .add_typed_default(DataType::Data, DataValue::Str("Hello, MTP!".into())) .add_typed_default(DataType::Flags, DataValue::BoolTrue) .add_typed_default(DataType::Value, DataValue::Float(2, 12345)) - .add_typed_default(DataType::BinaryData, DataValue::Bytes(vec![0xDE, 0xAD, 0xBE, 0xEF, 0x42])) + .add_typed_default( + DataType::BinaryData, + DataValue::Bytes(vec![0xDE, 0xAD, 0xBE, 0xEF, 0x42]), + ) .add_typed_default( DataType::Items, DataValue::Array(vec![ @@ -60,7 +72,8 @@ pub fn build_demo_message(client_id: u64, keyring: &Keyring) -> CommunicationVal .add_typed_default(DataType::EncryptedPayload, dv_enc) .add_typed_default(DataType::SignedPayload, dv_sig) .add_typed_default(DataType::SecurePayload, dv_sec) - .with_sender(client_id) + .with_sender(client_id); + msg } pub async fn send_and_receive( @@ -72,7 +85,9 @@ pub async fn send_and_receive( conn.sender.send(&msg).await?; match conn.receiver.receive().await { - Ok(resp) => println!("Received: {resp}"), + Ok(resp) => { + println!("Received: {resp}"); + } Err(e) => eprintln!("Receive error: {e}"), } diff --git a/example-usage/server/src/clients.rs b/example-usage/server/src/clients.rs index 3417604..d9ad714 100644 --- a/example-usage/server/src/clients.rs +++ b/example-usage/server/src/clients.rs @@ -8,14 +8,18 @@ pub fn load_client_db( path: &str, ) -> Result<(Arc>>, Arc>), Box> { - let clients: Arc>> = - Arc::new(Mutex::new(if let Ok(data) = fs::read_to_string(path) { - serde_json::from_str(&data).unwrap_or_default() - } else { - HashMap::new() - })); - let next_id = Arc::new(Mutex::new( - clients.lock().unwrap().keys().max().unwrap_or(&999) + 1, - )); + let clients_map = match fs::read_to_string(path) { + Ok(data) => match serde_json::from_str(&data) { + Ok(clients) => clients, + Err(e) => { + eprintln!("Failed to parse {path}; starting with empty client database: {e}"); + HashMap::new() + } + }, + Err(_) => HashMap::new(), + }; + let clients: Arc>> = Arc::new(Mutex::new(clients_map)); + let next_value = clients.lock().unwrap().keys().max().unwrap_or(&999) + 1; + let next_id = Arc::new(Mutex::new(next_value)); Ok((clients, next_id)) } diff --git a/example-usage/server/src/handlers.rs b/example-usage/server/src/handlers.rs index 768293b..ff542c0 100644 --- a/example-usage/server/src/handlers.rs +++ b/example-usage/server/src/handlers.rs @@ -1,5 +1,7 @@ use mtp::codec::{CommunicationType, CommunicationValue, DataType, DataTypeId, DataValue, TypeMap}; -use mtp::crypto::{ChaCha20Poly1305, CryptoError, SignatureScheme, SignaturePublicKey, verify_ed25519}; +use mtp::crypto::{ + ChaCha20Poly1305, CryptoError, SignaturePublicKey, SignatureScheme, verify_ed25519, +}; struct Ed25519Verifier(SignaturePublicKey); @@ -71,6 +73,7 @@ pub fn process_and_respond( enc_status = format!("EncryptedPayload decrypted OK ({} entries)", entries.len()); } } else { + eprintln!(" EncryptedPayload decryption failed"); enc_status = String::from("EncryptedPayload: decryption FAILED"); } } @@ -83,13 +86,14 @@ pub fn process_and_respond( if dv.verify_into_container(&verifier).is_some() { if let Some(entries) = dv.as_container() { println!(" Verified SignedPayload: {:?}", entries); - sig_status = - format!("SignedPayload verified OK ({} entries)", entries.len()); + sig_status = format!("SignedPayload verified OK ({} entries)", entries.len()); } } else { + eprintln!(" SignedPayload verification failed"); sig_status = String::from("SignedPayload: verification FAILED"); } } else { + eprintln!(" SignedPayload cannot be verified; no client public key available"); sig_status = String::from("SignedPayload: no client public key available"); } } @@ -99,7 +103,9 @@ pub fn process_and_respond( if let Some(pk_bundle) = client_pk { let verifier = Ed25519Verifier(pk_bundle.sig_cl_public_key.clone()); let mut dv = secure.clone(); - if dv.decrypt_signed_encrypted_container(&cipher, b"demo-aad").is_some() + if dv + .decrypt_signed_encrypted_container(&cipher, b"demo-aad") + .is_some() && dv.verify_into_container(&verifier).is_some() { if let Some(entries) = dv.as_container() { @@ -110,9 +116,11 @@ pub fn process_and_respond( ); } } else { + eprintln!(" SecurePayload decryption/verification failed"); secure_status = String::from("SecurePayload: decryption/verification FAILED"); } } else { + eprintln!(" SecurePayload cannot be verified; no client public key available"); secure_status = String::from("SecurePayload: no client public key available"); } } diff --git a/example-usage/server/src/keys.rs b/example-usage/server/src/keys.rs index fdebfd6..133eeba 100644 --- a/example-usage/server/src/keys.rs +++ b/example-usage/server/src/keys.rs @@ -1,9 +1,11 @@ use std::fs; -use mtp::crypto::{Ed25519Signer, Keyring, MlDsaSigner}; use mtp::crypto::kem::HybridKem; +use mtp::crypto::{Ed25519Signer, Keyring, MlDsaSigner}; -pub fn load_or_generate_host_keys(path: &str) -> Result<(u64, Keyring), Box> { +pub fn load_or_generate_host_keys( + path: &str, +) -> Result<(u64, Keyring), Box> { if let Ok(data) = fs::read_to_string(path) { let json: serde_json::Value = serde_json::from_str(&data)?; let hid = json["host_id"].as_u64().unwrap_or(1); @@ -29,10 +31,7 @@ pub fn load_or_generate_host_keys(path: &str) -> Result<(u64, Keyring), Box Result<(), Box> { let public_key_bundle_hex = hex::encode(host_keyring.public_key_bundle().as_bytes()); - fs::write( - "host_public_key_bundle.hex", - &public_key_bundle_hex, - )?; + fs::write("host_public_key_bundle.hex", &public_key_bundle_hex)?; fs::create_dir_all("web-client/public")?; fs::write( "web-client/public/host_public_key_bundle.hex", diff --git a/example-usage/server/src/main.rs b/example-usage/server/src/main.rs index 4b0bd5d..8cf949e 100644 --- a/example-usage/server/src/main.rs +++ b/example-usage/server/src/main.rs @@ -40,7 +40,13 @@ async fn main() -> Result<(), Box> { let clients_for_get = clients.clone(); let get_existing_user = Box::new(move |id: u64| -> Option { - clients_for_get.lock().unwrap().get(&id).cloned() + let result = clients_for_get.lock().unwrap().get(&id).cloned(); + if result.is_some() { + println!("Auth lookup: client ID {id} found"); + } else { + eprintln!("Auth lookup: unknown client ID {id}"); + } + result }); let clients_for_register = clients.clone(); @@ -52,7 +58,13 @@ async fn main() -> Result<(), Box> { let id = *nid; *nid += 1; db.insert(id, bundle); - std::fs::write(&clients_path, serde_json::to_string_pretty(&*db).unwrap()).ok(); + match serde_json::to_string_pretty(&*db) { + Ok(json) => match std::fs::write(&clients_path, json) { + Ok(()) => {} + Err(e) => eprintln!("Failed to persist client database to {clients_path}: {e}"), + }, + Err(e) => eprintln!("Failed to serialize client database after registering {id}: {e}"), + } println!("Registered new client with ID: {}", id); id }); diff --git a/example-usage/web-client/src/main.ts b/example-usage/web-client/src/main.ts index 4c0290b..661b1cd 100644 --- a/example-usage/web-client/src/main.ts +++ b/example-usage/web-client/src/main.ts @@ -6,6 +6,7 @@ import init, { ed25519_generate, keyring_from_ed25519, build_demo_message, + format_frame, } from "mtp-wasm"; const STATUS = document.getElementById("status")!; @@ -57,7 +58,9 @@ function hexToBytes(value: string): Uint8Array { } function saveKeys() { - if (!keyringBytes) return; + if (!keyringBytes) { + return; + } let hostPublicKey: number[] | undefined; try { @@ -106,7 +109,7 @@ async function loadHostPublicKey() { HOST_PUBLIC_KEY.value = hostPublicKey; saveKeys(); - log("Loaded host public key bundle from public file."); + log(`Loaded host public key bundle (${hostPublicKey.length / 2} bytes).`); } catch { // Manual paste still works when the server has not exported the file yet. } @@ -138,11 +141,11 @@ function createClient(): WasmClient { (state: number) => log(`[state] ${ConnectionState[state] ?? state}`, "state"), (data: Uint8Array) => { - const decoder = new TextDecoder(); - log( - `[message] ${data.length} bytes: ${decoder.decode(data)}`, - "received", - ); + try { + log(`Received: ${format_frame(data)}`, "received"); + } catch (e) { + log(`[message parse error] ${e}`, "error"); + } }, (err: any) => log(`[error] ${err}`, "error"), ); @@ -181,7 +184,8 @@ async function connect() { await loadDevCertHash(); const client = createClient(); - const config = new ConnectionConfig(SERVER_URL.value.trim()); + const serverUrl = SERVER_URL.value.trim(); + const config = new ConnectionConfig(serverUrl); if (devCertHash) { log(`Pinning WebTransport certificate hash: sha-256:${devCertHash}`); config.server_certificate_hashes = [`sha-256:${devCertHash}`]; @@ -210,6 +214,7 @@ async function connect() { log("\nSending demo message..."); const frame = build_demo_message(activeClientId, keyringBytes); + log(`Sending: ${format_frame(frame)}`, "state"); await client.send(frame); log(`Sent ${frame.length} bytes`); @@ -236,6 +241,10 @@ GENERATE_KEYPAIR.addEventListener("click", () => { CONNECT.addEventListener("click", () => { connect().catch((e) => { log(`Fatal error: ${e}`, "error"); + log( + `[fatal context] clientId=${clientId?.toString() ?? "unregistered"}, server=${SERVER_URL.value.trim()}, hostPkChars=${HOST_PUBLIC_KEY.value.replace(/[^0-9a-fA-F]/g, "").length}, keyringBytes=${keyringBytes?.length ?? 0}, certHash=${devCertHash || "none"}`, + "error", + ); console.error(e); }); }); diff --git a/flake.nix b/flake.nix index 5f62d97..18375a9 100644 --- a/flake.nix +++ b/flake.nix @@ -7,18 +7,16 @@ flake-utils.url = "github:numtide/flake-utils"; }; - outputs = - { - self, - nixpkgs, - rust-overlay, - flake-utils, - }: + outputs = { + self, + nixpkgs, + rust-overlay, + flake-utils, + }: flake-utils.lib.eachDefaultSystem ( - 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 = [ @@ -26,68 +24,80 @@ "clippy" "rustfmt" ]; - targets = [ "wasm32-unknown-unknown" ]; + targets = ["wasm32-unknown-unknown"]; }; - in - { - devShells.default = pkgs.mkShell { - name = "mtp-dev"; + in { + devShells = { + default = pkgs.mkShell { + name = "mtp-dev"; - buildInputs = with pkgs; [ - rustToolchain - wasm-pack - pkg-config - openssl - ]; + buildInputs = with pkgs; [ + rustToolchain + wasm-pack + pkg-config + openssl + ]; - MTP_TYPE_MAPS = "${toString ./example-usage/type-maps.yaml}"; + MTP_TYPE_MAPS = "${toString ./example-usage/type-maps.yaml}"; - shellHook = '' - cert_dir="$PWD/example-usage/dev-cert" - public_dir="$PWD/example-usage/web-client/public" - cert_key="$cert_dir/key.pem" - cert_pem="$cert_dir/cert.pem" - cert_hash="$cert_dir/sha256.txt" + shellHook = '' + repo_root="$(git rev-parse --show-toplevel 2>/dev/null || pwd)" + cert_dir="$repo_root/example-usage/dev-cert" + public_dir="$repo_root/example-usage/web-client/public" + cert_key="$cert_dir/key.pem" + cert_pem="$cert_dir/cert.pem" + cert_hash="$cert_dir/sha256.txt" - mkdir -p "$cert_dir" - mkdir -p "$public_dir" + mkdir -p "$cert_dir" + mkdir -p "$public_dir" - if [ ! -f "$cert_key" ] || [ ! -f "$cert_pem" ]; then - openssl ecparam -name prime256v1 -genkey -noout -out "$cert_key" - openssl req -new -x509 \ - -sha256 \ - -key "$cert_key" \ - -out "$cert_pem" \ - -days 13 \ - -subj "/CN=localhost" \ - -addext "subjectAltName=DNS:localhost,IP:127.0.0.1" \ - -addext "basicConstraints=critical,CA:FALSE" \ - -addext "keyUsage=critical,digitalSignature" \ - -addext "extendedKeyUsage=serverAuth" - cert_status="generated" - else - cert_status="cached" - fi + if [ ! -f "$cert_key" ] || [ ! -f "$cert_pem" ]; then + openssl ecparam -name prime256v1 -genkey -noout -out "$cert_key" + openssl req -new -x509 \ + -sha256 \ + -key "$cert_key" \ + -out "$cert_pem" \ + -days 13 \ + -subj "/CN=localhost" \ + -addext "subjectAltName=DNS:localhost,IP:127.0.0.1" \ + -addext "basicConstraints=critical,CA:FALSE" \ + -addext "keyUsage=critical,digitalSignature" \ + -addext "extendedKeyUsage=serverAuth" + cert_status="generated" + else + cert_status="cached" + fi - openssl x509 -in "$cert_pem" -outform der \ - | openssl dgst -sha256 -binary \ - | od -An -tx1 -v \ - | tr -d ' \n' > "$cert_hash" - cp "$cert_hash" "$public_dir/mtp_dev_cert_hash.txt" + openssl x509 -in "$cert_pem" -outform der \ + | openssl dgst -sha256 -binary \ + | od -An -tx1 -v \ + | tr -d ' \n' > "$cert_hash" + cp "$cert_hash" "$public_dir/mtp_dev_cert_hash.txt" - export MTP_DEV_CERT="$cert_pem" - export MTP_DEV_KEY="$cert_key" - export MTP_DEV_CERT_HASH="$(cat "$cert_hash")" + export MTP_DEV_CERT="$cert_pem" + export MTP_DEV_KEY="$cert_key" + export MTP_DEV_CERT_HASH="$(cat "$cert_hash")" - echo "MTP dev shell" - echo " rustc : $(rustc --version)" - echo " cargo : $(cargo --version)" - echo " wasm-pack : $(wasm-pack --version 2>/dev/null || echo 'not found')" - echo " targets: $(rustc --print target-list | grep wasm32 | tr '\n' ' ')" - echo " MTP_TYPE_MAPS = $MTP_TYPE_MAPS" - echo " dev cert: $MTP_DEV_CERT ($cert_status)" - echo " cert sha256: $MTP_DEV_CERT_HASH" - ''; + echo "MTP dev shell" + echo " rustc : $(rustc --version)" + echo " cargo : $(cargo --version)" + echo " wasm-pack : $(wasm-pack --version 2>/dev/null || echo 'not found')" + echo " targets: $(rustc --print target-list | grep wasm32 | tr '\n' ' ')" + echo " MTP_TYPE_MAPS = $MTP_TYPE_MAPS" + echo " dev cert: $MTP_DEV_CERT ($cert_status)" + echo " cert sha256: $MTP_DEV_CERT_HASH" + ''; + }; + autoStart = pkgs.mkShell { + name = "autoStart"; + buildInputs = with pkgs; [ + mprocs + ]; + shellHook = '' + nix develop --command bash -c "mprocs 'cd wasm && wasm-pack build --target web --out-dir pkg && cd ../example-usage/web-client && bun dev' 'cargo b && cd example-usage && cargo r --bin server'" + exit + ''; + }; }; # Ad-hoc WASM build using wasm-pack diff --git a/host/src/lib.rs b/host/src/lib.rs index cbf0348..f5c3f4b 100644 --- a/host/src/lib.rs +++ b/host/src/lib.rs @@ -176,7 +176,9 @@ impl MTPHost { _ => vec![], }; - let (assigned_id, client_bundle) = if msg.get_type() == mtp_codec::CommunicationTypeId(15) { + let (assigned_id, client_bundle, response_type) = if msg.get_type() + == mtp_codec::CommunicationTypeId(15) + { // LOGIN let cid = match msg.get_data(DataTypeId(6)) { DataValue::UnsignedNumber(n) => *n as u64, @@ -237,7 +239,11 @@ impl MTPHost { } /* ===== End Signature ===== */ - (cid, bundle) + ( + cid, + bundle, + mtp_codec::CommunicationType::IdentificationResponse, + ) } else if msg.get_type() == mtp_codec::CommunicationTypeId(17) { // REGISTER let bundle = match msg.get_data(DataTypeId(9)) { @@ -284,7 +290,11 @@ impl MTPHost { /* ===== End Signature ===== */ let new_id = (self.config.complete_register)(bundle.clone()); - (new_id, bundle) + ( + new_id, + bundle, + mtp_codec::CommunicationType::RegisterResponse, + ) } else { sender.close(); return None; @@ -304,16 +314,15 @@ impl MTPHost { /* ===== Signature ===== */ let host_sig = host_signer.sign(&host_sig_payload).ok()?; - let mut response = - CommunicationValue::new(mtp_codec::CommunicationType::IdentificationResponse) - .add_typed_default(DataType::Connected, DataValue::BoolTrue) - .add_typed_default( - DataType::ClientNonce, - DataValue::UnsignedNumber(client_nonce), - ) - .add_typed_default(DataType::Timestamp, DataValue::UnsignedNumber(new_nonce)) - .add_typed_default(DataType::Signature, DataValue::Bytes(host_sig)) - .add_typed_default(DataType::Id, DataValue::UnsignedNumber(assigned_id as u128)); + let mut response = CommunicationValue::new(response_type) + .add_typed_default(DataType::Connected, DataValue::BoolTrue) + .add_typed_default( + DataType::ClientNonce, + DataValue::UnsignedNumber(client_nonce), + ) + .add_typed_default(DataType::Timestamp, DataValue::UnsignedNumber(new_nonce)) + .add_typed_default(DataType::Signature, DataValue::Bytes(host_sig)) + .add_typed_default(DataType::Id, DataValue::UnsignedNumber(assigned_id as u128)); if !self .config @@ -335,6 +344,7 @@ impl MTPHost { /* ===== End Signature ===== */ sender.send(&response).await.ok()?; + sender.finish_stream().await.ok()?; // 3. Version negotiation let negotiated = self.registry.negotiate(&[client_version])?; diff --git a/transport/src/connection.rs b/transport/src/connection.rs index ed6bbd8..49fb57b 100644 --- a/transport/src/connection.rs +++ b/transport/src/connection.rs @@ -61,7 +61,7 @@ enum ReceivedFrame { pub struct Sender { send_guard: Mutex<()>, - stream_guard: Mutex>, + stream_guard: Arc>>, handle: Arc, connection: Connection, policy: Arc, @@ -71,7 +71,7 @@ impl Sender { pub fn new(connection: Connection, handle: Arc, policy: Arc) -> Self { Self { send_guard: Mutex::new(()), - stream_guard: Mutex::new(None), + stream_guard: Arc::new(Mutex::new(None)), handle, connection, policy, @@ -275,6 +275,18 @@ impl Sender { } } + pub async fn finish_stream(&self) -> Result<(), CommunicationError> { + let _send_lock = self.send_guard.lock().await; + let mut stream_opt = self.stream_guard.lock().await; + if let Some(mut stream) = stream_opt.take() { + timeout(self.policy.write_timeout, stream.finish()) + .await + .map_err(|_| CommunicationError::StreamError)? + .map_err(|_| CommunicationError::StreamError)?; + } + Ok(()) + } + pub fn handle(&self) -> &Arc { &self.handle } @@ -283,6 +295,7 @@ impl Sender { let connection = self.connection.clone(); let handle = self.handle.clone(); let policy = self.policy.clone(); + let stream_guard = self.stream_guard.clone(); tokio::spawn(async move { if connection.quic_connection().close_reason().is_some() || handle.is_closed() { @@ -290,6 +303,14 @@ impl Sender { return; } + if let Some(mut stream) = stream_guard.lock().await.take() { + match timeout(policy.write_timeout, stream.finish()).await { + Ok(Ok(())) => {} + Ok(Err(e)) => log::warn!("[Sender] persistent stream finish failed: {e}"), + Err(_) => log::warn!("[Sender] persistent stream finish timed out"), + } + } + let _ = Self::send_close_frame(&connection, &policy).await; handle.close(Some(CommunicationError::StreamClosed)); diff --git a/wasm/src/client.rs b/wasm/src/client.rs index b5feaae..0730600 100644 --- a/wasm/src/client.rs +++ b/wasm/src/client.rs @@ -3,9 +3,7 @@ use std::rc::Rc; use wasm_bindgen::prelude::*; -use mtp_codec::{ - CommunicationType, CommunicationValue, DataType, DataValue, PROTOCOL_VERSION, -}; +use mtp_codec::{CommunicationType, CommunicationValue, DataType, DataValue, PROTOCOL_VERSION}; use mtp_type_map::{CommunicationTypeId, DataTypeId}; use mtp_crypto::SignatureScheme; @@ -13,6 +11,31 @@ use mtp_crypto::SignatureScheme; use crate::error::js_error; use crate::transport::WasmTransport; +fn raw_frame_preview(bytes: &[u8]) -> String { + let shown = bytes.len().min(256); + let mut preview = hex::encode(&bytes[..shown]); + if bytes.len() > shown { + preview.push_str("..."); + } + format!("{} bytes, hex={preview}", bytes.len()) +} + +fn unexpected_response_type_error( + context: &str, + expected_type: CommunicationTypeId, + response_type: CommunicationTypeId, + response: &[u8], + parsed: &CommunicationValue, +) -> JsValue { + js_error(&format!( + "unexpected response type during {context}: expected {:?}, got {:?}; raw {}; parsed {}", + expected_type, + response_type, + raw_frame_preview(response), + parsed + )) +} + #[wasm_bindgen] #[derive(Debug, Clone, Copy, PartialEq, Eq)] pub enum ConnectionState { @@ -33,17 +56,27 @@ pub struct ConnectionConfig { impl ConnectionConfig { #[wasm_bindgen(constructor)] pub fn new(url: String) -> Self { - Self { url, server_certificate_hashes: None, client_id: 0 } + Self { + url, + server_certificate_hashes: None, + client_id: 0, + } } #[wasm_bindgen(getter)] - pub fn url(&self) -> String { self.url.clone() } + pub fn url(&self) -> String { + self.url.clone() + } #[wasm_bindgen(setter)] - pub fn set_client_id(&mut self, id: u64) { self.client_id = id; } + pub fn set_client_id(&mut self, id: u64) { + self.client_id = id; + } #[wasm_bindgen(getter)] - pub fn client_id(&self) -> u64 { self.client_id } + pub fn client_id(&self) -> u64 { + self.client_id + } #[wasm_bindgen(setter)] pub fn set_server_certificate_hashes(&mut self, hashes: Vec) { @@ -79,24 +112,29 @@ impl WasmClient { #[wasm_bindgen] pub fn is_supported() -> bool { - js_sys::Reflect::has(&js_sys::global(), &JsValue::from_str("WebTransport")) - .unwrap_or(false) + 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 } + pub fn state(&self) -> u8 { + self.state.get() as u8 + } /// Unauthenticated connect (sends basic Identification, enables receive loop). #[wasm_bindgen] pub async fn connect(&mut self, config: &ConnectionConfig) -> Result<(), JsValue> { self.set_state(ConnectionState::Connecting); - let transport = WasmTransport::connect(&config.url, config.server_certificate_hashes.clone()).await?; + let transport = + WasmTransport::connect(&config.url, config.server_certificate_hashes.clone()).await?; let inner = transport.inner().clone(); let version_str = format!("{}", PROTOCOL_VERSION); let ident = CommunicationValue::new(CommunicationType::Identification) .add_typed_default(DataType::Version, DataValue::Str(version_str)) - .add_typed_default(DataType::Id, DataValue::UnsignedNumber(config.client_id as u128)); + .add_typed_default( + DataType::Id, + DataValue::UnsignedNumber(config.client_id as u128), + ); transport.send_frame(&ident.to_bytes()).await?; self.transport = Some(transport); @@ -106,7 +144,9 @@ impl WasmClient { let on_msg = self.on_message.clone(); let on_err = self.on_error.clone(); wasm_bindgen_futures::spawn_local(async move { - WasmTransport::from_inner(inner).receive_loop(on_msg, on_err).await; + WasmTransport::from_inner(inner) + .receive_loop(on_msg, on_err) + .await; state.set(ConnectionState::Disconnected); }); Ok(()) @@ -137,8 +177,7 @@ impl WasmClient { let version_str = format!("{}", PROTOCOL_VERSION); let mut nonce_bytes = [0u8; 16]; - getrandom::fill(&mut nonce_bytes) - .map_err(|_| js_error("rng failed"))?; + getrandom::fill(&mut nonce_bytes).map_err(|_| js_error("rng failed"))?; let client_nonce = u128::from_be_bytes(nonce_bytes); // Build signature payload: version || client_id || client_nonce @@ -149,17 +188,22 @@ impl WasmClient { let signer = mtp_crypto::Ed25519Signer::new(&keyring.sig_cl_secret_key) .map_err(|e| js_error(&format!("signer creation failed: {}", e)))?; - let signature = signer.sign(&sig_payload) + let signature = signer + .sign(&sig_payload) .map_err(|e| js_error(&format!("signature failed: {}", e)))?; let frame = CommunicationValue::new(CommunicationType::Identification) .add_typed_default(DataType::Version, DataValue::Str(version_str)) .add_typed_default(DataType::Id, DataValue::UnsignedNumber(client_id as u128)) - .add_typed_default(DataType::ClientNonce, DataValue::UnsignedNumber(client_nonce)) + .add_typed_default( + DataType::ClientNonce, + DataValue::UnsignedNumber(client_nonce), + ) .add_typed_default(DataType::Signature, DataValue::Bytes(signature)) .to_bytes(); - let transport = WasmTransport::connect(&config.url, config.server_certificate_hashes.clone()).await?; + let transport = + WasmTransport::connect(&config.url, config.server_certificate_hashes.clone()).await?; let inner = transport.inner().clone(); transport.send_frame(&frame).await?; @@ -171,7 +215,13 @@ impl WasmClient { let resp_type = resp_comm.get_type(); let expected_type = CommunicationTypeId(16); // IdentificationResponse if resp_type != expected_type { - return Err(js_error("unexpected response type")); + return Err(unexpected_response_type_error( + "auth_connect", + expected_type, + resp_type, + &response, + &resp_comm, + )); } if resp_comm.get_data(DataTypeId(11)) != &DataValue::BoolTrue { @@ -197,7 +247,9 @@ impl WasmClient { let on_msg = self.on_message.clone(); let on_err = self.on_error.clone(); wasm_bindgen_futures::spawn_local(async move { - WasmTransport::from_inner(inner).receive_loop(on_msg, on_err).await; + WasmTransport::from_inner(inner) + .receive_loop(on_msg, on_err) + .await; state.set(ConnectionState::Disconnected); }); @@ -227,8 +279,7 @@ impl WasmClient { let version_str = format!("{}", PROTOCOL_VERSION); let mut nonce_bytes = [0u8; 16]; - getrandom::fill(&mut nonce_bytes) - .map_err(|_| js_error("rng failed"))?; + getrandom::fill(&mut nonce_bytes).map_err(|_| js_error("rng failed"))?; let client_nonce = u128::from_be_bytes(nonce_bytes); let pk_bytes = keyring.public_key_bundle().as_bytes(); @@ -241,17 +292,22 @@ impl WasmClient { let signer = mtp_crypto::Ed25519Signer::new(&keyring.sig_cl_secret_key) .map_err(|e| js_error(&format!("signer creation failed: {}", e)))?; - let signature = signer.sign(&sig_payload) + let signature = signer + .sign(&sig_payload) .map_err(|e| js_error(&format!("signature failed: {}", e)))?; let frame = CommunicationValue::new(CommunicationType::Register) .add_typed_default(DataType::Version, DataValue::Str(version_str)) - .add_typed_default(DataType::ClientNonce, DataValue::UnsignedNumber(client_nonce)) + .add_typed_default( + DataType::ClientNonce, + DataValue::UnsignedNumber(client_nonce), + ) .add_typed_default(DataType::PublicKeys, DataValue::Bytes(pk_bytes)) .add_typed_default(DataType::Signature, DataValue::Bytes(signature)) .to_bytes(); - let transport = WasmTransport::connect(&config.url, config.server_certificate_hashes.clone()).await?; + let transport = + WasmTransport::connect(&config.url, config.server_certificate_hashes.clone()).await?; let inner = transport.inner().clone(); transport.send_frame(&frame).await?; @@ -262,7 +318,13 @@ impl WasmClient { let resp_type = resp_comm.get_type(); let expected_type = CommunicationTypeId(18); // RegisterResponse if resp_type != expected_type { - return Err(js_error("unexpected response type")); + return Err(unexpected_response_type_error( + "auth_register", + expected_type, + resp_type, + &response, + &resp_comm, + )); } if resp_comm.get_data(DataTypeId(11)) != &DataValue::BoolTrue { @@ -286,7 +348,9 @@ impl WasmClient { let on_msg = self.on_message.clone(); let on_err = self.on_error.clone(); wasm_bindgen_futures::spawn_local(async move { - WasmTransport::from_inner(inner).receive_loop(on_msg, on_err).await; + WasmTransport::from_inner(inner) + .receive_loop(on_msg, on_err) + .await; state.set(ConnectionState::Disconnected); }); @@ -303,16 +367,17 @@ impl WasmClient { #[wasm_bindgen] pub fn disconnect(&mut self) { - if let Some(t) = &self.transport { t.close(); } + if let Some(t) = &self.transport { + t.close(); + } self.transport = None; self.set_state(ConnectionState::Disconnected); } fn set_state(&self, new_state: ConnectionState) { self.state.set(new_state); - let _ = self.on_state_change.call1( - &JsValue::NULL, - &JsValue::from(new_state as u8), - ); + let _ = self + .on_state_change + .call1(&JsValue::NULL, &JsValue::from(new_state as u8)); } } diff --git a/wasm/src/message.rs b/wasm/src/message.rs index 52699e8..49e39f5 100644 --- a/wasm/src/message.rs +++ b/wasm/src/message.rs @@ -1,27 +1,23 @@ use wasm_bindgen::prelude::*; -use mtp_codec::{ - CommunicationType, CommunicationValue, DataType, DataTypeId, DataValue, -}; +use mtp_codec::{CommunicationType, CommunicationValue, DataType, DataTypeId, DataValue}; +use mtp_crypto::{ChaCha20Poly1305, Ed25519Signer, Keyring, SigAlgorithm, derive_encryption_key}; use mtp_type_map::communication_type_name; -use mtp_crypto::{ - ChaCha20Poly1305, Ed25519Signer, Keyring, SigAlgorithm, - derive_encryption_key, -}; use crate::error::js_error; /// Build a simple Ping frame with description, timestamp, and optional data. #[wasm_bindgen] -pub fn build_ping_frame( - client_id: u64, - description: &str, - timestamp: u64, - data: &[u8], -) -> Vec { +pub fn build_ping_frame(client_id: u64, description: &str, timestamp: u64, data: &[u8]) -> Vec { let mut msg = CommunicationValue::new(CommunicationType::Ping) - .add_typed_default(DataType::Description, DataValue::Str(description.to_string())) - .add_typed_default(DataType::Timestamp, DataValue::UnsignedNumber(timestamp as u128)) + .add_typed_default( + DataType::Description, + DataValue::Str(description.to_string()), + ) + .add_typed_default( + DataType::Timestamp, + DataValue::UnsignedNumber(timestamp as u128), + ) .with_sender(client_id); if !data.is_empty() { @@ -55,7 +51,8 @@ pub fn build_demo_message(client_id: u64, keyring_bytes: &[u8]) -> Result Result Result { let obj = js_sys::Object::new(); - let _ = js_sys::Reflect::set(&obj, &JsValue::from_str("_id"), &JsValue::from(comm.get_id())); + let _ = js_sys::Reflect::set( + &obj, + &JsValue::from_str("_id"), + &JsValue::from(comm.get_id()), + ); let type_name = communication_type_name(comm.get_type().0).unwrap_or("Unknown"); let _ = js_sys::Reflect::set( @@ -200,12 +212,21 @@ pub fn parse_response_frame(frame: &[u8]) -> Result { } } - let stringified = js_sys::JSON::stringify(&obj) - .map_err(|_| js_error("JSON stringify failed"))?; - stringified.as_string() + let stringified = + js_sys::JSON::stringify(&obj).map_err(|_| js_error("JSON stringify failed"))?; + stringified + .as_string() .ok_or_else(|| js_error("JSON stringify result not a string")) } +/// Parse any MTP frame into the human-readable CommunicationValue display form. +#[wasm_bindgen] +pub fn format_frame(frame: &[u8]) -> Result { + let comm = CommunicationValue::from_bytes(frame) + .map_err(|e| js_error(&format!("parse failed: {}", e)))?; + Ok(comm.to_string()) +} + #[cfg(test)] #[cfg(target_arch = "wasm32")] mod tests { @@ -219,8 +240,14 @@ mod tests { assert_eq!(cv.get_type(), CommunicationTypeId(19)); // Ping assert_eq!(cv.get_sender(), 42); - assert_eq!(cv.get_data(DataTypeId(4)), &DataValue::Str("test-ping".into())); - assert_eq!(cv.get_data(DataTypeId(5)), &DataValue::UnsignedNumber(1234567890)); + assert_eq!( + cv.get_data(DataTypeId(4)), + &DataValue::Str("test-ping".into()) + ); + assert_eq!( + cv.get_data(DataTypeId(5)), + &DataValue::UnsignedNumber(1234567890) + ); } #[wasm_bindgen_test] @@ -231,9 +258,15 @@ mod tests { assert_eq!(cv.get_type(), CommunicationTypeId(19)); assert_eq!(cv.get_sender(), 99); - assert_eq!(cv.get_data(DataTypeId(4)), &DataValue::Str("with-data".into())); + assert_eq!( + cv.get_data(DataTypeId(4)), + &DataValue::Str("with-data".into()) + ); assert_eq!(cv.get_data(DataTypeId(5)), &DataValue::UnsignedNumber(555)); - assert_eq!(cv.get_data(DataTypeId(6)), &DataValue::Bytes(payload.to_vec())); + assert_eq!( + cv.get_data(DataTypeId(6)), + &DataValue::Bytes(payload.to_vec()) + ); } #[wasm_bindgen_test] @@ -264,7 +297,10 @@ mod tests { assert_eq!(cv.get_type(), CommunicationTypeId(19)); // Ping assert_eq!(cv.get_sender(), 7); - assert_eq!(cv.get_data(DataTypeId(4)), &DataValue::Str("MTP WASM Demo".into())); + assert_eq!( + cv.get_data(DataTypeId(4)), + &DataValue::Str("MTP WASM Demo".into()) + ); } #[wasm_bindgen_test] @@ -287,11 +323,13 @@ mod tests { let result = parse_auth_response(&resp).expect("parse failed"); let connected = js_sys::Reflect::get(&result, &"connected".into()) - .ok().and_then(|v| v.as_bool()); + .ok() + .and_then(|v| v.as_bool()); assert_eq!(connected, Some(true)); let id = js_sys::Reflect::get(&result, &"assignedId".into()) - .ok().and_then(|v| v.as_f64()); + .ok() + .and_then(|v| v.as_f64()); assert_eq!(id, Some(42.0)); } @@ -304,7 +342,8 @@ mod tests { let result = parse_auth_response(&resp).expect("parse failed"); let connected = js_sys::Reflect::get(&result, &"connected".into()) - .ok().and_then(|v| v.as_bool()); + .ok() + .and_then(|v| v.as_bool()); assert_eq!(connected, Some(false)); // rejected should have no assignedId diff --git a/wasm/src/transport.rs b/wasm/src/transport.rs index 29ac985..1648d58 100644 --- a/wasm/src/transport.rs +++ b/wasm/src/transport.rs @@ -5,6 +5,8 @@ use web_sys::{WebTransport, WebTransportHash, WebTransportOptions}; use crate::error::js_error; +const CLOSE_FRAME_LEN: u32 = u32::MAX; + /// Given a `SendStream` (old API with `.writable` or new API where stream IS a WritableStream), /// return the object to call `.getWriter()` on. fn resolve_stream_writable(send_stream: &JsValue) -> Result { @@ -199,10 +201,17 @@ impl WasmTransport { } if buffer.len() >= 4 { - let frame_len = - u32::from_be_bytes([buffer[0], buffer[1], buffer[2], buffer[3]]) as usize; - if 4 + frame_len <= buffer.len() { - return Ok(buffer[4..4 + frame_len].to_vec()); + let frame_len = u32::from_be_bytes([buffer[0], buffer[1], buffer[2], buffer[3]]); + if frame_len == CLOSE_FRAME_LEN { + return Err(js_error("connection closed before frame")); + } + + let frame_len = frame_len as usize; + let Some(frame_end) = 4usize.checked_add(frame_len) else { + return Err(js_error("invalid frame length")); + }; + if frame_end <= buffer.len() { + return Ok(buffer[4..frame_end].to_vec()); } } } @@ -230,13 +239,7 @@ impl WasmTransport { let result = match read_fn.call0(&reader_val) { Ok(p) => match JsFuture::from(p.unchecked_into::()).await { Ok(v) => v, - Err(e) => { - let _ = on_error.call1( - &JsValue::NULL, - &JsValue::from_str(&format!("read stream failed: {:?}", e)), - ); - break; - } + Err(_) => break, }, Err(_) => break, }; @@ -302,29 +305,43 @@ impl WasmTransport { // Extract all complete frames from the buffer while buffer.len() >= 4 { - let frame_len = - u32::from_be_bytes([buffer[0], buffer[1], buffer[2], buffer[3]]) as usize; - if 4 + frame_len > buffer.len() { + let frame_len = u32::from_be_bytes([buffer[0], buffer[1], buffer[2], buffer[3]]); + if frame_len == CLOSE_FRAME_LEN { + return Ok(()); + } + + let frame_len = frame_len as usize; + let Some(frame_end) = 4usize.checked_add(frame_len) else { + return Err(js_error("invalid frame length")); + }; + if frame_end > buffer.len() { break; } - let frame = buffer[4..4 + frame_len].to_vec(); + let frame = buffer[4..frame_end].to_vec(); let arr = js_sys::Uint8Array::from(&frame[..]); let _ = on_message.call1(&JsValue::NULL, &arr); - buffer.drain(..4 + frame_len); + buffer.drain(..frame_end); } } // Process any remaining complete frames after stream closes while buffer.len() >= 4 { - let frame_len = - u32::from_be_bytes([buffer[0], buffer[1], buffer[2], buffer[3]]) as usize; - if 4 + frame_len > buffer.len() { + let frame_len = u32::from_be_bytes([buffer[0], buffer[1], buffer[2], buffer[3]]); + if frame_len == CLOSE_FRAME_LEN { + return Ok(()); + } + + let frame_len = frame_len as usize; + let Some(frame_end) = 4usize.checked_add(frame_len) else { + return Err(js_error("invalid frame length")); + }; + if frame_end > buffer.len() { break; } - let frame = buffer[4..4 + frame_len].to_vec(); + let frame = buffer[4..frame_end].to_vec(); let arr = js_sys::Uint8Array::from(&frame[..]); let _ = on_message.call1(&JsValue::NULL, &arr); - buffer.drain(..4 + frame_len); + buffer.drain(..frame_end); } Ok(())