use super::wire_msg_header::WireMsgHeader;
use crate::messaging::{
data::{ClientDataResponse, ClientMsg},
system::{NodeDataResponse, NodeMsg},
AuthorityProof, Dst, Error, MsgId, MsgKind, MsgType, Result,
};
use bytes::{BufMut, Bytes, BytesMut};
use custom_debug::Debug;
use qp2p::UsrMsgBytes;
use serde::Serialize;
use xor_name::XorName;
#[derive(Clone, Debug)]
pub struct WireMsg {
pub header: WireMsgHeader,
#[debug(skip)]
pub serialized_header: Option<Bytes>,
#[debug(skip)]
pub payload: Bytes,
pub dst: Dst,
#[debug(skip)]
pub serialized_dst: Option<Bytes>,
}
impl PartialEq for WireMsg {
fn eq(&self, other: &Self) -> bool {
self.header == other.header && self.payload == other.payload
}
}
impl WireMsg {
pub fn serialize_msg_payload<T: Serialize>(msg: &T) -> Result<Bytes> {
let mut bytes = BytesMut::new().writer();
rmp_serde::encode::write(&mut bytes, &msg).map_err(|err| {
Error::Serialisation(format!(
"could not serialize message payload with Msgpack: {err}",
))
})?;
Ok(bytes.into_inner().freeze())
}
fn serialize_dst_payload(dst: &Dst) -> Result<Bytes> {
let mut bytes = BytesMut::new().writer();
rmp_serde::encode::write(&mut bytes, dst).map_err(|err| {
Error::Serialisation(format!(
"could not serialize dst payload with Msgpack: {err}",
))
})?;
Ok(bytes.into_inner().freeze())
}
pub fn serialize_msg_dst(&self) -> Result<Bytes> {
Self::serialize_dst_payload(&self.dst)
}
pub fn single_src_node(name: XorName, dst: Dst, msg: NodeMsg) -> Result<WireMsg> {
let msg_payload = WireMsg::serialize_msg_payload(&msg)
.map_err(|_| Error::Serialisation("Could not serialise node msg".to_string()))?;
let wire_msg = WireMsg::new_msg(
MsgId::new(),
msg_payload,
MsgKind::Node {
name,
is_ae: msg.is_ae(),
is_join: msg.is_join(),
},
dst,
);
Ok(wire_msg)
}
pub fn new_msg(msg_id: MsgId, payload: Bytes, auth: MsgKind, dst: Dst) -> Self {
Self {
header: WireMsgHeader::new(msg_id, auth),
dst,
payload,
serialized_dst: None,
serialized_header: None,
}
}
pub fn from(bytes: UsrMsgBytes) -> Result<Self> {
let (header_bytes, dst_bytes, payload) = bytes;
let header = WireMsgHeader::from(header_bytes.clone())?;
let dst: Dst = rmp_serde::from_slice(&dst_bytes).map_err(|err| {
Error::FailedToParse(format!(
"Message dst couldn't be deserialized from the dst bytes: {err}",
))
})?;
Ok(Self {
header,
dst,
payload,
serialized_dst: Some(dst_bytes),
serialized_header: Some(header_bytes),
})
}
pub fn serialize(&self) -> Result<UsrMsgBytes> {
let header = if let Some(bytes) = &self.serialized_header {
bytes.clone()
} else {
self.header.serialize()?
};
let dst = if let Some(bytes) = &self.serialized_dst {
bytes.clone()
} else {
self.serialize_msg_dst()?
};
Ok((header, dst, self.payload.clone()))
}
pub fn serialize_and_cache_bytes(&mut self) -> Result<UsrMsgBytes> {
let header = if let Some(hdr_bytes) = &self.serialized_header {
hdr_bytes.clone()
} else {
let hdr_bytes = self.header.serialize()?;
self.serialized_header = Some(hdr_bytes.clone());
hdr_bytes
};
let dst = if let Some(dst_bytes) = &self.serialized_dst {
dst_bytes.clone()
} else {
let dst_bytes = self.serialize_msg_dst()?;
self.serialized_dst = Some(dst_bytes.clone());
dst_bytes
};
Ok((header, dst, self.payload.clone()))
}
pub fn serialize_with_new_dst(&self, dst: &Dst) -> Result<UsrMsgBytes> {
let header = if let Some(bytes) = &self.serialized_header {
bytes.clone()
} else {
self.header.serialize()?
};
let dst = Self::serialize_dst_payload(dst)?;
Ok((header, dst, self.payload.clone()))
}
pub fn into_msg(&self) -> Result<MsgType> {
match self.header.msg_envelope.kind.clone() {
MsgKind::Client(auth) => {
let msg: ClientMsg = rmp_serde::from_slice(&self.payload).map_err(|err| {
Error::FailedToParse(format!("Data message payload as Msgpack: {err}"))
})?;
let auth = AuthorityProof::verify(auth, &self.payload)?;
Ok(MsgType::Client {
msg_id: self.header.msg_envelope.msg_id,
auth,
dst: self.dst,
msg,
})
}
MsgKind::ClientDataResponse(_) => {
let msg: ClientDataResponse =
rmp_serde::from_slice(&self.payload).map_err(|err| {
Error::FailedToParse(format!("Data message payload as Msgpack: {err}"))
})?;
Ok(MsgType::ClientDataResponse {
msg_id: self.header.msg_envelope.msg_id,
msg,
})
}
MsgKind::Node { .. } => {
let msg: NodeMsg = rmp_serde::from_slice(&self.payload).map_err(|err| {
Error::FailedToParse(format!("Node signed message payload as Msgpack: {err}"))
})?;
Ok(MsgType::Node {
msg_id: self.header.msg_envelope.msg_id,
dst: self.dst,
msg,
})
}
MsgKind::NodeDataResponse(_) => {
let msg: NodeDataResponse =
rmp_serde::from_slice(&self.payload).map_err(|err| {
Error::FailedToParse(format!("Data message payload as Msgpack: {err}"))
})?;
Ok(MsgType::NodeDataResponse {
msg_id: self.header.msg_envelope.msg_id,
msg,
})
}
}
}
pub fn msg_id(&self) -> MsgId {
self.header.msg_envelope.msg_id
}
pub fn kind(&self) -> &MsgKind {
&self.header.msg_envelope.kind
}
pub fn dst_section_key(&self) -> bls::PublicKey {
self.dst.section_key
}
pub fn dst(&self) -> &Dst {
&self.dst
}
pub fn deserialize(bytes: UsrMsgBytes) -> Result<MsgType> {
Self::from(bytes)?.into_msg()
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::{
messaging::{
data::{ClientMsg, DataQuery, DataQueryVariant},
system::NodeMsg,
AuthorityProof, ClientAuth, MsgId,
},
types::{ChunkAddress, Keypair},
};
use bls::SecretKey;
use eyre::Result;
#[test]
fn serialisation_node_msg() -> Result<()> {
let dst = Dst {
name: xor_name::rand::random(),
section_key: SecretKey::random().public_key(),
};
let msg_id = MsgId::new();
let msg = NodeMsg::HandoverAE(100);
let payload = WireMsg::serialize_msg_payload(&msg)?;
let kind = MsgKind::Node {
name: Default::default(),
is_join: true,
is_ae: false,
};
let wire_msg = WireMsg::new_msg(msg_id, payload, kind, dst);
let serialized = wire_msg.serialize()?;
let deserialized = WireMsg::from(serialized)?;
assert_eq!(deserialized, wire_msg);
assert_eq!(deserialized.msg_id(), wire_msg.msg_id());
assert_eq!(deserialized.dst(), &dst);
assert_eq!(deserialized.dst_section_key(), dst.section_key);
assert_eq!(
deserialized.into_msg()?,
MsgType::Node {
msg_id: wire_msg.msg_id(),
dst,
msg,
}
);
Ok(())
}
#[test]
fn serialisation_client_msg() -> Result<()> {
let src_client_keypair = Keypair::new_ed25519();
let dst = Dst {
name: xor_name::rand::random(),
section_key: SecretKey::random().public_key(),
};
let msg_id = MsgId::new();
let client_msg = ClientMsg::Query(DataQuery {
node_index: 0,
variant: DataQueryVariant::GetChunk(ChunkAddress(xor_name::rand::random())),
});
let payload = WireMsg::serialize_msg_payload(&client_msg)?;
let auth = ClientAuth {
public_key: src_client_keypair.public_key(),
signature: src_client_keypair.sign(&payload),
};
let auth_proof = AuthorityProof::verify(auth.clone(), &payload)?;
let kind = MsgKind::Client(auth);
let wire_msg = WireMsg::new_msg(msg_id, payload, kind, dst);
let serialized = wire_msg.serialize()?;
let deserialized = WireMsg::from(serialized)?;
assert_eq!(deserialized, wire_msg);
assert_eq!(deserialized.msg_id(), wire_msg.msg_id());
assert_eq!(deserialized.dst(), &dst);
assert_eq!(deserialized.dst_section_key(), dst.section_key);
assert_eq!(
deserialized.into_msg()?,
MsgType::Client {
msg_id: wire_msg.msg_id(),
auth: auth_proof,
dst,
msg: client_msg,
}
);
Ok(())
}
}