use crate::error::ZmqError;
use crate::message::{Msg, MsgFlags};
use crate::protocol::zmtp::{ZmtpCodec, manual_parser::ZmtpManualParser};
use crate::security::IDataCipher;
use bytes::{Buf, BufMut, Bytes, BytesMut};
use tokio_util::codec::Encoder;
pub(crate) trait ISecureFramer: Send + Sync + 'static {
fn try_read_msg(&mut self, network_buffer: &mut BytesMut) -> Result<Option<Msg>, ZmqError>;
fn write_msg_multipart(&mut self, msgs: Vec<Msg>) -> Result<Bytes, ZmqError>;
fn write_msg_batch(&mut self, batch: &[Vec<Msg>]) -> Result<Bytes, ZmqError>;
fn is_passthrough(&self) -> bool {
false
}
fn write_msg_split(&mut self, msg: Msg) -> Result<(Bytes, Option<Bytes>), ZmqError> {
let merged = self.write_msg_multipart(vec![msg])?;
Ok((merged, None))
}
fn frame_vectored(&mut self, batch: &[Vec<Msg>]) -> Result<Vec<Bytes>, ZmqError> {
Ok(vec![self.write_msg_batch(batch)?])
}
fn try_read_msgs_from_bytes(
&mut self,
data: Bytes,
accumulator: &mut BytesMut,
) -> Result<Vec<Msg>, ZmqError> {
accumulator.extend_from_slice(&data);
let mut msgs = Vec::new();
while let Some(msg) = self.try_read_msg(accumulator)? {
msgs.push(msg);
}
Ok(msgs)
}
}
pub(crate) struct NullFramer {
parser: ZmtpManualParser,
coalesce_buffer: BytesMut,
header_slab: BytesMut,
}
impl NullFramer {
pub(crate) fn new(max_msg_size: i64) -> Self {
Self {
parser: ZmtpManualParser::new(max_msg_size),
coalesce_buffer: BytesMut::with_capacity(65536),
header_slab: BytesMut::with_capacity(4096),
}
}
}
impl ISecureFramer for NullFramer {
fn is_passthrough(&self) -> bool {
true
}
fn try_read_msg(&mut self, network_buffer: &mut BytesMut) -> Result<Option<Msg>, ZmqError> {
self.parser.decode_from_buffer(network_buffer)
}
fn write_msg_multipart(&mut self, msgs: Vec<Msg>) -> Result<Bytes, ZmqError> {
let mut codec = ZmtpCodec::new();
let mut buffer = BytesMut::new();
for msg in msgs {
codec.encode(msg, &mut buffer)?;
}
Ok(buffer.freeze())
}
fn write_msg_batch(&mut self, batch: &[Vec<Msg>]) -> Result<Bytes, ZmqError> {
self.coalesce_buffer.clear();
let mut codec = ZmtpCodec::new();
for msgs in batch {
for msg in msgs {
codec.encode(msg.clone(), &mut self.coalesce_buffer)?;
}
}
Ok(self.coalesce_buffer.split().freeze())
}
fn write_msg_split(&mut self, msg: Msg) -> Result<(Bytes, Option<Bytes>), ZmqError> {
let payload = msg.data_bytes().unwrap_or_default();
let payload_len = payload.len();
let is_more = msg.flags().contains(MsgFlags::MORE);
let mut hdr = BytesMut::with_capacity(9);
if payload_len <= 255 {
hdr.put_u8(if is_more { 0x01 } else { 0x00 }); hdr.put_u8(payload_len as u8);
} else {
hdr.put_u8(if is_more { 0x03 } else { 0x02 }); hdr.put_u64(payload_len as u64);
}
Ok((hdr.freeze(), Some(payload)))
}
fn frame_vectored(&mut self, batch: &[Vec<Msg>]) -> Result<Vec<Bytes>, ZmqError> {
let total_msgs: usize = batch.iter().map(|g| g.len()).sum();
let mut out = Vec::with_capacity(total_msgs * 2);
for group in batch {
for msg in group {
let payload = msg.data_bytes().unwrap_or_default();
let payload_len = payload.len();
let is_more = msg.flags().contains(MsgFlags::MORE);
self.header_slab.reserve(9);
if payload_len <= 255 {
self.header_slab.put_u8(if is_more { 0x01 } else { 0x00 });
self.header_slab.put_u8(payload_len as u8);
} else {
self.header_slab.put_u8(if is_more { 0x03 } else { 0x02 });
self.header_slab.put_u64(payload_len as u64);
}
out.push(self.header_slab.split().freeze());
if !payload.is_empty() {
out.push(payload); }
}
}
Ok(out)
}
}
pub(crate) struct LengthPrefixedFramer {
cipher: Box<dyn IDataCipher>,
parser: ZmtpManualParser,
decrypted_buffer: BytesMut,
coalesce_buffer: BytesMut,
}
impl LengthPrefixedFramer {
pub(crate) fn new(cipher: Box<dyn IDataCipher>, max_msg_size: i64) -> Self {
Self {
cipher,
parser: ZmtpManualParser::new(max_msg_size),
decrypted_buffer: BytesMut::with_capacity(65536 * 2),
coalesce_buffer: BytesMut::with_capacity(65536),
}
}
}
impl ISecureFramer for LengthPrefixedFramer {
fn try_read_msg(&mut self, network_buffer: &mut BytesMut) -> Result<Option<Msg>, ZmqError> {
loop {
if let Some(msg) = self.parser.decode_from_buffer(&mut self.decrypted_buffer)? {
return Ok(Some(msg));
}
if network_buffer.len() < 2 {
return Ok(None); }
let len = network_buffer.as_ref().get_u16() as usize;
if network_buffer.len() < 2 + len {
return Ok(None); }
network_buffer.advance(2); let encrypted_frame = network_buffer.split_to(len);
let plaintext = self.cipher.decrypt(&encrypted_frame)?;
self.decrypted_buffer.extend_from_slice(&plaintext);
}
}
fn write_msg_multipart(&mut self, msgs: Vec<Msg>) -> Result<Bytes, ZmqError> {
let mut codec = ZmtpCodec::new();
let mut plaintext_buffer = BytesMut::new();
for msg in msgs {
codec.encode(msg, &mut plaintext_buffer)?;
}
let ciphertext = self.cipher.encrypt(&plaintext_buffer)?;
let mut final_buffer = BytesMut::with_capacity(2 + ciphertext.len());
final_buffer.put_u16(ciphertext.len() as u16);
final_buffer.extend_from_slice(&ciphertext);
Ok(final_buffer.freeze())
}
fn write_msg_batch(&mut self, batch: &[Vec<Msg>]) -> Result<Bytes, ZmqError> {
self.coalesce_buffer.clear();
let mut codec = ZmtpCodec::new();
for msgs in batch {
for msg in msgs {
codec.encode(msg.clone(), &mut self.coalesce_buffer)?;
}
}
let ciphertext = self.cipher.encrypt(&self.coalesce_buffer)?;
let mut out = BytesMut::with_capacity(2 + ciphertext.len());
out.put_u16(ciphertext.len() as u16);
out.extend_from_slice(&ciphertext);
Ok(out.freeze())
}
}