use std::ops::Deref;
use std::sync::Arc;
use std::sync::atomic::{AtomicBool, Ordering};
use smol::io::{AsyncBufReadExt, AsyncWriteExt, BufReader};
use smol::lock::{Mutex, MutexGuard};
use smol::net::TcpStream;
use super::data_stream::DataStream;
use super::tls::SmolTlsStream;
use crate::Status;
use crate::command::Command;
use crate::types::{FtpError, FtpResult, Response};
pub(super) const TRANSFER_COMPLETE: &[Status] =
&[Status::ClosingDataConnection, Status::RequestedFileActionOk];
pub(super) type SharedControl<T> = Arc<Mutex<ControlChannel<T>>>;
#[derive(Debug)]
pub(super) struct ControlChannel<T>
where
T: SmolTlsStream + Send,
{
pub(super) reader: BufReader<DataStream<T>>,
pub(super) data_connection_open: bool,
pub(super) pending_transfer_reply: Arc<AtomicBool>,
response_line: Vec<u8>,
response_body: Vec<u8>,
}
#[cfg(feature = "async-secure")]
pub(super) fn into_exclusive<T>(control: SharedControl<T>) -> FtpResult<ControlChannel<T>>
where
T: SmolTlsStream + Send,
{
Arc::try_unwrap(control)
.map(Mutex::into_inner)
.map_err(|_| FtpError::DataConnectionAlreadyOpen)
}
impl<T> ControlChannel<T>
where
T: SmolTlsStream + Send,
{
pub(super) fn new(stream: DataStream<T>) -> Self {
Self {
reader: BufReader::new(stream),
data_connection_open: false,
pending_transfer_reply: Arc::new(AtomicBool::new(false)),
response_line: Vec::new(),
response_body: Vec::new(),
}
}
pub(super) fn shared(stream: DataStream<T>) -> SharedControl<T> {
Arc::new(Mutex::new(Self::new(stream)))
}
pub(super) fn socket(&self) -> &TcpStream {
self.reader.get_ref().get_ref()
}
pub(super) async fn perform(&mut self, command: Command) -> FtpResult<()> {
let command = command.to_string();
crate::command::validate_command_line(&command)?;
trace!("CC OUT: {}", command.trim_end_matches("\r\n"));
self.reader
.get_mut()
.write_all(command.as_bytes())
.await
.map_err(FtpError::ConnectionError)
}
pub(super) async fn read_response(&mut self, expected_code: Status) -> FtpResult<Response> {
self.read_response_in(&[expected_code]).await
}
pub(super) async fn read_response_in(
&mut self,
expected_code: &[Status],
) -> FtpResult<Response> {
loop {
let bytes_read = self
.reader
.read_until(b'\n', &mut self.response_line)
.await
.map_err(FtpError::ConnectionError)?;
if bytes_read == 0 && !self.response_body.is_empty() {
return Err(FtpError::ConnectionError(std::io::Error::new(
std::io::ErrorKind::UnexpectedEof,
"connection closed during multiline response",
)));
}
self.response_body.extend_from_slice(&self.response_line);
trace!("CC IN: {line:?}", line = self.response_line);
if self.response_body.len() < 5 {
self.response_line.clear();
self.response_body.clear();
return Err(FtpError::BadResponse);
}
let opening_code = match code_from_buffer(&self.response_body, 3) {
Ok(code) => code,
Err(err) => {
self.response_line.clear();
self.response_body.clear();
return Err(err);
}
};
let opening = &self.response_body[..3];
let line = &self.response_line;
let terminal = line.len() >= 4
&& line[..3].iter().all(u8::is_ascii_digit)
&& (line[3] == b' '
|| (expected_code.contains(&Status::System)
&& line[..3] == *opening
&& line[3] == b'-'));
if terminal {
let code = if line[..3] == *opening {
opening_code
} else {
code_from_buffer(line, 3)?
};
let status = Status::from(code);
self.response_line.clear();
let response = Response::new(status, std::mem::take(&mut self.response_body));
return if expected_code.contains(&status) {
Ok(response)
} else {
Err(FtpError::UnexpectedResponse(response))
};
}
self.response_line.clear();
}
}
pub(super) async fn read_line(&mut self, line: &mut Vec<u8>) -> FtpResult<usize> {
self.reader
.read_until(0x0A, line.as_mut())
.await
.map_err(FtpError::ConnectionError)?;
Ok(line.len())
}
pub(super) fn guard_multiple_data_connections(&self) -> FtpResult<()> {
if self.data_connection_open {
Err(FtpError::DataConnectionAlreadyOpen)
} else {
Ok(())
}
}
pub(super) async fn complete_transfer(&mut self) -> FtpResult<()> {
self.data_connection_open = false;
trace!("data connection closed; reading transfer reply");
let reply = self.read_response_in(TRANSFER_COMPLETE).await.map(|_| ());
self.pending_transfer_reply.store(false, Ordering::Release);
reply
}
pub(super) async fn drain_pending_transfer_reply(&mut self) -> FtpResult<()> {
if !self.pending_transfer_reply.load(Ordering::Acquire) {
return Ok(());
}
debug!("reading the reply of a transfer stream dropped without finish()");
match self.complete_transfer().await {
Ok(()) => Ok(()),
Err(FtpError::UnexpectedResponse(response)) => {
warn!("a dropped transfer stream failed: {response}");
Ok(())
}
Err(err) => Err(err),
}
}
}
fn code_from_buffer(buf: &[u8], len: usize) -> FtpResult<u32> {
if buf.len() < len {
return Err(FtpError::BadResponse);
}
let buffer = buf[0..len].to_vec();
let as_string = String::from_utf8(buffer).map_err(|_| FtpError::BadResponse)?;
as_string.parse::<u32>().map_err(|_| FtpError::BadResponse)
}
#[derive(Debug)]
pub struct ControlSocket<'a, T>
where
T: SmolTlsStream + Send,
{
guard: MutexGuard<'a, ControlChannel<T>>,
}
impl<'a, T> ControlSocket<'a, T>
where
T: SmolTlsStream + Send,
{
pub(super) fn new(guard: MutexGuard<'a, ControlChannel<T>>) -> Self {
Self { guard }
}
}
impl<T> Deref for ControlSocket<'_, T>
where
T: SmolTlsStream + Send,
{
type Target = TcpStream;
fn deref(&self) -> &Self::Target {
self.guard.socket()
}
}
#[cfg(all(test, feature = "async-secure"))]
mod tls_transition_tests {
use std::io::{BufRead, BufReader, Read, Write};
use std::net::TcpListener;
use std::sync::Arc;
use std::thread;
use std::time::Duration;
use super::super::tls::{AsyncNoTlsStream, AsyncTlsConnector};
use crate::smol::AsyncFtpStream;
use crate::{FtpError, FtpResult};
#[derive(Debug)]
struct UnusedConnector;
#[async_trait::async_trait]
impl AsyncTlsConnector for UnusedConnector {
type Stream = AsyncNoTlsStream;
async fn connect(&self, _: &str, _: smol::net::TcpStream) -> FtpResult<Self::Stream> {
panic!("TLS must not start while a transfer owns the control connection")
}
}
#[test]
fn should_reject_tls_changes_before_sending_commands() {
smol::block_on(async {
for secure in [false, true] {
let listener = TcpListener::bind("127.0.0.1:0").unwrap();
let address = listener.local_addr().unwrap();
let server = thread::spawn(move || {
let (socket, _) = listener.accept().unwrap();
socket
.set_read_timeout(Some(Duration::from_secs(5)))
.unwrap();
let mut reader = BufReader::new(socket);
reader.get_mut().write_all(b"220 ready\r\n").unwrap();
let mut command = String::new();
reader.read_line(&mut command).unwrap();
if !command.is_empty() {
let reply: &[u8] = if secure {
b"234 start TLS\r\n"
} else {
b"200 cleared\r\n"
};
reader.get_mut().write_all(reply).unwrap();
reader.read_to_end(&mut Vec::new()).unwrap();
}
command
});
let ftp = AsyncFtpStream::connect(address).await.unwrap();
let transfer_control = Arc::clone(&ftp.control);
let result = if secure {
ftp.into_secure(UnusedConnector, "localhost").await
} else {
ftp.clear_command_channel().await
};
assert!(matches!(result, Err(FtpError::DataConnectionAlreadyOpen)));
drop(transfer_control);
assert_eq!(server.join().unwrap(), "");
}
});
}
}