use bytes::BytesMut;
use std::collections::HashMap;
use std::sync::Arc;
use std::time::Duration;
use tokio::io::{AsyncRead, AsyncReadExt, AsyncWrite, AsyncWriteExt};
use tokio::sync::{Mutex, broadcast, mpsc, oneshot};
use crate::error::ClientError;
use imap_core::ast::{Response, Status};
use imap_core::parser::{MAX_LITERAL_SIZE, parse_response};
type TaggedReply = oneshot::Sender<Result<Vec<u8>, ClientError>>;
type PendingCommands = Arc<Mutex<HashMap<String, TaggedReply>>>;
const EVENT_CHANNEL_CAP: usize = 1024;
const MAX_FRAME_SIZE: usize = MAX_LITERAL_SIZE + 64 * 1024;
enum WriteRequest {
Command {
bytes: Vec<u8>,
tag: String,
reply_tx: TaggedReply,
},
Raw { bytes: Vec<u8> },
}
pub struct RawClient {
write_tx: mpsc::Sender<WriteRequest>,
event_tx: broadcast::Sender<Vec<u8>>,
tag_counter: u64,
pub default_timeout: Duration,
}
impl RawClient {
pub fn new<S>(stream: S) -> Self
where
S: AsyncRead + AsyncWrite + Send + 'static,
{
let (read_half, write_half) = tokio::io::split(stream);
let (write_tx, write_rx) = mpsc::channel(32);
let (event_tx, _) = broadcast::channel(EVENT_CHANNEL_CAP);
let pending_commands = Arc::new(Mutex::new(HashMap::new()));
tokio::spawn(read_loop(
read_half,
Arc::clone(&pending_commands),
event_tx.clone(),
MAX_FRAME_SIZE,
));
tokio::spawn(write_loop(
write_half,
write_rx,
Arc::clone(&pending_commands),
));
Self {
write_tx,
event_tx,
tag_counter: 1,
default_timeout: Duration::from_secs(30),
}
}
pub fn events(&self) -> broadcast::Receiver<Vec<u8>> {
self.event_tx.subscribe()
}
fn next_tag(&mut self) -> String {
let tag = format!("A{:04}", self.tag_counter);
self.tag_counter = self.tag_counter.wrapping_add(1);
tag
}
pub async fn execute_command(&mut self, cmd: &str) -> Result<Vec<u8>, ClientError> {
self.execute_command_with_timeout(cmd, self.default_timeout)
.await
}
pub async fn execute_command_with_timeout(
&mut self,
cmd: &str,
timeout: Duration,
) -> Result<Vec<u8>, ClientError> {
let (_tag, rx) = self.send_command_async(cmd).await?;
match tokio::time::timeout(timeout, rx).await {
Ok(Ok(res)) => res,
Ok(Err(_)) => Err(ClientError::ConnectionClosed),
Err(_) => Err(ClientError::Timeout),
}
}
pub async fn send_command_async(
&mut self,
cmd: &str,
) -> Result<(String, oneshot::Receiver<Result<Vec<u8>, ClientError>>), ClientError> {
let tag = self.next_tag();
let bytes = format!("{} {}\r\n", tag, cmd).into_bytes();
let (reply_tx, reply_rx) = oneshot::channel();
self.write_tx
.send(WriteRequest::Command {
bytes,
tag: tag.clone(),
reply_tx,
})
.await
.map_err(|_| ClientError::ConnectionClosed)?;
Ok((tag, reply_rx))
}
pub async fn send_raw(&mut self, bytes: Vec<u8>) -> Result<(), ClientError> {
self.write_tx
.send(WriteRequest::Raw { bytes })
.await
.map_err(|_| ClientError::ConnectionClosed)
}
pub fn writer(&self) -> WriterHandle {
WriterHandle {
write_tx: self.write_tx.clone(),
}
}
}
#[derive(Clone)]
pub struct WriterHandle {
write_tx: mpsc::Sender<WriteRequest>,
}
impl WriterHandle {
pub async fn send_raw(&self, bytes: Vec<u8>) -> Result<(), ClientError> {
self.write_tx
.send(WriteRequest::Raw { bytes })
.await
.map_err(|_| ClientError::ConnectionClosed)
}
}
async fn write_loop<W>(
mut write_half: W,
mut rx: mpsc::Receiver<WriteRequest>,
pending_commands: PendingCommands,
) where
W: AsyncWrite + Unpin,
{
while let Some(req) = rx.recv().await {
match req {
WriteRequest::Command {
bytes,
tag,
reply_tx,
} => {
pending_commands.lock().await.insert(tag, reply_tx);
if write_half.write_all(&bytes).await.is_err() {
break;
}
}
WriteRequest::Raw { bytes } => {
if write_half.write_all(&bytes).await.is_err() {
break;
}
}
}
}
}
async fn fail_all_pending(pending: &PendingCommands, make_err: impl Fn() -> ClientError) {
let mut map = pending.lock().await;
for (_, tx) in map.drain() {
let _ = tx.send(Err(make_err()));
}
}
async fn read_loop<R>(
mut read_half: R,
pending_commands: PendingCommands,
event_tx: broadcast::Sender<Vec<u8>>,
max_frame_size: usize,
) where
R: AsyncRead + Unpin,
{
let mut buffer = BytesMut::with_capacity(8192);
loop {
match read_half.read_buf(&mut buffer).await {
Ok(0) => {
fail_all_pending(&pending_commands, || ClientError::ConnectionClosed).await;
break;
}
Ok(_) => {
while !buffer.is_empty() {
let routing = match parse_response(&buffer) {
Ok((remaining, response)) => {
let consumed = buffer.len() - remaining.len();
let routing = match &response {
Response::Status(s) => s
.tag
.map(|tag| (tag.to_string(), s.status, s.text.to_string())),
_ => None,
};
(consumed, routing)
}
Err(imap_core::error::ParseError::Incomplete) => break,
Err(_) => {
buffer.clear();
break;
}
};
let (consumed, routing) = routing;
let frame = buffer.split_to(consumed).to_vec();
dispatch_frame(routing, frame, &pending_commands, &event_tx).await;
}
if buffer.len() > max_frame_size {
fail_all_pending(&pending_commands, || ClientError::FrameTooLarge {
max: max_frame_size,
})
.await;
break;
}
}
Err(_) => {
fail_all_pending(&pending_commands, || ClientError::ConnectionClosed).await;
break;
}
}
}
}
async fn dispatch_frame(
routing: Option<(String, Status, String)>,
frame: Vec<u8>,
pending_commands: &PendingCommands,
event_tx: &broadcast::Sender<Vec<u8>>,
) {
if let Some((tag, status, text)) = routing {
let mut map = pending_commands.lock().await;
if let Some(tx) = map.remove(&tag) {
let result = match status {
Status::Ok => Ok(frame),
Status::No | Status::Bad => Err(ClientError::CommandFailed(text)),
Status::Bye => Err(ClientError::ConnectionClosed),
Status::PreAuth => Err(ClientError::CommandFailed(text)),
};
let _ = tx.send(result);
return;
}
}
let _ = event_tx.send(frame);
}
#[cfg(test)]
mod tests {
use super::*;
use std::time::Duration;
use tokio::io::{AsyncReadExt, AsyncWriteExt, duplex};
#[tokio::test]
async fn test_tagged_response_matching() {
let (client_io, mut server_io) = duplex(1024);
let mut client = RawClient::new(client_io);
let command_task = tokio::spawn(async move { client.execute_command("NOOP").await });
let mut buf = [0u8; 1024];
let n = server_io.read(&mut buf).await.unwrap();
let cmd = String::from_utf8_lossy(&buf[..n]);
assert!(cmd.contains("NOOP"));
let tag = cmd.split_whitespace().next().unwrap();
server_io
.write_all(format!("{} OK NOOP completed\r\n", tag).as_bytes())
.await
.unwrap();
let result = command_task.await.unwrap().unwrap();
assert!(String::from_utf8_lossy(&result).contains("OK"));
}
#[tokio::test]
async fn test_no_response_becomes_error() {
let (client_io, mut server_io) = duplex(1024);
let mut client = RawClient::new(client_io);
let command_task = tokio::spawn(async move { client.execute_command("LOGIN x y").await });
let mut buf = [0u8; 1024];
let n = server_io.read(&mut buf).await.unwrap();
let tag = String::from_utf8_lossy(&buf[..n])
.split_whitespace()
.next()
.unwrap()
.to_owned();
server_io
.write_all(format!("{} NO authentication failed\r\n", tag).as_bytes())
.await
.unwrap();
let result = command_task.await.unwrap();
match result {
Err(ClientError::CommandFailed(text)) => {
assert_eq!(text, "authentication failed")
}
other => panic!("expected CommandFailed, got {:?}", other),
}
}
#[tokio::test]
async fn test_bad_response_becomes_error() {
let (client_io, mut server_io) = duplex(1024);
let mut client = RawClient::new(client_io);
let command_task = tokio::spawn(async move { client.execute_command("BOGUS").await });
let mut buf = [0u8; 1024];
let n = server_io.read(&mut buf).await.unwrap();
let tag = String::from_utf8_lossy(&buf[..n])
.split_whitespace()
.next()
.unwrap()
.to_owned();
server_io
.write_all(format!("{} BAD unknown command\r\n", tag).as_bytes())
.await
.unwrap();
let result = command_task.await.unwrap();
assert!(matches!(result, Err(ClientError::CommandFailed(_))));
}
#[tokio::test]
async fn test_untagged_event_broadcasting() {
let (client_io, mut server_io) = duplex(1024);
let client = RawClient::new(client_io);
let mut events = client.events();
server_io.write_all(b"* 5 EXISTS\r\n").await.unwrap();
let event = events.recv().await.unwrap();
assert_eq!(String::from_utf8_lossy(&event), "* 5 EXISTS\r\n");
}
#[tokio::test]
async fn test_partial_read_reassembly() {
let (client_io, mut server_io) = duplex(1024);
let mut client = RawClient::new(client_io);
let command_task = tokio::spawn(async move { client.execute_command("NOOP").await });
let mut buf = [0u8; 1024];
let n = server_io.read(&mut buf).await.unwrap();
let tag = String::from_utf8_lossy(&buf[..n])
.split_whitespace()
.next()
.unwrap()
.to_string();
let response = format!("{} OK NOOP completed\r\n", tag);
for byte in response.as_bytes() {
server_io.write_all(&[*byte]).await.unwrap();
tokio::task::yield_now().await;
}
let result = command_task.await.unwrap().unwrap();
assert!(String::from_utf8_lossy(&result).contains("OK"));
}
#[tokio::test]
async fn test_command_timeout() {
let (client_io, _server_io) = duplex(1024);
let mut client = RawClient::new(client_io);
client.default_timeout = Duration::from_millis(50);
let result = client.execute_command("NOOP").await;
assert!(matches!(result, Err(ClientError::Timeout)));
}
#[tokio::test]
async fn test_search_parsing() {
let (client_io, mut server_io) = duplex(1024);
let mut client = RawClient::new(client_io);
let mut events = client.events();
let command_task =
tokio::spawn(async move { client.execute_command("SEARCH FROM \"alice\"").await });
let mut buf = [0u8; 1024];
let n = server_io.read(&mut buf).await.unwrap();
let cmd = String::from_utf8_lossy(&buf[..n]);
let tag = cmd.split_whitespace().next().unwrap();
server_io.write_all(b"* SEARCH 1 2 3\r\n").await.unwrap();
server_io
.write_all(format!("{} OK SEARCH completed\r\n", tag).as_bytes())
.await
.unwrap();
let _result = command_task.await.unwrap().unwrap();
let event = events.recv().await.unwrap();
assert_eq!(String::from_utf8_lossy(&event), "* SEARCH 1 2 3\r\n");
}
#[tokio::test]
async fn test_send_raw() {
let (client_io, mut server_io) = duplex(1024);
let mut client = RawClient::new(client_io);
client.send_raw(b"DONE\r\n".to_vec()).await.unwrap();
let mut buf = [0u8; 1024];
let n = server_io.read(&mut buf).await.unwrap();
assert_eq!(&buf[..n], b"DONE\r\n");
}
#[tokio::test]
async fn test_connection_closed_on_eof() {
let (client_io, server_io) = duplex(1024);
let mut client = RawClient::new(client_io);
client.default_timeout = Duration::from_secs(1);
let task = tokio::spawn(async move { client.execute_command("NOOP").await });
tokio::time::sleep(Duration::from_millis(20)).await;
drop(server_io);
let result = task.await.unwrap();
assert!(matches!(
result,
Err(ClientError::ConnectionClosed) | Err(ClientError::Timeout)
));
}
#[tokio::test]
async fn test_connection_closed_immediate() {
let (client_io, server_io) = duplex(1024);
let mut client = RawClient::new(client_io);
drop(server_io);
let result = client.execute_command("NOOP").await;
assert!(matches!(
result,
Err(ClientError::ConnectionClosed) | Err(ClientError::Timeout)
));
}
#[tokio::test]
async fn test_unterminated_frame_is_bounded() {
let (client_io, mut server_io) = duplex(8192);
let pending: PendingCommands = Arc::new(Mutex::new(HashMap::new()));
let (event_tx, _event_rx) = broadcast::channel(EVENT_CHANNEL_CAP);
let (reply_tx, reply_rx) = oneshot::channel();
pending.lock().await.insert("A0001".to_string(), reply_tx);
let max_frame_size = 64;
let loop_task = tokio::spawn(read_loop(
client_io,
Arc::clone(&pending),
event_tx,
max_frame_size,
));
server_io.write_all(b"* OK ").await.unwrap();
server_io
.write_all(&vec![b'a'; max_frame_size * 2])
.await
.unwrap();
let result = reply_rx.await.unwrap();
assert!(
matches!(result, Err(ClientError::FrameTooLarge { max }) if max == max_frame_size),
"expected FrameTooLarge, got {result:?}"
);
loop_task.await.unwrap();
}
}