use bytes::{Buf, BufMut, BytesMut};
use std::io;
use tokio_util::codec::{Decoder, Encoder};
use crate::frame::Frame;
use crate::parser::{parse_frame_slice, unescape_header_value};
fn escape_header_value(input: &str) -> String {
let mut result = String::with_capacity(input.len());
for ch in input.chars() {
match ch {
'\\' => result.push_str("\\\\"),
'\r' => result.push_str("\\r"),
'\n' => result.push_str("\\n"),
':' => result.push_str("\\c"),
_ => result.push(ch),
}
}
result
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum StompItem {
Frame(Frame),
Heartbeat,
}
pub struct StompCodec {
}
impl StompCodec {
pub fn new() -> Self {
Self {}
}
}
impl Default for StompCodec {
fn default() -> Self {
Self::new()
}
}
impl Decoder for StompCodec {
type Item = StompItem;
type Error = io::Error;
fn decode(&mut self, src: &mut BytesMut) -> Result<Option<Self::Item>, Self::Error> {
if let Some(&b'\n') = src.chunk().first() {
src.advance(1);
return Ok(Some(StompItem::Heartbeat));
}
let chunk = src.chunk();
match parse_frame_slice(chunk) {
Ok(Some((cmd_bytes, headers, body, consumed))) => {
src.advance(consumed);
let command = String::from_utf8(cmd_bytes).map_err(|e| {
io::Error::new(
io::ErrorKind::InvalidData,
format!("invalid utf8 in command: {}", e),
)
})?;
let mut hdrs: Vec<(String, String)> = Vec::new();
for (k, v) in headers {
let k_unescaped = unescape_header_value(&k).map_err(|e| {
io::Error::new(
io::ErrorKind::InvalidData,
format!("invalid escape in header key: {}", e),
)
})?;
let ks = String::from_utf8(k_unescaped).map_err(|e| {
io::Error::new(
io::ErrorKind::InvalidData,
format!("invalid utf8 in header key: {}", e),
)
})?;
let v_unescaped = unescape_header_value(&v).map_err(|e| {
io::Error::new(
io::ErrorKind::InvalidData,
format!("invalid escape in header value: {}", e),
)
})?;
let vs = String::from_utf8(v_unescaped).map_err(|e| {
io::Error::new(
io::ErrorKind::InvalidData,
format!("invalid utf8 in header value: {}", e),
)
})?;
hdrs.push((ks, vs));
}
let body = body.unwrap_or_default();
let frame = Frame {
command,
headers: hdrs,
body,
};
Ok(Some(StompItem::Frame(frame)))
}
Ok(None) => Ok(None),
Err(e) => Err(io::Error::new(
io::ErrorKind::InvalidData,
format!("parse error: {}", e),
)),
}
}
}
impl Encoder<StompItem> for StompCodec {
type Error = io::Error;
fn encode(&mut self, item: StompItem, dst: &mut BytesMut) -> Result<(), Self::Error> {
match item {
StompItem::Heartbeat => {
dst.put_u8(b'\n');
}
StompItem::Frame(frame) => {
dst.extend_from_slice(frame.command.as_bytes());
dst.put_u8(b'\n');
let mut headers = frame.headers;
let has_cl = headers
.iter()
.any(|(k, _)| k.to_lowercase() == "content-length");
if !has_cl {
let include_cl =
frame.body.contains(&0) || std::str::from_utf8(&frame.body).is_err();
if include_cl {
headers.push(("content-length".to_string(), frame.body.len().to_string()));
}
}
for (k, v) in headers {
let escaped_key = escape_header_value(&k);
let escaped_val = escape_header_value(&v);
dst.extend_from_slice(escaped_key.as_bytes());
dst.put_u8(b':');
dst.extend_from_slice(escaped_val.as_bytes());
dst.put_u8(b'\n');
}
dst.put_slice(b"\n");
dst.extend_from_slice(&frame.body);
dst.put_u8(0);
}
}
Ok(())
}
}