use crate::error::ZmqError;
use bytes::{BufMut, BytesMut};
use std::convert::TryInto;
use tracing;
pub const GREETING_LENGTH: usize = 64;
pub const MECHANISM_LENGTH: usize = 20;
pub const SIGNATURE_LENGTH: usize = 10;
pub const GREETING_VERSION_MAJOR_BYTE: u8 = 0x03;
pub const GREETING_VERSION_MINOR_BYTE: u8 = 0x00;
pub const V2_REVISION: u8 = 0x01;
pub const V3_REVISION: u8 = 0x03;
pub const V2_SOCKET_TYPE_PAIR: u8 = 0;
pub const V2_SOCKET_TYPE_PUB: u8 = 1;
pub const V2_SOCKET_TYPE_SUB: u8 = 2;
pub const V2_SOCKET_TYPE_REQ: u8 = 3;
pub const V2_SOCKET_TYPE_REP: u8 = 4;
pub const V2_SOCKET_TYPE_DEALER: u8 = 5; pub const V2_SOCKET_TYPE_ROUTER: u8 = 6; pub const V2_SOCKET_TYPE_PULL: u8 = 7;
pub const V2_SOCKET_TYPE_PUSH: u8 = 8;
pub const V2_SOCKET_TYPE_XPUB: u8 = 9;
pub const V2_SOCKET_TYPE_XSUB: u8 = 10;
const VERSION_MAJOR_OFFSET: usize = 10;
const VERSION_MINOR_OFFSET: usize = 11;
pub const MECHANISM_OFFSET: usize = 12;
pub const AS_SERVER_OFFSET: usize = MECHANISM_OFFSET + MECHANISM_LENGTH; const PADDING_OFFSET: usize = AS_SERVER_OFFSET + 1; const PADDING_LENGTH: usize = GREETING_LENGTH - PADDING_OFFSET;
pub fn encode_signature(buffer: &mut BytesMut) {
buffer.reserve(SIGNATURE_LENGTH);
buffer.put_u8(0xFF);
buffer.put_bytes(0, 8);
buffer.put_u8(0x7F);
}
pub fn encode_v3_tail(mechanism: &[u8; MECHANISM_LENGTH], as_server: bool, buffer: &mut BytesMut) {
buffer.reserve(GREETING_LENGTH - SIGNATURE_LENGTH - 1);
buffer.put_u8(GREETING_VERSION_MINOR_BYTE);
buffer.put_slice(mechanism);
buffer.put_u8(as_server as u8);
buffer.put_bytes(0, PADDING_LENGTH);
}
pub fn peek_revision(buf: &[u8]) -> Result<u8, ZmqError> {
if buf.len() < SIGNATURE_LENGTH + 1 {
return Err(ZmqError::ProtocolViolation(
"Greeting too short to determine protocol revision".into(),
));
}
if buf[0] != 0xFF || buf[SIGNATURE_LENGTH - 1] != 0x7F {
return Err(ZmqError::ProtocolViolation(
"Invalid ZMTP signature in greeting".into(),
));
}
Ok(buf[VERSION_MAJOR_OFFSET])
}
pub fn socket_type_code(name: &str) -> Option<u8> {
Some(match name {
"PAIR" => V2_SOCKET_TYPE_PAIR,
"PUB" => V2_SOCKET_TYPE_PUB,
"SUB" => V2_SOCKET_TYPE_SUB,
"REQ" => V2_SOCKET_TYPE_REQ,
"REP" => V2_SOCKET_TYPE_REP,
"DEALER" => V2_SOCKET_TYPE_DEALER,
"ROUTER" => V2_SOCKET_TYPE_ROUTER,
"PULL" => V2_SOCKET_TYPE_PULL,
"PUSH" => V2_SOCKET_TYPE_PUSH,
"XPUB" => V2_SOCKET_TYPE_XPUB,
"XSUB" => V2_SOCKET_TYPE_XSUB,
_ => return None,
})
}
pub fn socket_type_name_from_code(code: u8) -> Option<&'static str> {
Some(match code {
V2_SOCKET_TYPE_PAIR => "PAIR",
V2_SOCKET_TYPE_PUB => "PUB",
V2_SOCKET_TYPE_SUB => "SUB",
V2_SOCKET_TYPE_REQ => "REQ",
V2_SOCKET_TYPE_REP => "REP",
V2_SOCKET_TYPE_DEALER => "DEALER",
V2_SOCKET_TYPE_ROUTER => "ROUTER",
V2_SOCKET_TYPE_PULL => "PULL",
V2_SOCKET_TYPE_PUSH => "PUSH",
V2_SOCKET_TYPE_XPUB => "XPUB",
V2_SOCKET_TYPE_XSUB => "XSUB",
_ => return None,
})
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct ZmtpGreeting {
pub version: (u8, u8),
pub mechanism: [u8; MECHANISM_LENGTH], pub as_server: bool,
}
impl ZmtpGreeting {
pub fn encode(mechanism: &[u8; MECHANISM_LENGTH], as_server: bool, buffer: &mut BytesMut) {
buffer.reserve(GREETING_LENGTH);
encode_signature(buffer);
buffer.put_u8(GREETING_VERSION_MAJOR_BYTE);
buffer.put_u8(GREETING_VERSION_MINOR_BYTE);
buffer.put_slice(mechanism);
buffer.put_u8(as_server as u8);
let current_len = buffer.len();
let padding_len = GREETING_LENGTH - current_len;
if padding_len > 0 {
buffer.put_bytes(0, padding_len);
}
debug_assert_eq!(buffer.len(), GREETING_LENGTH);
}
pub fn decode(buffer: &mut BytesMut) -> Result<Option<Self>, ZmqError> {
if buffer.len() < GREETING_LENGTH {
return Ok(None); }
let data = buffer.split_to(GREETING_LENGTH);
if data[0] != 0xFF {
tracing::error!("Greeting does not start with 0xFF (got {:#04x})", data[0]);
return Err(ZmqError::ProtocolViolation(
"Greeting does not start with 0xFF".into(),
));
}
for i in 0..PADDING_LENGTH {
let idx = PADDING_OFFSET + i;
if data[idx] != 0x00 {
tracing::error!(
"Invalid ZMTP greeting: non-zero padding byte at index {}: {:#04x}",
idx,
data[idx]
);
return Err(ZmqError::ProtocolViolation(
"Non-zero byte in greeting padding".into(),
));
}
}
let major_version = data[VERSION_MAJOR_OFFSET];
let minor_version = data[VERSION_MINOR_OFFSET];
if major_version != GREETING_VERSION_MAJOR_BYTE {
return Err(ZmqError::ProtocolViolation(format!(
"Unsupported ZMTP major version {}.{}",
major_version, minor_version
)));
}
let version = (major_version, minor_version);
let mechanism_slice = &data[MECHANISM_OFFSET..MECHANISM_OFFSET + MECHANISM_LENGTH];
let mechanism: [u8; MECHANISM_LENGTH] = mechanism_slice.try_into().unwrap();
let as_server_byte = data[AS_SERVER_OFFSET];
let as_server = match as_server_byte {
0x00 => false,
0x01 => true,
_ => return Err(ZmqError::ProtocolViolation("Invalid as-server flag".into())),
};
tracing::debug!(?version, mechanism_name = %std::str::from_utf8(&mechanism).unwrap_or("").trim_end_matches('\0'), as_server, "Parsed ZMTP Greeting");
Ok(Some(Self {
version,
mechanism,
as_server,
}))
}
pub fn mechanism_name(&self) -> &str {
let first_null = self
.mechanism
.iter()
.position(|&b| b == 0)
.unwrap_or(MECHANISM_LENGTH);
std::str::from_utf8(&self.mechanism[..first_null]).unwrap_or("<invalid_utf8>")
}
}
#[cfg(test)]
mod additional_greeting_tests {
use super::*;
use bytes::BytesMut;
#[test]
fn test_decode_valid_greeting() {
let mut buf = BytesMut::new();
let mechanism = [0u8; MECHANISM_LENGTH];
ZmtpGreeting::encode(&mechanism, true, &mut buf);
let result = ZmtpGreeting::decode(&mut buf)
.expect("Decoding valid greeting should succeed")
.expect("Should return Some(ZmtpGreeting)");
assert_eq!(result.version, (3, 0));
assert_eq!(result.mechanism, mechanism);
assert!(result.as_server);
}
#[test]
fn test_decode_invalid_signature() {
let mut buf = BytesMut::zeroed(64);
buf[0] = 0x00;
let result = ZmtpGreeting::decode(&mut buf);
assert!(matches!(result, Err(ZmqError::ProtocolViolation(_))));
}
#[test]
fn test_decode_dirty_padding() {
let mut buf = BytesMut::new();
let mechanism = [0u8; MECHANISM_LENGTH];
ZmtpGreeting::encode(&mechanism, false, &mut buf);
buf[45] = 0xAA;
let result = ZmtpGreeting::decode(&mut buf);
assert!(matches!(result, Err(ZmqError::ProtocolViolation(_))));
}
#[test]
fn test_decode_unsupported_version() {
let mut buf = BytesMut::new();
let mechanism = [0u8; MECHANISM_LENGTH];
ZmtpGreeting::encode(&mechanism, false, &mut buf);
buf[10] = 2;
let result = ZmtpGreeting::decode(&mut buf);
assert!(matches!(result, Err(ZmqError::ProtocolViolation(_))));
}
#[test]
fn test_decode_invalid_as_server_flag() {
let mut buf = BytesMut::new();
let mechanism = [0u8; MECHANISM_LENGTH];
ZmtpGreeting::encode(&mechanism, false, &mut buf);
buf[32] = 0x02;
let result = ZmtpGreeting::decode(&mut buf);
assert!(matches!(result, Err(ZmqError::ProtocolViolation(_))));
}
#[test]
fn test_encode_signature_is_canonical() {
let mut buf = BytesMut::new();
encode_signature(&mut buf);
assert_eq!(buf.len(), SIGNATURE_LENGTH);
assert_eq!(buf[0], 0xFF);
assert!(buf[1..9].iter().all(|&b| b == 0));
assert_eq!(buf[9], 0x7F);
}
#[test]
fn test_peek_revision_v3() {
let mut buf = BytesMut::new();
encode_signature(&mut buf);
buf.put_u8(V3_REVISION);
assert_eq!(peek_revision(&buf).unwrap(), V3_REVISION);
}
#[test]
fn test_peek_revision_v2() {
let mut buf = BytesMut::new();
encode_signature(&mut buf);
buf.put_u8(V2_REVISION);
assert_eq!(peek_revision(&buf).unwrap(), V2_REVISION);
}
#[test]
fn test_peek_revision_too_short() {
let mut buf = BytesMut::new();
encode_signature(&mut buf); assert!(matches!(peek_revision(&buf), Err(ZmqError::ProtocolViolation(_))));
}
#[test]
fn test_peek_revision_bad_signature() {
let mut buf = BytesMut::zeroed(11);
buf[0] = 0x00; assert!(matches!(peek_revision(&buf), Err(ZmqError::ProtocolViolation(_))));
}
#[test]
fn test_socket_type_code_roundtrip() {
for name in [
"PAIR", "PUB", "SUB", "REQ", "REP", "DEALER", "ROUTER", "PULL", "PUSH", "XPUB", "XSUB",
] {
let code = socket_type_code(name).expect("known type has a code");
assert_eq!(socket_type_name_from_code(code), Some(name));
}
assert!(socket_type_code("UNKNOWN").is_none());
assert!(socket_type_name_from_code(200).is_none());
}
#[test]
fn test_encode_v3_tail_completes_greeting() {
let mut buf = BytesMut::new();
encode_signature(&mut buf);
buf.put_u8(V3_REVISION);
encode_v3_tail(&[0u8; MECHANISM_LENGTH], true, &mut buf);
assert_eq!(buf.len(), GREETING_LENGTH);
let decoded = ZmtpGreeting::decode(&mut buf).unwrap().unwrap();
assert_eq!(decoded.version.0, V3_REVISION);
assert!(decoded.as_server);
}
}