use tokio::io::{AsyncReadExt, AsyncWriteExt};
use crate::error::{ConnectError, ErrorKind};
use super::IoStream;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct ProtocolHeader {
pub protocol_id: u8,
pub major: u8,
pub minor: u8,
pub revision: u8,
}
impl ProtocolHeader {
pub const AMQP: ProtocolHeader = ProtocolHeader {
protocol_id: 0,
major: 1,
minor: 0,
revision: 0,
};
pub const TLS: ProtocolHeader = ProtocolHeader {
protocol_id: 2,
major: 1,
minor: 0,
revision: 0,
};
pub const SASL: ProtocolHeader = ProtocolHeader {
protocol_id: 3,
major: 1,
minor: 0,
revision: 0,
};
pub fn to_bytes(self) -> [u8; 8] {
[
b'A',
b'M',
b'Q',
b'P',
self.protocol_id,
self.major,
self.minor,
self.revision,
]
}
pub fn from_bytes(bytes: [u8; 8]) -> Result<Self, ConnectError> {
if &bytes[0..4] != b"AMQP" {
return Err(ConnectError::msg(
ErrorKind::ProtocolViolation,
format!("bad protocol header magic: {:02x?}", &bytes[0..4]),
));
}
Ok(ProtocolHeader {
protocol_id: bytes[4],
major: bytes[5],
minor: bytes[6],
revision: bytes[7],
})
}
pub async fn write<S: IoStream>(self, stream: &mut S) -> Result<(), ConnectError> {
stream.write_all(&self.to_bytes()).await?;
stream.flush().await?;
Ok(())
}
pub async fn read<S: IoStream>(stream: &mut S) -> Result<Self, ConnectError> {
let mut buf = [0u8; 8];
stream.read_exact(&mut buf).await?;
ProtocolHeader::from_bytes(buf)
}
pub async fn negotiate<S: IoStream>(self, stream: &mut S) -> Result<(), ConnectError> {
self.write(stream).await?;
let peer = ProtocolHeader::read(stream).await?;
if peer != self {
return Err(ConnectError::msg(
ErrorKind::ProtocolViolation,
format!(
"protocol header mismatch: sent {:?}, peer offered {:?}",
self.to_bytes(),
peer.to_bytes()
),
));
}
Ok(())
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn amqp_header_bytes() {
assert_eq!(ProtocolHeader::AMQP.to_bytes(), *b"AMQP\x00\x01\x00\x00");
assert_eq!(ProtocolHeader::SASL.to_bytes(), *b"AMQP\x03\x01\x00\x00");
assert_eq!(ProtocolHeader::TLS.to_bytes(), *b"AMQP\x02\x01\x00\x00");
}
#[test]
fn roundtrip_and_magic_check() {
let h = ProtocolHeader::AMQP;
assert_eq!(ProtocolHeader::from_bytes(h.to_bytes()).unwrap(), h);
assert!(ProtocolHeader::from_bytes(*b"XXXX\x00\x01\x00\x00").is_err());
}
#[tokio::test]
async fn negotiate_over_duplex_matches() {
let (mut a, mut b) = tokio::io::duplex(64);
let server = tokio::spawn(async move {
let h = ProtocolHeader::read(&mut b).await.unwrap();
h.write(&mut b).await.unwrap();
});
ProtocolHeader::AMQP.negotiate(&mut a).await.unwrap();
server.await.unwrap();
}
}