use std::pin::Pin;
use std::sync::Arc;
use std::sync::atomic::{AtomicBool, Ordering};
use std::task::{Context, Poll};
use tokio::io::{AsyncRead, AsyncWrite, AsyncWriteExt, ReadBuf};
use super::control::SharedControl;
use super::data_stream::DataStream;
use super::tls::TokioTlsStream;
use crate::types::{FtpError, FtpResult};
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub(super) enum Direction {
Upload,
Download,
}
#[must_use = "call `finish().await` to close the data connection and read the transfer reply"]
#[derive(Debug)]
pub struct TransferStream<T>
where
T: TokioTlsStream + Send,
{
data: Option<DataStream<T>>,
control: SharedControl<T>,
direction: Direction,
pending_reply: Arc<AtomicBool>,
finished: bool,
}
impl<T> TransferStream<T>
where
T: TokioTlsStream + Send,
{
pub(super) fn new(
data: DataStream<T>,
control: SharedControl<T>,
direction: Direction,
pending_reply: Arc<AtomicBool>,
) -> Self {
Self {
data: Some(data),
control,
direction,
pending_reply,
finished: false,
}
}
pub fn get_ref(&self) -> &DataStream<T> {
self.data
.as_ref()
.expect("data stream is present until the transfer is finished")
}
pub fn get_mut(&mut self) -> &mut DataStream<T> {
self.data
.as_mut()
.expect("data stream is present until the transfer is finished")
}
pub async fn finish(mut self) -> FtpResult<()> {
let Some(mut data) = self.data.take() else {
return Ok(());
};
let shutdown = data.shutdown().await;
drop(data);
let reply = self.control.lock().await.complete_transfer().await;
self.finished = true;
if self.direction == Direction::Upload {
shutdown.map_err(FtpError::ConnectionError)?;
}
reply
}
pub(super) fn detach(mut self) -> DataStream<T> {
self.finished = true;
self.data
.take()
.expect("data stream is present until the transfer is finished")
}
fn is_finished(&self) -> bool {
self.finished
}
}
impl<T> AsyncRead for TransferStream<T>
where
T: TokioTlsStream + Send,
{
fn poll_read(
self: Pin<&mut Self>,
cx: &mut Context<'_>,
buf: &mut ReadBuf<'_>,
) -> Poll<std::io::Result<()>> {
Pin::new(Pin::into_inner(self).get_mut()).poll_read(cx, buf)
}
}
impl<T> AsyncWrite for TransferStream<T>
where
T: TokioTlsStream + Send,
{
fn poll_write(
self: Pin<&mut Self>,
cx: &mut Context<'_>,
buf: &[u8],
) -> Poll<std::io::Result<usize>> {
Pin::new(Pin::into_inner(self).get_mut()).poll_write(cx, buf)
}
fn poll_flush(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<std::io::Result<()>> {
Pin::new(Pin::into_inner(self).get_mut()).poll_flush(cx)
}
fn poll_shutdown(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<std::io::Result<()>> {
Pin::new(Pin::into_inner(self).get_mut()).poll_shutdown(cx)
}
}
impl<T> Drop for TransferStream<T>
where
T: TokioTlsStream + Send,
{
fn drop(&mut self) {
if self.is_finished() {
return;
}
drop(self.data.take());
self.pending_reply.store(true, Ordering::Release);
warn!("transfer stream dropped without finish(); its reply is read by the next command");
}
}
#[cfg(test)]
mod tests {
use std::io::{BufRead, BufReader, Read, Write};
use std::net::TcpListener;
use std::sync::mpsc::{self, Sender};
use std::thread::{self, JoinHandle};
use std::time::Duration;
use crate::tokio::AsyncFtpStream;
async fn delayed_transfer_reply(
prefix: &'static [u8],
suffix: &'static [u8],
) -> (AsyncFtpStream, Sender<()>, JoinHandle<()>) {
let listener = TcpListener::bind("127.0.0.1:0").unwrap();
let address = listener.local_addr().unwrap();
let (release, wait) = mpsc::channel();
let server = thread::spawn(move || {
let (socket, _) = listener.accept().unwrap();
socket
.set_read_timeout(Some(Duration::from_secs(5)))
.unwrap();
let mut control = BufReader::new(socket);
control.get_mut().write_all(b"220 ready\r\n").unwrap();
let mut command = String::new();
control.read_line(&mut command).unwrap();
assert_eq!(command, "PASV\r\n");
let data_listener = TcpListener::bind("127.0.0.1:0").unwrap();
let port = data_listener.local_addr().unwrap().port();
write!(
control.get_mut(),
"227 passive (127,0,0,1,{high},{low})\r\n",
high = port / 256,
low = port % 256
)
.unwrap();
command.clear();
control.read_line(&mut command).unwrap();
assert_eq!(command, "STOR test.bin\r\n");
let (mut data, _) = data_listener.accept().unwrap();
data.set_read_timeout(Some(Duration::from_secs(5))).unwrap();
control.get_mut().write_all(b"150 send data\r\n").unwrap();
data.read_to_end(&mut Vec::new()).unwrap();
control.get_mut().write_all(prefix).unwrap();
wait.recv_timeout(Duration::from_secs(5)).unwrap();
control.get_mut().write_all(suffix).unwrap();
command.clear();
control.read_line(&mut command).unwrap();
assert_eq!(command, "NOOP\r\n");
control.get_mut().write_all(b"200 noop\r\n").unwrap();
});
let ftp = AsyncFtpStream::connect(address).await.unwrap();
(ftp, release, server)
}
#[tokio::test]
async fn should_recover_drop_while_control_socket_is_locked() {
let (mut ftp, release, server) = delayed_transfer_reply(b"", b"226 complete\r\n").await;
let transfer = ftp.put_with_stream("test.bin").await.unwrap();
let guard = ftp.get_ref().await;
drop(transfer);
drop(guard);
release.send(()).unwrap();
ftp.noop().await.unwrap();
ftp.control()
.await
.unwrap()
.guard_multiple_data_connections()
.unwrap();
server.join().unwrap();
}
#[tokio::test]
async fn should_recover_cancelled_finish_waiting_for_control_lock() {
use std::future::{Future, poll_fn};
use std::task::Poll;
let (mut ftp, release, server) = delayed_transfer_reply(b"", b"226 complete\r\n").await;
let transfer = ftp.put_with_stream("test.bin").await.unwrap();
let guard = ftp.get_ref().await;
let mut finish = Box::pin(transfer.finish());
poll_fn(|cx| {
assert!(finish.as_mut().poll(cx).is_pending());
Poll::Ready(())
})
.await;
drop(finish);
drop(guard);
release.send(()).unwrap();
ftp.noop().await.unwrap();
ftp.control()
.await
.unwrap()
.guard_multiple_data_connections()
.unwrap();
server.join().unwrap();
}
#[tokio::test]
async fn should_resume_partial_reply_after_cancelled_finish() {
let (mut ftp, release, server) = delayed_transfer_reply(
b"226-transferred\r\nintermediate line\r\n226 comp",
b"lete\r\n",
)
.await;
let transfer = ftp.put_with_stream("test.bin").await.unwrap();
assert!(
tokio::time::timeout(Duration::from_millis(50), transfer.finish())
.await
.is_err()
);
release.send(()).unwrap();
ftp.noop().await.unwrap();
ftp.control()
.await
.unwrap()
.guard_multiple_data_connections()
.unwrap();
server.join().unwrap();
}
#[tokio::test]
async fn should_resume_partial_reply_after_cancelled_drain() {
let (mut ftp, release, server) = delayed_transfer_reply(
b"226-transferred\r\nintermediate line\r\n226 comp",
b"lete\r\n",
)
.await;
let transfer = ftp.put_with_stream("test.bin").await.unwrap();
drop(transfer);
assert!(
tokio::time::timeout(Duration::from_millis(50), ftp.noop())
.await
.is_err()
);
release.send(()).unwrap();
ftp.noop().await.unwrap();
ftp.control()
.await
.unwrap()
.guard_multiple_data_connections()
.unwrap();
server.join().unwrap();
}
}