From 3ade8fb480dbac073c0f5fafa347ce3cc3d07d0d Mon Sep 17 00:00:00 2001 From: John Driscoll Date: Wed, 19 Aug 2026 23:06:16 -0500 Subject: [PATCH] feat(wasm-mps): make message pool safe for E2E-encrypted per-recipient delivery Ticket: HSM-548 --- packages/wasm-mps/src/lib.rs | 172 ++++++++++++++++++++-------------- packages/wasm-mps/test/mps.ts | 38 ++++---- 2 files changed, 124 insertions(+), 86 deletions(-) diff --git a/packages/wasm-mps/src/lib.rs b/packages/wasm-mps/src/lib.rs index 8d95005bbf6..783c91fcce0 100644 --- a/packages/wasm-mps/src/lib.rs +++ b/packages/wasm-mps/src/lib.rs @@ -2,6 +2,8 @@ mod mps { + const MAX_MESSAGES: usize = 500; + use multi_party_schnorr::{ common::{ redpallas::{RedPallasPoint, RedPallasPointBytes}, @@ -22,7 +24,7 @@ mod mps { use rand::Rng; use serde::{Deserialize, Serialize}; use std::{ - io::{Cursor, Read}, + collections::{HashMap, VecDeque}, sync::Arc, }; use thiserror::Error; @@ -44,6 +46,9 @@ mod mps { #[error("Protocol Error")] ProtocolError, + + #[error("Unexpected Error")] + UnexpectedError, } /// Internal DKG state used for round 1. @@ -142,7 +147,7 @@ mod mps { pub struct MsgDerivationInit { pub share: Vec, pub pk: [u8; 32], - pub msg: Vec, + pub msg: HashMap>, pub state: Vec, } @@ -154,7 +159,7 @@ mod mps { } pub struct MsgDerivation { - pub msg: Vec, + pub msg: HashMap>, pub state: Vec, pub done: bool, pub ask: Option<[u8; 32]>, @@ -205,40 +210,6 @@ mod mps { result } - /// Serialize a message pool as a concatenation of individually-encoded messages. - /// This format supports simple byte concatenation to merge pools. - pub fn serialize_pool(prefix: &str, msgs: &[DrvMessage]) -> Result, MpsError> { - let mut buf = Vec::new(); - for msg in msgs { - buf.extend(add_prefix( - prefix, - &bincode::serde::encode_to_vec(msg, bincode::config::standard()) - .map_err(|_| MpsError::SerializationError)?, - )); - } - Ok(buf) - } - - /// Deserialize a pool produced by `serialize_pool`. - pub fn deserialize_pool(prefix: &str, data: &[u8]) -> Result, MpsError> { - let mut cursor = Cursor::new(data); - let mut msgs = Vec::new(); - while (cursor.position() as usize) < data.len() { - let mut buf_prefix = vec![0u8; prefix.len()]; - cursor - .read_exact(&mut buf_prefix) - .map_err(|_| MpsError::DeserializationError)?; - let _ = buf_prefix - .strip_prefix(prefix.as_bytes()) - .ok_or(MpsError::InvalidInput)?; - let msg: DrvMessage = - bincode::serde::decode_from_std_read(&mut cursor, bincode::config::standard()) - .map_err(|_| MpsError::DeserializationError)?; - msgs.push(msg); - } - Ok(msgs) - } - fn internal_dkg_round0_process( party_id: u8, decryption_key: &[u8; 32], @@ -827,6 +798,33 @@ mod mps { }) } + fn derivation_msgs_to_hashmap( + prefix: &str, + party_id: u8, + msgs: Vec, + ) -> Result>, MpsError> { + let mut msg_map: HashMap> = Default::default(); + for msg in msgs { + let idx = msg.receiver().ok_or(MpsError::ProtocolError)?; + if idx == party_id || idx >= 3 { + return Err(MpsError::UnexpectedError); + } + msg_map.entry(idx).or_default().push(msg); + } + let mut vec_map: HashMap> = Default::default(); + for (idx, msgs) in msg_map { + vec_map.insert( + idx, + add_prefix( + prefix, + &bincode::serde::encode_to_vec(msgs, bincode::config::standard()) + .map_err(|_| MpsError::SerializationError)?, + ), + ); + } + Ok(vec_map) + } + /// Process round 2 of RedPallas DKG; finalizes keyshare and starts derivation session. pub fn redpallas_dkg_round2_process( round2_messages: &[Vec; 2], @@ -846,16 +844,23 @@ mod mps { DerivationSession::new(party_id, *share.shamir_share(), *derivation_seed) .map_err(|_| MpsError::ProtocolError)?; - let drv = serialize_pool("mps-redpallas-dkg-derivation-message$", &initial_outgoing)?; + let msg = derivation_msgs_to_hashmap( + "mps-redpallas-dkg-derivation-message$", + party_id, + initial_outgoing, + )?; + let incoming: VecDeque = Default::default(); - let state = - bincode::serde::encode_to_vec((party_id, &drv_session), bincode::config::standard()) - .map_err(|_| MpsError::SerializationError)?; + let state = bincode::serde::encode_to_vec( + (party_id, &drv_session, &incoming), + bincode::config::standard(), + ) + .map_err(|_| MpsError::SerializationError)?; Ok(MsgDerivationInit { share: share_bytes, pk, - msg: drv, + msg, state: add_prefix("mps-redpallas-dkg-derivation-state$", &state), }) } @@ -870,30 +875,43 @@ mod mps { state: &[u8], ) -> Result { let state = rem_prefix("mps-redpallas-dkg-derivation-state$", state)?; - let (party_id, mut session): (u8, DerivationSession) = + let (party_id, mut session, mut incoming): (u8, DerivationSession, VecDeque) = bincode::serde::decode_from_slice(&state, bincode::config::standard()) .map(|(v, _)| v) .map_err(|_| MpsError::DeserializationError)?; - let mut pool = deserialize_pool("mps-redpallas-dkg-derivation-message$", messages)?; + if !messages.is_empty() { + let pool: Vec = bincode::serde::decode_from_slice( + &rem_prefix("mps-redpallas-dkg-derivation-message$", messages)?, + bincode::config::standard(), + ) + .map(|(v, _)| v) + .map_err(|_| MpsError::DeserializationError)?; + if incoming.len() + pool.len() > MAX_MESSAGES { + return Err(MpsError::InvalidInput); + } + for msg in &pool { + if msg.receiver() != Some(party_id) { + return Err(MpsError::InvalidInput); + } + } + incoming.extend(pool); + } - // Find and consume the first message in the pool addressed to this party. - let pos = pool - .iter() - .position(|msg| msg.receiver().is_none_or(|to| to == party_id)); - if let Some(idx) = pos { - let msg = pool.remove(idx); - let mut outgoing: Vec = Vec::new(); + let mut outgoing: Vec = Default::default(); + if let Some(msg) = incoming.pop_front() { let status = session .handle_messages(vec![msg], &mut outgoing) .map_err(|_| MpsError::ProtocolError)?; if let DerivationStatus::Aborted(_) = status { return Err(MpsError::ProtocolError); } - pool.extend(outgoing); } - - let new_messages = serialize_pool("mps-redpallas-dkg-derivation-message$", &pool)?; + let msg = derivation_msgs_to_hashmap( + "mps-redpallas-dkg-derivation-message$", + party_id, + outgoing, + )?; let (done, ask, nk, rivk, internal_ivk, external_ivk) = if let Some(keys) = session.derived_keys() { @@ -909,12 +927,14 @@ mod mps { (false, [0u8; 32], [0u8; 32], [0u8; 32], [0u8; 64], [0u8; 64]) }; - let new_state = - bincode::serde::encode_to_vec((party_id, &session), bincode::config::standard()) - .map_err(|_| MpsError::SerializationError)?; + let new_state = bincode::serde::encode_to_vec( + (party_id, &session, &incoming), + bincode::config::standard(), + ) + .map_err(|_| MpsError::SerializationError)?; Ok(MsgDerivation { - msg: new_messages, + msg, state: add_prefix("mps-redpallas-dkg-derivation-state$", &new_state), done, ask: if done { Some(ask) } else { None }, @@ -1502,7 +1522,8 @@ mod tests { } } -use js_sys::Array; +use js_sys::{Array, Object, Reflect, Uint8Array}; +use std::collections::HashMap; use wasm_bindgen::prelude::*; #[wasm_bindgen] @@ -1553,7 +1574,7 @@ impl Share { pub struct MsgDerivationInit { share: Vec, pk: Vec, - msg: Vec, + msg: Result, state: Vec, } @@ -1570,7 +1591,7 @@ impl MsgDerivationInit { } #[wasm_bindgen(getter)] - pub fn msg(&self) -> Vec { + pub fn msg(&self) -> Result { self.msg.clone() } @@ -1582,7 +1603,7 @@ impl MsgDerivationInit { #[wasm_bindgen] pub struct MsgDerivation { - msg: Vec, + msg: Result, state: Vec, done: bool, ask: Option>, @@ -1595,7 +1616,7 @@ pub struct MsgDerivation { #[wasm_bindgen] impl MsgDerivation { #[wasm_bindgen(getter)] - pub fn msg(&self) -> Vec { + pub fn msg(&self) -> Result { self.msg.clone() } @@ -1827,20 +1848,31 @@ pub fn redpallas_dkg_round2_process( Ok(MsgDerivationInit { share: result.share, pk: result.pk.to_vec(), - msg: result.msg, + msg: hashmap_to_js(result.msg), state: result.state, }) } +fn hashmap_to_js(map: HashMap>) -> Result { + let obj = Object::new(); + + for (key, value) in map { + Reflect::set( + &obj, + &JsValue::from(key), + &Uint8Array::from(value.as_slice()), + )?; + } + + Ok(obj.into()) +} + #[wasm_bindgen] -pub fn redpallas_derivation_process( - messages: &[u8], - state: &[u8], -) -> Result { - let result = mps::redpallas_derivation_process(messages, state).map_err(|e| e.to_string())?; +pub fn redpallas_derivation_process(message: &[u8], state: &[u8]) -> Result { + let result = mps::redpallas_derivation_process(message, state).map_err(|e| e.to_string())?; Ok(MsgDerivation { - msg: result.msg, + msg: hashmap_to_js(result.msg), state: result.state, done: result.done, ask: result.ask.map(|ask| ask.to_vec()), diff --git a/packages/wasm-mps/test/mps.ts b/packages/wasm-mps/test/mps.ts index 97031b53927..f4428a10e84 100644 --- a/packages/wasm-mps/test/mps.ts +++ b/packages/wasm-mps/test/mps.ts @@ -2,7 +2,7 @@ import assert from "assert"; import crypto from "crypto"; import * as mps from "../js"; import sodium from "libsodium-wrappers-sumo"; -import { makeImportShares, runDsg, runImportDkg } from "./utils.js"; +import { makeImportShares, runDsg } from "./utils.js"; await sodium.ready; @@ -807,10 +807,9 @@ describe("mps", function () { ), ); for (let i = 0; i < results3.length; i++) { - if (results3[i].msg.length) { - assert( - Buffer.from(results3[i].msg).slice(0, messagePrefix.length).equals(messagePrefix), - ); + const msg = results3[i].msg as Record; + for (const value of Object.values(msg)) { + assert(Buffer.from(value).slice(0, messagePrefix.length).equals(messagePrefix)); } assert(Buffer.from(results3[i].state).slice(0, statePrefix.length).equals(statePrefix)); } @@ -879,31 +878,38 @@ describe("mps", function () { ); }); + function enqueue(messages: Record, msg: unknown) { + for (const [recipient, value] of Object.entries(msg as Record)) { + messages[Number(recipient)].push(value); + } + } + it("runs derivation to completion", function () { this.timeout(30000); const messagePrefix = Buffer.from("mps-redpallas-dkg-derivation-message$"); const statePrefix = Buffer.from("mps-redpallas-dkg-derivation-state$"); - let message = Buffer.concat(results3.map((d) => Buffer.from(d.msg))); + const messages: Record = { 0: [], 1: [], 2: [] }; + for (const result of results3) { + enqueue(messages, result.msg); + } const states = results3.map((d) => d.state); const derivedKeys: Map = new Map(); - for (let round = 0; round < 500 && Array.from(derivedKeys.keys()).length < 3; round++) { + for (let round = 0; round < 500 && derivedKeys.size < 3; round++) { for (let party = 0; party < 3; party++) { - const result = mps.redpallas_derivation_process(message, states[party]); - if (result.msg.length) { - assert(Buffer.from(result.msg).slice(0, messagePrefix.length).equals(messagePrefix)); - } + const input = messages[party].length > 0 ? messages[party].shift() : new Uint8Array(0); + const result = mps.redpallas_derivation_process(input, states[party]); assert(Buffer.from(result.state).slice(0, statePrefix.length).equals(statePrefix)); - message = result.msg; states[party] = result.state; + for (const value of Object.values(result.msg as Record)) { + assert(Buffer.from(value).slice(0, messagePrefix.length).equals(messagePrefix)); + } + enqueue(messages, result.msg); if (result.done) { derivedKeys.set(party, result); } } } - assert.ok( - Array.from(derivedKeys.keys()).length == 3, - "derivation did not complete within 500 rounds", - ); + assert.ok(derivedKeys.size == 3, "derivation did not complete within 500 rounds"); for (let i = 0; i < 3; i++) { const k = derivedKeys.get(i); assert.equal(k.ask.length, 32);