use std::sync::Arc;
use aes_gcm::{Aes256Gcm, KeyInit as AesKeyInit, Nonce as AesNonce, aead::Aead as AesAead};
use bytes::Bytes;
use chacha20poly1305::{ChaCha20Poly1305, Nonce};
use getrandom::getrandom;
use serde::de::{Deserializer, MapAccess, SeqAccess, Visitor};
use serde::ser::SerializeMap;
use serde::{Deserialize, Serialize};
use serde_json::{Map, Value, json};
use sha2::{Digest, Sha256};
use thiserror::Error;
use x25519_dalek::{PublicKey, StaticSecret};
pub const ECDH_DOMAIN_TAG: &[u8] = b"wscall-ecdh-v1";
pub const ECDH_KEY_LEN: usize = 32;
pub struct EcdhKeypair {
secret: StaticSecret,
public: PublicKey,
}
impl EcdhKeypair {
pub fn generate() -> Result<Self, ProtocolError> {
let mut secret_bytes = [0u8; ECDH_KEY_LEN];
getrandom(&mut secret_bytes).map_err(|source| ProtocolError::Random(source.to_string()))?;
let secret = StaticSecret::from(secret_bytes);
let public = PublicKey::from(&secret);
Ok(Self { secret, public })
}
pub fn public_bytes(&self) -> [u8; ECDH_KEY_LEN] {
self.public.to_bytes()
}
pub fn derive_session_key(&self, peer_public: &[u8; ECDH_KEY_LEN]) -> [u8; 32] {
let peer = PublicKey::from(*peer_public);
let shared = self.secret.diffie_hellman(&peer);
derive_session_key(&shared.to_bytes())
}
}
pub fn derive_session_key(shared_secret: &[u8; 32]) -> [u8; 32] {
let mut hasher = Sha256::new();
hasher.update(ECDH_DOMAIN_TAG);
hasher.update(shared_secret);
let result = hasher.finalize();
let mut key = [0u8; 32];
key.copy_from_slice(&result);
key
}
pub fn parse_peer_public(bytes: &[u8]) -> Result<[u8; ECDH_KEY_LEN], ProtocolError> {
if bytes.len() != ECDH_KEY_LEN {
return Err(ProtocolError::InvalidEcdhPublicKey {
expected: ECDH_KEY_LEN,
actual: bytes.len(),
});
}
let mut key = [0u8; ECDH_KEY_LEN];
key.copy_from_slice(bytes);
Ok(key)
}
const AES256_NONCE_LEN: usize = 12;
const CHACHA20_NONCE_LEN: usize = 12;
pub const DEFAULT_MAX_FRAME_BYTES: usize = 100 * 1024 * 1024;
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
#[repr(u8)]
pub enum MessageType {
Api = 0x00,
Event = 0x01,
}
impl TryFrom<u8> for MessageType {
type Error = ProtocolError;
fn try_from(value: u8) -> Result<Self, Self::Error> {
match value {
0x00 => Ok(Self::Api),
0x01 => Ok(Self::Event),
_ => Err(ProtocolError::UnknownMessageType(value)),
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
#[repr(u8)]
pub enum EncryptionKind {
None = 0x00,
ChaCha20 = 0x01,
Aes256 = 0x02,
}
impl TryFrom<u8> for EncryptionKind {
type Error = ProtocolError;
fn try_from(value: u8) -> Result<Self, Self::Error> {
match value {
0x00 => Ok(Self::None),
0x01 => Ok(Self::ChaCha20),
0x02 => Ok(Self::Aes256),
_ => Err(ProtocolError::UnknownEncryption(value)),
}
}
}
#[derive(Debug, Clone)]
pub struct FileAttachment {
pub id: String,
pub name: String,
pub content_type: String,
pub data: Bytes,
}
impl FileAttachment {
pub fn inline_text(
id: impl Into<String>,
name: impl Into<String>,
content_type: impl Into<String>,
text: impl AsRef<str>,
) -> Self {
Self::inline_bytes(
id,
name,
content_type,
Bytes::from(text.as_ref().to_owned()),
)
}
pub fn inline_bytes(
id: impl Into<String>,
name: impl Into<String>,
content_type: impl Into<String>,
bytes: impl Into<Bytes>,
) -> Self {
Self {
id: id.into(),
name: name.into(),
content_type: content_type.into(),
data: bytes.into(),
}
}
pub fn size(&self) -> usize {
self.data.len()
}
pub fn param_ref(id: impl Into<String>) -> Value {
json!({ "$file": id.into() })
}
pub(crate) fn wire_size(&self) -> usize {
1 + self.id.len() + 1 + self.name.len() + 1 + self.content_type.len() + 4 + self.data.len()
}
pub(crate) fn write_wire(&self, buf: &mut Vec<u8>) {
buf.push(self.id.len() as u8);
buf.extend_from_slice(self.id.as_bytes());
buf.push(self.name.len() as u8);
buf.extend_from_slice(self.name.as_bytes());
buf.push(self.content_type.len() as u8);
buf.extend_from_slice(self.content_type.as_bytes());
buf.extend_from_slice(&(self.data.len() as u32).to_be_bytes());
buf.extend_from_slice(&self.data);
}
pub(crate) fn read_wire(data: &Bytes) -> Result<(Self, Bytes), ProtocolError> {
let mut pos = 0;
if pos >= data.len() {
return Err(ProtocolError::InvalidAttachmentEncoding(
"truncated id_len".into(),
));
}
let id_len = data[pos] as usize;
pos += 1;
if pos + id_len > data.len() {
return Err(ProtocolError::InvalidAttachmentEncoding(
"truncated id".into(),
));
}
let id = String::from_utf8_lossy(&data[pos..pos + id_len]).into_owned();
pos += id_len;
if pos >= data.len() {
return Err(ProtocolError::InvalidAttachmentEncoding(
"truncated name_len".into(),
));
}
let name_len = data[pos] as usize;
pos += 1;
if pos + name_len > data.len() {
return Err(ProtocolError::InvalidAttachmentEncoding(
"truncated name".into(),
));
}
let name = String::from_utf8_lossy(&data[pos..pos + name_len]).into_owned();
pos += name_len;
if pos >= data.len() {
return Err(ProtocolError::InvalidAttachmentEncoding(
"truncated ct_len".into(),
));
}
let ct_len = data[pos] as usize;
pos += 1;
if pos + ct_len > data.len() {
return Err(ProtocolError::InvalidAttachmentEncoding(
"truncated content_type".into(),
));
}
let content_type = String::from_utf8_lossy(&data[pos..pos + ct_len]).into_owned();
pos += ct_len;
if pos + 4 > data.len() {
return Err(ProtocolError::InvalidAttachmentEncoding(
"truncated data_len".into(),
));
}
let data_len =
u32::from_be_bytes([data[pos], data[pos + 1], data[pos + 2], data[pos + 3]]) as usize;
pos += 4;
if pos + data_len > data.len() {
return Err(ProtocolError::InvalidAttachmentEncoding(
"truncated data".into(),
));
}
let payload = data.slice(pos..pos + data_len);
pos += data_len;
Ok((
Self {
id,
name,
content_type,
data: payload,
},
data.slice(pos..),
))
}
}
impl Serialize for FileAttachment {
fn serialize<S: serde::Serializer>(&self, serializer: S) -> Result<S::Ok, S::Error> {
let mut map = serializer.serialize_map(Some(4))?;
map.serialize_entry("id", &self.id)?;
map.serialize_entry("name", &self.name)?;
map.serialize_entry("content_type", &self.content_type)?;
map.serialize_entry("data", &self.data.to_vec())?;
map.end()
}
}
impl<'de> Deserialize<'de> for FileAttachment {
fn deserialize<D: Deserializer<'de>>(deserializer: D) -> Result<Self, D::Error> {
#[derive(Deserialize)]
#[serde(field_identifier, rename_all = "snake_case")]
enum Field {
Id,
Name,
ContentType,
Data,
}
struct FileAttachmentVisitor;
impl<'de> Visitor<'de> for FileAttachmentVisitor {
type Value = FileAttachment;
fn expecting(&self, f: &mut std::fmt::Formatter) -> std::fmt::Result {
f.write_str("struct FileAttachment")
}
fn visit_map<A: MapAccess<'de>>(self, mut map: A) -> Result<Self::Value, A::Error> {
let mut id = None;
let mut name = None;
let mut content_type = None;
let mut data: Option<Bytes> = None;
while let Some(key) = map.next_key()? {
match key {
Field::Id => id = Some(map.next_value()?),
Field::Name => name = Some(map.next_value()?),
Field::ContentType => content_type = Some(map.next_value()?),
Field::Data => {
let bytes: Vec<u8> = map.next_value()?;
data = Some(Bytes::from(bytes));
}
}
}
Ok(FileAttachment {
id: id.ok_or_else(|| serde::de::Error::missing_field("id"))?,
name: name.ok_or_else(|| serde::de::Error::missing_field("name"))?,
content_type: content_type
.ok_or_else(|| serde::de::Error::missing_field("content_type"))?,
data: data.ok_or_else(|| serde::de::Error::missing_field("data"))?,
})
}
fn visit_seq<A: SeqAccess<'de>>(self, mut seq: A) -> Result<Self::Value, A::Error> {
let id = seq
.next_element()?
.ok_or_else(|| serde::de::Error::invalid_length(0, &self))?;
let name = seq
.next_element()?
.ok_or_else(|| serde::de::Error::invalid_length(1, &self))?;
let content_type = seq
.next_element()?
.ok_or_else(|| serde::de::Error::invalid_length(2, &self))?;
let bytes: Vec<u8> = seq
.next_element()?
.ok_or_else(|| serde::de::Error::invalid_length(3, &self))?;
Ok(FileAttachment {
id,
name,
content_type,
data: Bytes::from(bytes),
})
}
}
const FIELDS: &[&str] = &["id", "name", "content_type", "data"];
deserializer.deserialize_struct("FileAttachment", FIELDS, FileAttachmentVisitor)
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ErrorPayload {
pub code: String,
pub message: String,
pub status: u16,
#[serde(skip_serializing_if = "Option::is_none")]
pub details: Option<Value>,
}
pub const K_API_REQUEST: u8 = 0;
pub const K_EVENT_EMIT: u8 = 1;
pub const K_API_RESPONSE: u8 = 2;
pub const K_EVENT_ACK: u8 = 3;
#[derive(Debug, Clone)]
pub enum PacketBody {
ApiRequest {
request_id: u64,
route: String,
params: Value,
attachments: Vec<FileAttachment>,
metadata: Value,
},
ApiResponse {
request_id: u64,
ok: bool,
status: u16,
data: Value,
error: Option<ErrorPayload>,
metadata: Value,
},
EventEmit {
event_id: u64,
name: String,
data: Map<String, Value>,
attachments: Vec<FileAttachment>,
metadata: Value,
expect_ack: bool,
},
EventAck {
event_id: u64,
ok: bool,
receipt: Value,
error: Option<ErrorPayload>,
},
}
fn metadata_is_empty(v: &Value) -> bool {
matches!(v, Value::Null) || v.as_object().is_some_and(|m| m.is_empty())
}
impl Serialize for PacketBody {
fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
where
S: serde::Serializer,
{
match self {
Self::ApiRequest {
request_id,
route,
params,
attachments: _,
metadata,
} => {
let include_meta = !metadata_is_empty(metadata);
let field_count = 4 + usize::from(include_meta);
let mut s = serializer.serialize_map(Some(field_count))?;
s.serialize_entry("k", &K_API_REQUEST)?;
s.serialize_entry("i", request_id)?;
s.serialize_entry("r", route)?;
s.serialize_entry("p", params)?;
if include_meta {
s.serialize_entry("m", metadata)?;
}
s.end()
}
Self::ApiResponse {
request_id,
ok,
status,
data,
error,
metadata,
} => {
let include_meta = !metadata_is_empty(metadata);
let field_count = 5 + usize::from(error.is_some()) + usize::from(include_meta);
let mut s = serializer.serialize_map(Some(field_count))?;
s.serialize_entry("k", &K_API_RESPONSE)?;
s.serialize_entry("i", request_id)?;
s.serialize_entry("o", ok)?;
s.serialize_entry("s", status)?;
s.serialize_entry("d", data)?;
if let Some(err) = error {
s.serialize_entry("er", err)?;
}
if include_meta {
s.serialize_entry("m", metadata)?;
}
s.end()
}
Self::EventEmit {
event_id,
name,
data,
attachments: _,
metadata,
expect_ack,
} => {
let include_meta = !metadata_is_empty(metadata);
let field_count = 5 + usize::from(include_meta);
let mut s = serializer.serialize_map(Some(field_count))?;
s.serialize_entry("k", &K_EVENT_EMIT)?;
s.serialize_entry("i", event_id)?;
s.serialize_entry("n", name)?;
s.serialize_entry("d", data)?;
if include_meta {
s.serialize_entry("m", metadata)?;
}
s.serialize_entry("e", expect_ack)?;
s.end()
}
Self::EventAck {
event_id,
ok,
receipt,
error,
} => {
let field_count = 3 + usize::from(error.is_some()) + 1;
let mut s = serializer.serialize_map(Some(field_count))?;
s.serialize_entry("k", &K_EVENT_ACK)?;
s.serialize_entry("i", event_id)?;
s.serialize_entry("o", ok)?;
s.serialize_entry("rc", receipt)?;
if let Some(err) = error {
s.serialize_entry("er", err)?;
}
s.end()
}
}
}
}
fn take_u64<E: serde::de::Error>(map: &mut Map<String, Value>, key: &str) -> Result<u64, E> {
map.remove(key).and_then(|v| v.as_u64()).ok_or_else(|| {
E::custom(format!(
"missing or non-numeric '{key}' field in packet body"
))
})
}
fn take_string<E: serde::de::Error>(map: &mut Map<String, Value>, key: &str) -> Result<String, E> {
match map.remove(key) {
Some(Value::String(s)) => Ok(s),
_ => Err(E::custom(format!(
"missing or non-string '{key}' field in packet body"
))),
}
}
fn take_object<E: serde::de::Error>(
map: &mut Map<String, Value>,
key: &str,
) -> Result<Map<String, Value>, E> {
match map.remove(key) {
None => Ok(Map::new()),
Some(Value::Object(o)) => Ok(o),
Some(_) => Err(E::custom(format!("'{key}' field must be a JSON object"))),
}
}
fn take_bool(map: &mut Map<String, Value>, key: &str) -> bool {
map.remove(key).and_then(|v| v.as_bool()).unwrap_or(false)
}
fn take_u16(map: &mut Map<String, Value>, key: &str) -> u16 {
map.remove(key).and_then(|v| v.as_u64()).unwrap_or(0) as u16
}
fn take_attachments<E: serde::de::Error>(
map: &mut Map<String, Value>,
) -> Result<Vec<FileAttachment>, E> {
match map.remove("a") {
None => Ok(Vec::new()),
Some(value) => serde_json::from_value(value).map_err(E::custom),
}
}
fn take_error<E: serde::de::Error>(
map: &mut Map<String, Value>,
) -> Result<Option<ErrorPayload>, E> {
match map.remove("er") {
None => Ok(None),
Some(value) => serde_json::from_value(value).map(Some).map_err(E::custom),
}
}
impl<'de> Deserialize<'de> for PacketBody {
fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
where
D: Deserializer<'de>,
{
let mut map = match Value::deserialize(deserializer)? {
Value::Object(map) => map,
_ => {
return Err(serde::de::Error::custom(
"packet body must be a JSON object",
));
}
};
let k = map.get("k").and_then(|v| v.as_u64()).ok_or_else(|| {
serde::de::Error::custom("missing or non-numeric 'k' field in packet body")
})?;
let k = u8::try_from(k).map_err(|_| {
serde::de::Error::custom(format!("'k' value {k} is outside the valid range"))
})?;
match k {
K_API_REQUEST => Ok(Self::ApiRequest {
request_id: take_u64(&mut map, "i")?,
route: take_string(&mut map, "r")?,
params: map.remove("p").unwrap_or(Value::Null),
attachments: take_attachments(&mut map)?,
metadata: map.remove("m").unwrap_or(Value::Null),
}),
K_EVENT_EMIT => Ok(Self::EventEmit {
event_id: take_u64(&mut map, "i")?,
name: take_string(&mut map, "n")?,
data: take_object(&mut map, "d")?,
attachments: take_attachments(&mut map)?,
metadata: map.remove("m").unwrap_or(Value::Null),
expect_ack: take_bool(&mut map, "e"),
}),
K_API_RESPONSE => Ok(Self::ApiResponse {
request_id: take_u64(&mut map, "i")?,
ok: take_bool(&mut map, "o"),
status: take_u16(&mut map, "s"),
data: map.remove("d").unwrap_or(Value::Null),
error: take_error(&mut map)?,
metadata: map.remove("m").unwrap_or(Value::Null),
}),
K_EVENT_ACK => Ok(Self::EventAck {
event_id: take_u64(&mut map, "i")?,
ok: take_bool(&mut map, "o"),
receipt: map.remove("rc").unwrap_or(Value::Null),
error: take_error(&mut map)?,
}),
_ => Err(serde::de::Error::custom(format!("unknown 'k' value: {k}"))),
}
}
}
impl PacketBody {
pub fn message_type(&self) -> MessageType {
match self {
Self::ApiRequest { .. } | Self::ApiResponse { .. } => MessageType::Api,
Self::EventEmit { .. } | Self::EventAck { .. } => MessageType::Event,
}
}
pub fn attachments(&self) -> &[FileAttachment] {
match self {
Self::ApiRequest { attachments, .. } => attachments,
Self::EventEmit { attachments, .. } => attachments,
_ => &[],
}
}
pub fn set_attachments(&mut self, attachments: Vec<FileAttachment>) {
match self {
Self::ApiRequest { attachments: a, .. } => *a = attachments,
Self::EventEmit { attachments: a, .. } => *a = attachments,
_ => {}
}
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct PacketEnvelope {
pub message_type: MessageType,
pub encryption: EncryptionKind,
pub body: PacketBody,
}
impl PacketEnvelope {
pub fn new(body: PacketBody) -> Self {
Self {
message_type: body.message_type(),
encryption: EncryptionKind::None,
body,
}
}
pub fn with_encryption(body: PacketBody, encryption: EncryptionKind) -> Self {
Self {
message_type: body.message_type(),
encryption,
body,
}
}
}
#[derive(Clone)]
pub struct FrameCodec {
aes256_cipher: Option<Arc<Aes256Gcm>>,
chacha20_cipher: Option<Arc<ChaCha20Poly1305>>,
wire_encryption: EncryptionKind,
max_frame_bytes: usize,
}
impl Default for FrameCodec {
fn default() -> Self {
Self {
aes256_cipher: None,
chacha20_cipher: None,
wire_encryption: EncryptionKind::None,
max_frame_bytes: DEFAULT_MAX_FRAME_BYTES,
}
}
}
impl std::fmt::Debug for FrameCodec {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("FrameCodec")
.field("aes256", &self.aes256_cipher.is_some())
.field("chacha20", &self.chacha20_cipher.is_some())
.field("wire_encryption", &self.wire_encryption)
.field("max_frame_bytes", &self.max_frame_bytes)
.finish()
}
}
impl FrameCodec {
pub fn plaintext() -> Self {
Self::default()
}
pub fn with_chacha20_key(self, key: [u8; 32]) -> Self {
let cipher = ChaCha20Poly1305::new_from_slice(&key)
.expect("a 32-byte key is always valid for ChaCha20-Poly1305");
Self {
chacha20_cipher: Some(Arc::new(cipher)),
wire_encryption: EncryptionKind::ChaCha20,
..self
}
}
pub fn with_aes256_key(self, key: [u8; 32]) -> Self {
let cipher =
Aes256Gcm::new_from_slice(&key).expect("a 32-byte key is always valid for AES-256-GCM");
Self {
aes256_cipher: Some(Arc::new(cipher)),
wire_encryption: EncryptionKind::Aes256,
..self
}
}
pub fn wire_encryption(&self) -> EncryptionKind {
self.wire_encryption
}
pub fn with_max_frame_bytes(mut self, max: usize) -> Self {
self.max_frame_bytes = max;
self
}
pub fn max_frame_bytes(&self) -> usize {
self.max_frame_bytes
}
pub fn encode(&self, packet: &PacketEnvelope) -> Result<Vec<u8>, ProtocolError> {
let max_payload = self.max_frame_bytes.saturating_sub(5);
let attachments = packet.body.attachments();
let att_size: usize = attachments.iter().map(|a| a.wire_size()).sum();
let message_type = packet.message_type as u8;
if self.wire_encryption == EncryptionKind::None {
let mut frame = Vec::with_capacity(4 + 1 + 4 + 1 + att_size + 64);
frame.extend_from_slice(&[0u8; 4]); frame.push(message_type);
frame.extend_from_slice(&[0u8; 4]); serde_json::to_writer(&mut frame, &packet.body)?;
let json_len = frame.len() - 9;
frame.push(attachments.len() as u8);
for att in attachments {
att.write_wire(&mut frame);
}
let payload_len = frame.len() - 5;
if payload_len > max_payload {
return Err(ProtocolError::PayloadTooLarge {
actual: payload_len,
max: max_payload,
});
}
let frame_len = (frame.len() - 4) as u32;
frame[0..4].copy_from_slice(&frame_len.to_be_bytes());
frame[5..9].copy_from_slice(&(json_len as u32).to_be_bytes());
return Ok(frame);
}
let mut composite = Vec::with_capacity(4 + 1 + att_size + 64);
composite.extend_from_slice(&[0u8; 4]); serde_json::to_writer(&mut composite, &packet.body)?;
let json_len = composite.len() - 4;
composite[0..4].copy_from_slice(&(json_len as u32).to_be_bytes());
composite.push(attachments.len() as u8);
for att in attachments {
att.write_wire(&mut composite);
}
if composite.len() > max_payload {
return Err(ProtocolError::PayloadTooLarge {
actual: composite.len(),
max: max_payload,
});
}
let payload = match self.wire_encryption {
EncryptionKind::ChaCha20 => self.encrypt_chacha20(&composite)?,
EncryptionKind::Aes256 => self.encrypt_aes256(&composite)?,
EncryptionKind::None => unreachable!("plaintext is handled by the fast path above"),
};
if payload.len() > max_payload {
return Err(ProtocolError::PayloadTooLarge {
actual: payload.len(),
max: max_payload,
});
}
let frame_len = 1 + payload.len();
let mut frame = Vec::with_capacity(4 + frame_len);
frame.extend_from_slice(&(frame_len as u32).to_be_bytes());
frame.push(message_type);
frame.extend_from_slice(&payload);
Ok(frame)
}
pub fn decode(&self, frame: Bytes) -> Result<PacketEnvelope, ProtocolError> {
if frame.len() < 5 {
return Err(ProtocolError::FrameTooShort);
}
let declared = u32::from_be_bytes([frame[0], frame[1], frame[2], frame[3]]) as usize;
let actual = frame.len() - 4;
if declared != actual {
return Err(ProtocolError::FrameLengthMismatch { declared, actual });
}
if frame.len() > self.max_frame_bytes {
return Err(ProtocolError::FrameTooLarge {
actual: frame.len(),
max: self.max_frame_bytes,
});
}
let message_type = MessageType::try_from(frame[4])?;
let composite: Bytes = match self.wire_encryption {
EncryptionKind::None => frame.slice(5..),
EncryptionKind::ChaCha20 => Bytes::from(self.decrypt_chacha20(&frame[5..])?),
EncryptionKind::Aes256 => Bytes::from(self.decrypt_aes256(&frame[5..])?),
};
if composite.len() < 5 {
return Err(ProtocolError::FrameTooShort);
}
let meta_len =
u32::from_be_bytes([composite[0], composite[1], composite[2], composite[3]]) as usize;
if composite.len() < 4 + meta_len + 1 {
return Err(ProtocolError::FrameTooShort);
}
let json_slice = &composite[4..4 + meta_len];
let att_count = composite[4 + meta_len] as usize;
let mut att_data = composite.slice(4 + meta_len + 1..);
let mut attachments = Vec::with_capacity(att_count);
for _ in 0..att_count {
let (att, rest) = FileAttachment::read_wire(&att_data)?;
attachments.push(att);
att_data = rest;
}
let mut body: PacketBody = serde_json::from_slice(json_slice)?;
body.set_attachments(attachments);
if body.message_type() != message_type {
return Err(ProtocolError::MessageTypeMismatch);
}
Ok(PacketEnvelope {
message_type,
encryption: self.wire_encryption,
body,
})
}
fn encrypt_chacha20(&self, payload: &[u8]) -> Result<Vec<u8>, ProtocolError> {
let cipher = self
.chacha20_cipher
.as_ref()
.ok_or(ProtocolError::MissingEncryptionKey("chacha20"))?;
let mut nonce_bytes = [0_u8; CHACHA20_NONCE_LEN];
getrandom(&mut nonce_bytes).map_err(|source| ProtocolError::Random(source.to_string()))?;
let ciphertext = cipher
.encrypt(Nonce::from_slice(&nonce_bytes), payload)
.map_err(|_| ProtocolError::EncryptionFailed("chacha20"))?;
let mut encoded = Vec::with_capacity(CHACHA20_NONCE_LEN + ciphertext.len());
encoded.extend_from_slice(&nonce_bytes);
encoded.extend_from_slice(&ciphertext);
Ok(encoded)
}
fn decrypt_chacha20(&self, payload: &[u8]) -> Result<Vec<u8>, ProtocolError> {
if payload.len() < CHACHA20_NONCE_LEN {
return Err(ProtocolError::EncryptedPayloadTooShort {
algorithm: "chacha20",
expected_min: CHACHA20_NONCE_LEN,
actual: payload.len(),
});
}
let cipher = self
.chacha20_cipher
.as_ref()
.ok_or(ProtocolError::MissingEncryptionKey("chacha20"))?;
let (nonce_bytes, ciphertext) = payload.split_at(CHACHA20_NONCE_LEN);
cipher
.decrypt(Nonce::from_slice(nonce_bytes), ciphertext)
.map_err(|_| ProtocolError::DecryptionFailed("chacha20"))
}
fn encrypt_aes256(&self, payload: &[u8]) -> Result<Vec<u8>, ProtocolError> {
let cipher = self
.aes256_cipher
.as_ref()
.ok_or(ProtocolError::MissingEncryptionKey("aes256"))?;
let mut nonce_bytes = [0_u8; AES256_NONCE_LEN];
getrandom(&mut nonce_bytes).map_err(|source| ProtocolError::Random(source.to_string()))?;
let ciphertext = cipher
.encrypt(AesNonce::from_slice(&nonce_bytes), payload)
.map_err(|_| ProtocolError::EncryptionFailed("aes256"))?;
let mut encoded = Vec::with_capacity(AES256_NONCE_LEN + ciphertext.len());
encoded.extend_from_slice(&nonce_bytes);
encoded.extend_from_slice(&ciphertext);
Ok(encoded)
}
fn decrypt_aes256(&self, payload: &[u8]) -> Result<Vec<u8>, ProtocolError> {
if payload.len() < AES256_NONCE_LEN {
return Err(ProtocolError::EncryptedPayloadTooShort {
algorithm: "aes256",
expected_min: AES256_NONCE_LEN,
actual: payload.len(),
});
}
let cipher = self
.aes256_cipher
.as_ref()
.ok_or(ProtocolError::MissingEncryptionKey("aes256"))?;
let (nonce_bytes, ciphertext) = payload.split_at(AES256_NONCE_LEN);
cipher
.decrypt(AesNonce::from_slice(nonce_bytes), ciphertext)
.map_err(|_| ProtocolError::DecryptionFailed("aes256"))
}
}
pub fn encode_frame(packet: &PacketEnvelope) -> Result<Vec<u8>, ProtocolError> {
FrameCodec::plaintext().encode(packet)
}
pub fn decode_frame(frame: Bytes) -> Result<PacketEnvelope, ProtocolError> {
FrameCodec::plaintext().decode(frame)
}
#[derive(Debug, Error)]
pub enum ProtocolError {
#[error("frame too short")]
FrameTooShort,
#[error("frame length mismatch: declared={declared}, actual={actual}")]
FrameLengthMismatch { declared: usize, actual: usize },
#[error("payload too large: actual={actual}, max={max}")]
PayloadTooLarge { actual: usize, max: usize },
#[error("frame too large: actual={actual}, max={max}")]
FrameTooLarge { actual: usize, max: usize },
#[error("unknown message type: {0:#x}")]
UnknownMessageType(u8),
#[error("unknown encryption kind: {0:#x}")]
UnknownEncryption(u8),
#[error("unsupported encryption kind: {0:#x}")]
UnsupportedEncryption(u8),
#[error("missing encryption key for {0}")]
MissingEncryptionKey(&'static str),
#[error("invalid encryption key for {0}")]
InvalidEncryptionKey(&'static str),
#[error("secure random generation failed: {0}")]
Random(String),
#[error(
"encrypted payload too short for {algorithm}: expected at least {expected_min}, actual={actual}"
)]
EncryptedPayloadTooShort {
algorithm: &'static str,
expected_min: usize,
actual: usize,
},
#[error("encryption failed for {0}")]
EncryptionFailed(&'static str),
#[error("decryption failed for {0}")]
DecryptionFailed(&'static str),
#[error("invalid ECDH public key: expected {expected} bytes, got {actual}")]
InvalidEcdhPublicKey { expected: usize, actual: usize },
#[error("ECDH handshake failed: {0}")]
EcdhHandshake(String),
#[error("message type does not match packet body")]
MessageTypeMismatch,
#[error("invalid attachment encoding: {0}")]
InvalidAttachmentEncoding(String),
#[error("json error: {0}")]
Json(#[from] serde_json::Error),
}
#[cfg(test)]
mod tests {
use super::{
Bytes, EncryptionKind, FileAttachment, FrameCodec, PacketBody, PacketEnvelope,
ProtocolError, decode_frame, encode_frame,
};
use serde_json::json;
const TEST_KEY: [u8; 32] = [0x11; 32];
#[test]
fn plaintext_helpers_still_work() {
let packet = PacketEnvelope::new(PacketBody::EventAck {
event_id: 1,
ok: true,
receipt: json!({ "ok": true }),
error: None,
});
let encoded = encode_frame(&packet).expect("encode plaintext");
let decoded = decode_frame(Bytes::from(encoded)).expect("decode plaintext");
assert!(matches!(decoded.encryption, EncryptionKind::None));
}
#[test]
fn aes256_roundtrip_works() {
let codec = FrameCodec::plaintext().with_aes256_key(TEST_KEY);
let packet = PacketEnvelope::with_encryption(
PacketBody::ApiResponse {
request_id: 1,
ok: true,
status: 200,
data: json!({ "message": "encrypted" }),
error: None,
metadata: json!({}),
},
EncryptionKind::Aes256,
);
let encoded = codec.encode(&packet).expect("encode aes256");
let decoded = codec.decode(Bytes::from(encoded)).expect("decode aes256");
assert!(matches!(decoded.encryption, EncryptionKind::Aes256));
}
#[test]
fn attachment_roundtrip_plaintext() {
let att = FileAttachment::inline_bytes(
"f1",
"test.bin",
"application/octet-stream",
vec![1, 2, 3, 4, 5],
);
let packet = PacketEnvelope::new(PacketBody::ApiRequest {
request_id: 42,
route: "files.upload".to_string(),
params: json!({ "file": { "$file": "f1" } }),
attachments: vec![att],
metadata: json!({}),
});
let encoded = encode_frame(&packet).expect("encode with attachment");
let decoded = decode_frame(Bytes::from(encoded)).expect("decode with attachment");
let atts = decoded.body.attachments();
assert_eq!(atts.len(), 1);
assert_eq!(atts[0].id, "f1");
assert_eq!(atts[0].name, "test.bin");
assert_eq!(atts[0].content_type, "application/octet-stream");
assert_eq!(&atts[0].data[..], &[1, 2, 3, 4, 5]);
}
#[test]
fn attachment_roundtrip_encrypted() {
let codec = FrameCodec::plaintext().with_chacha20_key(TEST_KEY);
let att = FileAttachment::inline_text("f2", "hello.txt", "text/plain", "hello world");
let packet = PacketEnvelope::with_encryption(
PacketBody::EventEmit {
event_id: 7,
name: "chat.message".to_string(),
data: json!({ "text": "see attached" })
.as_object()
.unwrap()
.clone(),
attachments: vec![att],
metadata: json!({}),
expect_ack: true,
},
EncryptionKind::ChaCha20,
);
let encoded = codec
.encode(&packet)
.expect("encode encrypted with attachment");
let decoded = codec
.decode(Bytes::from(encoded))
.expect("decode encrypted with attachment");
let atts = decoded.body.attachments();
assert_eq!(atts.len(), 1);
assert_eq!(atts[0].id, "f2");
assert_eq!(&atts[0].data[..], b"hello world");
}
#[test]
fn encode_rejects_payloads_over_limit() {
let codec = FrameCodec::plaintext().with_max_frame_bytes(1024);
let packet = PacketEnvelope::new(PacketBody::ApiResponse {
request_id: 999,
ok: true,
status: 200,
data: json!({ "blob": "a".repeat(2048) }),
error: None,
metadata: json!({}),
});
let error = codec
.encode(&packet)
.expect_err("oversized payload should fail");
assert!(matches!(error, ProtocolError::PayloadTooLarge { .. }));
}
#[test]
fn decode_rejects_frames_over_limit() {
let codec = FrameCodec::plaintext().with_max_frame_bytes(64);
let packet = PacketEnvelope::new(PacketBody::ApiResponse {
request_id: 1,
ok: true,
status: 200,
data: json!({ "msg": "a]".repeat(50) }),
error: None,
metadata: json!({}),
});
let encoded = FrameCodec::plaintext()
.encode(&packet)
.expect("encode with default limit");
assert!(encoded.len() > 64);
let error = codec
.decode(Bytes::from(encoded))
.expect_err("oversized frame should fail");
assert!(matches!(error, ProtocolError::FrameTooLarge { .. }));
}
}