use bytes::{Buf, BytesMut};
use std::collections::HashMap;
use tokio_util::codec::{Decoder, Encoder};
use crate::error::AmiError;
const MAX_MESSAGE_SIZE: usize = 64 * 1024;
const MAX_HEADERS: usize = 512;
#[derive(Debug, Clone, PartialEq)]
pub struct RawAmiMessage {
pub headers: Vec<(String, String)>,
pub output: Vec<String>,
pub channel_variables: HashMap<String, String>,
}
impl RawAmiMessage {
pub fn get(&self, key: &str) -> Option<&str> {
self.headers
.iter()
.find(|(k, _)| k.eq_ignore_ascii_case(key))
.map(|(_, v)| v.as_str())
}
pub fn get_all(&self, key: &str) -> Vec<&str> {
self.headers
.iter()
.filter(|(k, _)| k.eq_ignore_ascii_case(key))
.map(|(_, v)| v.as_str())
.collect()
}
pub fn is_response(&self) -> bool {
self.get("Response").is_some()
}
pub fn is_event(&self) -> bool {
self.get("Event").is_some()
}
pub fn get_variable(&self, name: &str) -> Option<&str> {
self.channel_variables.get(name).map(|s| s.as_str())
}
pub fn to_map(&self) -> HashMap<String, String> {
self.headers
.iter()
.map(|(k, v)| (k.clone(), v.clone()))
.collect()
}
}
#[derive(Debug)]
pub struct AmiCodec {
banner_consumed: bool,
}
impl AmiCodec {
pub fn new() -> Self {
Self {
banner_consumed: false,
}
}
}
impl Default for AmiCodec {
fn default() -> Self {
Self::new()
}
}
impl Decoder for AmiCodec {
type Item = RawAmiMessage;
type Error = AmiError;
fn decode(&mut self, src: &mut BytesMut) -> Result<Option<Self::Item>, Self::Error> {
if !self.banner_consumed {
if let Some(pos) = find_crlf(src) {
let line = &src[..pos];
if !line.starts_with(b"Asterisk Call Manager") {
let preview = String::from_utf8_lossy(&line[..line.len().min(64)]);
return Err(AmiError::Protocol(
asterisk_rs_core::error::ProtocolError::MalformedMessage {
details: format!("expected AMI banner, got: {}", preview),
},
));
}
src.advance(pos + 2); self.banner_consumed = true;
} else {
return Ok(None); }
}
const END_MARKER: &[u8] = b"--END COMMAND--";
loop {
let first_blank = match find_double_crlf(src) {
Some(pos) => pos,
None => return Ok(None),
};
let frame_end = if is_follows_response(&src[..first_blank]) {
match find_subsequence(src, END_MARKER) {
Some(marker_pos) => {
let after_marker = marker_pos + END_MARKER.len();
if src.len() < after_marker + 2 {
return Ok(None);
}
if &src[after_marker..after_marker + 2] != b"\r\n" {
return Ok(None);
}
after_marker + 2
}
None => return Ok(None),
}
} else {
first_blank + 4
};
if frame_end > MAX_MESSAGE_SIZE {
return Err(AmiError::Protocol(
asterisk_rs_core::error::ProtocolError::MalformedMessage {
details: format!("message exceeds {} byte limit", MAX_MESSAGE_SIZE),
},
));
}
let message_bytes = &src[..frame_end];
let mut headers = Vec::new();
let mut output = Vec::new();
let mut channel_variables = HashMap::new();
for line in message_bytes.split(|&b| b == b'\n') {
let line = line.strip_suffix(b"\r").unwrap_or(line);
if line.is_empty() {
continue;
}
if line == END_MARKER {
continue;
}
if let Some(colon_pos) = line.iter().position(|&b| b == b':') {
if headers.len() + channel_variables.len() >= MAX_HEADERS {
return Err(AmiError::Protocol(
asterisk_rs_core::error::ProtocolError::MalformedMessage {
details: format!("message exceeds {} header limit", MAX_HEADERS),
},
));
}
let key = String::from_utf8_lossy(&line[..colon_pos])
.trim()
.to_string();
let value_start = colon_pos + 1;
let value = if value_start < line.len() {
String::from_utf8_lossy(&line[value_start..])
.trim()
.to_string()
} else {
String::new()
};
if let Some(var_name) = key
.strip_prefix("ChanVariable(")
.and_then(|s| s.strip_suffix(')'))
{
channel_variables.insert(var_name.to_string(), value);
} else {
headers.push((key, value));
}
} else {
output.push(String::from_utf8_lossy(line).into_owned());
}
}
src.advance(frame_end);
if headers.is_empty() {
continue;
}
return Ok(Some(RawAmiMessage {
headers,
output,
channel_variables,
}));
}
}
}
impl Encoder<RawAmiMessage> for AmiCodec {
type Error = AmiError;
fn encode(&mut self, item: RawAmiMessage, dst: &mut BytesMut) -> Result<(), Self::Error> {
let contains_line_terminator = |s: &str| s.bytes().any(|b| b == b'\r' || b == b'\n');
for (key, value) in &item.headers {
if contains_line_terminator(key) {
return Err(AmiError::Protocol(
asterisk_rs_core::error::ProtocolError::MalformedMessage {
details: format!("header key contains illegal line terminator: {:?}", key),
},
));
}
if contains_line_terminator(value) {
return Err(AmiError::Protocol(
asterisk_rs_core::error::ProtocolError::MalformedMessage {
details: "header value contains illegal line terminator".to_owned(),
},
));
}
}
for (name, value) in &item.channel_variables {
if contains_line_terminator(name) {
return Err(AmiError::Protocol(
asterisk_rs_core::error::ProtocolError::MalformedMessage {
details: format!(
"channel variable name contains illegal line terminator: {:?}",
name
),
},
));
}
if contains_line_terminator(value) {
return Err(AmiError::Protocol(
asterisk_rs_core::error::ProtocolError::MalformedMessage {
details: "channel variable value contains illegal line terminator"
.to_owned(),
},
));
}
}
for (key, value) in &item.headers {
dst.extend_from_slice(key.as_bytes());
dst.extend_from_slice(b": ");
dst.extend_from_slice(value.as_bytes());
dst.extend_from_slice(b"\r\n");
}
for (name, value) in &item.channel_variables {
dst.extend_from_slice(b"ChanVariable(");
dst.extend_from_slice(name.as_bytes());
dst.extend_from_slice(b"): ");
dst.extend_from_slice(value.as_bytes());
dst.extend_from_slice(b"\r\n");
}
dst.extend_from_slice(b"\r\n"); Ok(())
}
}
fn find_crlf(buf: &[u8]) -> Option<usize> {
buf.windows(2).position(|w| w == b"\r\n")
}
fn find_double_crlf(buf: &[u8]) -> Option<usize> {
buf.windows(4).position(|w| w == b"\r\n\r\n")
}
fn is_follows_response(header_bytes: &[u8]) -> bool {
header_bytes.split(|&b| b == b'\n').any(|line| {
let line = line.strip_suffix(b"\r").unwrap_or(line);
if let Some(colon_pos) = line.iter().position(|&b| b == b':') {
let key = &line[..colon_pos];
let value = &line[colon_pos + 1..];
let value_trimmed = value.strip_prefix(b" ").unwrap_or(value);
key.eq_ignore_ascii_case(b"response") && value_trimmed.eq_ignore_ascii_case(b"follows")
} else {
false
}
})
}
fn find_subsequence(haystack: &[u8], needle: &[u8]) -> Option<usize> {
haystack.windows(needle.len()).position(|w| w == needle)
}