use std::collections::HashMap;
use std::sync::Arc;
use std::sync::atomic::{AtomicU32, Ordering};
use std::time::Duration;
use tokio::sync::{Mutex, mpsc};
use tracing::{debug, info, warn};
use super::connection::{SshConnection, SshOptions};
use super::packet::SftpPacket;
pub const SFTP_VERSION_MIN: u32 = 3;
pub const SFTP_VERSION_MAX: u32 = 6;
pub const DEFAULT_OPERATION_TIMEOUT: Duration = Duration::from_secs(60);
pub const MAX_PENDING_REQUESTS: usize = 256;
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct SftpExtension {
pub name: String,
pub data: String,
}
impl std::fmt::Display for SftpExtension {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
write!(f, "{}={}", self.name, self.data)
}
}
pub struct SftpSession {
channel: Arc<Mutex<russh::Channel<russh::client::Msg>>>,
response_rx: Arc<Mutex<mpsc::Receiver<Vec<u8>>>>,
server_version: u32,
options: Arc<SshOptions>,
next_request_id: AtomicU32,
extensions: Vec<SftpExtension>,
created_at: std::time::Instant,
operation_count: AtomicU32,
read_timeout: Duration,
recv_buffer: Arc<Mutex<Vec<u8>>>,
}
impl std::fmt::Debug for SftpSession {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("SftpSession")
.field("server_version", &self.server_version)
.field("options", &self.options)
.field("extensions", &self.extensions)
.field("age_secs", &self.age().as_secs())
.finish()
}
}
impl Clone for SftpSession {
fn clone(&self) -> Self {
Self {
channel: Arc::clone(&self.channel),
response_rx: Arc::clone(&self.response_rx),
server_version: self.server_version,
options: Arc::clone(&self.options),
next_request_id: AtomicU32::new(self.next_request_id.load(Ordering::Relaxed)),
extensions: self.extensions.clone(),
created_at: self.created_at,
operation_count: AtomicU32::new(self.operation_count.load(Ordering::Relaxed)),
read_timeout: self.read_timeout,
recv_buffer: Arc::clone(&self.recv_buffer),
}
}
}
impl SftpSession {
pub async fn open(conn: &mut SshConnection) -> Result<Self, String> {
debug!("[SFTP] Initializing SFTP session (pure Rust)...");
let (channel, data_rx) = conn
.open_sftp_channel()
.await
.map_err(|e| format!("Failed to open SFTP channel: {}", e))?;
let options = Arc::clone(conn.options());
let read_timeout = options.read_timeout;
let channel = Arc::new(Mutex::new(channel));
let response_rx = Arc::new(Mutex::new(data_rx));
let recv_buffer = Arc::new(Mutex::new(Vec::new()));
let init_pkt = SftpPacket::Init { version: 3 };
let encoded = init_pkt
.encode()
.map_err(|e| format!("Failed to encode SFTP INIT packet: {}", e))?;
{
let ch = channel.lock().await;
ch.data(encoded.as_slice())
.await
.map_err(|e| format!("Failed to send SFTP INIT packet: {}", e))?;
}
debug!("[SFTP] Sent SFTP INIT (version=3), awaiting VERSION...");
let version_response =
Self::recv_packet_from_channel(&response_rx, &recv_buffer, read_timeout).await?;
let (server_version, extensions) = match version_response {
SftpPacket::Version {
version,
extensions,
} => {
if version < SFTP_VERSION_MIN {
return Err(format!(
"Server SFTP version too low: v{} (minimum v{})",
version, SFTP_VERSION_MIN
));
}
(version, extensions)
}
other => {
return Err(format!(
"Expected VERSION packet after INIT, got type={}",
other.packet_type()
));
}
};
info!(
"[SFTP] Session established (server version=v{}, supports v{}-v{}, {} extensions)",
server_version,
SFTP_VERSION_MIN,
SFTP_VERSION_MAX,
extensions.len()
);
Ok(Self {
channel,
response_rx,
server_version,
options,
next_request_id: AtomicU32::new(1), extensions: extensions
.into_iter()
.map(|(name, data)| SftpExtension { name, data })
.collect(),
created_at: std::time::Instant::now(),
operation_count: AtomicU32::new(0),
read_timeout,
recv_buffer,
})
}
pub async fn send_packet(&self, pkt: &SftpPacket) -> Result<(), String> {
let encoded = pkt.encode().map_err(|e| format!("Encode error: {}", e))?;
let ch = self.channel.lock().await;
ch.data(encoded.as_slice())
.await
.map_err(|e| format!("Channel write error: {}", e))
}
pub async fn recv_packet(&self) -> Result<SftpPacket, String> {
Self::recv_packet_from_channel(&self.response_rx, &self.recv_buffer, self.read_timeout)
.await
}
async fn recv_packet_from_channel(
response_rx: &Arc<Mutex<mpsc::Receiver<Vec<u8>>>>,
recv_buffer: &Arc<Mutex<Vec<u8>>>,
timeout: Duration,
) -> Result<SftpPacket, String> {
loop {
{
let mut buf = recv_buffer.lock().await;
if !buf.is_empty() {
match SftpPacket::decode(buf.as_slice()) {
Ok((pkt, consumed)) => {
let remaining = buf.split_off(consumed);
*buf = remaining;
return Ok(pkt);
}
Err(_) => {
}
}
}
}
let mut rx = response_rx.lock().await;
match tokio::time::timeout(timeout, rx.recv()).await {
Ok(Some(chunk)) => {
drop(rx); let mut buf = recv_buffer.lock().await;
buf.extend_from_slice(&chunk);
}
Ok(None) => {
return Err("Channel closed by server".to_string());
}
Err(_) => {
return Err(format!("Receive timed out after {}s", timeout.as_secs()));
}
}
}
}
pub async fn request(&self, mut pkt: SftpPacket) -> Result<SftpPacket, String> {
let req_id = self.allocate_request_id();
Self::set_request_id_in_packet(&mut pkt, req_id);
self.send_packet(&pkt).await?;
let resp = self.recv_packet().await?;
Ok(resp)
}
fn set_request_id_in_packet(pkt: &mut SftpPacket, new_request_id: u32) {
match pkt {
SftpPacket::Open { request_id, .. } => *request_id = new_request_id,
SftpPacket::Close { request_id, .. } => *request_id = new_request_id,
SftpPacket::Read { request_id, .. } => *request_id = new_request_id,
SftpPacket::Write { request_id, .. } => *request_id = new_request_id,
SftpPacket::Stat { request_id, .. } => *request_id = new_request_id,
SftpPacket::Lstat { request_id, .. } => *request_id = new_request_id,
SftpPacket::Fstat { request_id, .. } => *request_id = new_request_id,
SftpPacket::Setstat { request_id, .. } => *request_id = new_request_id,
SftpPacket::Fsetstat { request_id, .. } => *request_id = new_request_id,
SftpPacket::Opendir { request_id, .. } => *request_id = new_request_id,
SftpPacket::Readdir { request_id, .. } => *request_id = new_request_id,
SftpPacket::Realpath { request_id, .. } => *request_id = new_request_id,
SftpPacket::Remove { request_id, .. } => *request_id = new_request_id,
SftpPacket::Mkdir { request_id, .. } => *request_id = new_request_id,
SftpPacket::Rmdir { request_id, .. } => *request_id = new_request_id,
SftpPacket::Rename { request_id, .. } => *request_id = new_request_id,
SftpPacket::Readlink { request_id, .. } => *request_id = new_request_id,
SftpPacket::Symlink { request_id, .. } => *request_id = new_request_id,
_ => {} }
}
pub fn server_version(&self) -> u32 {
self.server_version
}
pub fn supports_version(&self, v: u32) -> bool {
self.server_version >= v
}
pub fn options(&self) -> &Arc<SshOptions> {
&self.options
}
pub fn age(&self) -> Duration {
self.created_at.elapsed()
}
pub fn operation_count(&self) -> u32 {
self.operation_count.load(Ordering::Relaxed)
}
pub fn read_timeout(&self) -> Duration {
self.read_timeout
}
pub fn allocate_request_id(&self) -> u32 {
let id = self.next_request_id.fetch_add(1, Ordering::Relaxed);
if id == 0 {
warn!("[SFTP] Request ID counter wrapped around");
self.next_request_id.store(1, Ordering::Relaxed);
1
} else {
id
}
}
#[cfg(test)]
pub fn reset_request_id_counter(&self) {
self.next_request_id.store(1, Ordering::Relaxed);
}
pub fn record_operation(&self) {
self.operation_count.fetch_add(1, Ordering::Relaxed);
}
pub fn has_extension(&self, name: &str) -> bool {
self.extensions.iter().any(|ext| ext.name == name)
}
pub fn get_extension(&self, name: &str) -> Option<&SftpExtension> {
self.extensions.iter().find(|ext| ext.name == name)
}
pub fn extensions(&self) -> Vec<SftpExtension> {
self.extensions.clone()
}
pub fn add_extension(&mut self, name: impl Into<String>, data: impl Into<String>) {
self.extensions.push(SftpExtension {
name: name.into(),
data: data.into(),
});
}
pub fn parse_extensions_from_version(&mut self, extensions: Vec<(String, String)>) {
self.extensions = extensions
.into_iter()
.map(|(name, data)| SftpExtension { name, data })
.collect();
if !self.extensions.is_empty() {
debug!(
"[SFTP] Server advertises {} extensions: {}",
self.extensions.len(),
self.extensions
.iter()
.map(|e| e.name.as_str())
.collect::<Vec<_>>()
.join(", ")
);
}
}
pub fn diagnostics(&self) -> String {
format!(
"SftpSession{{version=v{}, age={}s, ops={}, extensions=[{}]}}",
self.server_version,
self.age().as_secs(),
self.operation_count(),
self.extensions
.iter()
.map(|e| e.name.as_str())
.collect::<Vec<_>>()
.join(", ")
)
}
}
pub struct PendingRequestTracker {
pending: HashMap<u32, tokio::sync::oneshot::Sender<SftpPacket>>,
max_pending: usize,
}
impl PendingRequestTracker {
pub fn new(max_pending: usize) -> Self {
Self {
pending: HashMap::with_capacity(max_pending),
max_pending,
}
}
pub fn register(
&mut self,
request_id: u32,
) -> Result<tokio::sync::oneshot::Receiver<SftpPacket>, PendingRequestError> {
if self.pending.len() >= self.max_pending {
return Err(PendingRequestError::AtCapacity {
current: self.pending.len(),
max: self.max_pending,
});
}
if self.pending.contains_key(&request_id) {
return Err(PendingRequestError::DuplicateId(request_id));
}
let (tx, rx) = tokio::sync::oneshot::channel();
self.pending.insert(request_id, tx);
Ok(rx)
}
pub fn deliver_response(&mut self, response: SftpPacket) -> Result<bool, PendingRequestError> {
let request_id = match response.request_id() {
Some(id) => id,
None => return Err(PendingRequestError::NoRequestId),
};
match self.pending.remove(&request_id) {
Some(sender) => {
let _ = sender.send(response);
Ok(true)
}
None => Ok(false),
}
}
pub fn cancel(&mut self, request_id: u32) -> bool {
self.pending.remove(&request_id).is_some()
}
pub fn cancel_all(&mut self) -> usize {
let count = self.pending.len();
self.pending.clear();
count
}
pub fn len(&self) -> usize {
self.pending.len()
}
pub fn is_empty(&self) -> bool {
self.pending.is_empty()
}
}
impl Default for PendingRequestTracker {
fn default() -> Self {
Self::new(MAX_PENDING_REQUESTS)
}
}
#[derive(Debug, Clone, thiserror::Error)]
pub enum PendingRequestError {
#[error("Tracker at capacity ({current}/{max})")]
AtCapacity { current: usize, max: usize },
#[error("Duplicate request ID: {0}")]
DuplicateId(u32),
#[error("Response packet has no request ID")]
NoRequestId,
#[error("Request timed out after {timeout_secs}s")]
TimedOut { request_id: u32, timeout_secs: u64 },
}
#[cfg(test)]
mod tests {
use super::*;
use crate::sftp::packet::status_code_description;
use crate::sftp::packet::{SSH_FX_OK, SftpFileAttrs};
#[test]
fn test_version_constants() {
const _: () = assert!(SFTP_VERSION_MIN <= SFTP_VERSION_MAX);
assert_eq!(SFTP_VERSION_MIN, 3);
assert_eq!(SFTP_VERSION_MAX, 6);
}
#[test]
fn test_default_operation_timeout() {
assert_eq!(DEFAULT_OPERATION_TIMEOUT, Duration::from_secs(60));
}
#[test]
fn test_sftp_extension_creation_and_display() {
let ext = SftpExtension {
name: "hardlink@openssh.com".to_string(),
data: "1".to_string(),
};
assert_eq!(ext.name, "hardlink@openssh.com");
assert_eq!(ext.data, "1");
assert_eq!(format!("{}", ext), "hardlink@openssh.com=1");
}
#[test]
fn test_sftp_extension_equality() {
let ext1 = SftpExtension {
name: "fsync@openssh.com".to_string(),
data: "2".to_string(),
};
let ext2 = SftpExtension {
name: "fsync@openssh.com".to_string(),
data: "2".to_string(),
};
let ext3 = SftpExtension {
name: "fsync@openssh.com".to_string(),
data: "3".to_string(),
};
assert_eq!(ext1, ext2);
assert_ne!(ext1, ext3);
}
#[test]
fn test_session_extensions_empty_by_default() {
let exts: Vec<SftpExtension> = Vec::new();
assert!(exts.is_empty());
assert_eq!(exts.len(), 0);
}
#[test]
fn test_request_id_allocation_monotonic() {
let counter = AtomicU32::new(1);
let id1 = counter.fetch_add(1, Ordering::Relaxed);
let id2 = counter.fetch_add(1, Ordering::Relaxed);
let id3 = counter.fetch_add(1, Ordering::Relaxed);
assert_eq!(id1, 1);
assert_eq!(id2, 2);
assert_eq!(id3, 3);
assert!(id2 > id1);
assert!(id3 > id2);
}
#[test]
fn test_request_id_uniqueness_across_allocations() {
let counter = AtomicU32::new(1);
let mut ids = std::collections::HashSet::new();
for _ in 0..1000 {
let id = counter.fetch_add(1, Ordering::Relaxed);
assert!(ids.insert(id), "Request ID {} was duplicated!", id);
}
assert_eq!(ids.len(), 1000);
}
#[test]
fn test_set_request_id_for_all_packet_types() {
let cases: Vec<SftpPacket> = vec![
SftpPacket::Open {
request_id: 0,
filename: "/tmp/test".into(),
flags: 0,
attrs: SftpFileAttrs::default(),
},
SftpPacket::Close {
request_id: 0,
handle: vec![1, 2, 3],
},
SftpPacket::Read {
request_id: 0,
handle: vec![1],
offset: 0,
length: 1024,
},
SftpPacket::Write {
request_id: 0,
handle: vec![1],
offset: 0,
data: vec![0],
},
SftpPacket::Stat {
request_id: 0,
path: "/test".into(),
},
SftpPacket::Lstat {
request_id: 0,
path: "/test".into(),
},
SftpPacket::Fstat {
request_id: 0,
handle: vec![1],
},
SftpPacket::Opendir {
request_id: 0,
path: "/dir".into(),
},
SftpPacket::Readdir {
request_id: 0,
handle: vec![1],
},
SftpPacket::Realpath {
request_id: 0,
path: ".".into(),
},
SftpPacket::Remove {
request_id: 0,
filename: "/file".into(),
},
SftpPacket::Mkdir {
request_id: 0,
path: "/new".into(),
attrs: SftpFileAttrs::default(),
},
SftpPacket::Rmdir {
request_id: 0,
path: "/old".into(),
},
SftpPacket::Rename {
request_id: 0,
old_path: "/a".into(),
new_path: "/b".into(),
},
SftpPacket::Readlink {
request_id: 0,
path: "/link".into(),
},
SftpPacket::Symlink {
request_id: 0,
link_path: "/l".into(),
target_path: "/t".into(),
},
];
for mut pkt in cases {
SftpSession::set_request_id_in_packet(&mut pkt, 42);
assert_eq!(
pkt.request_id(),
Some(42),
"Failed for {:?}",
pkt.packet_type()
);
}
}
#[tokio::test]
async fn test_tracker_register_and_deliver() {
let mut tracker = PendingRequestTracker::new(16);
let rx = tracker.register(42).unwrap();
assert_eq!(tracker.len(), 1);
let response = SftpPacket::Handle {
request_id: 42,
handle: vec![0x01, 0x02, 0x03],
};
let delivered = tracker.deliver_response(response).unwrap();
assert!(delivered);
assert!(tracker.is_empty());
let received = rx.await.unwrap();
assert_eq!(received.request_id(), Some(42));
}
#[tokio::test]
async fn test_tracker_cancel_removes_pending() {
let mut tracker = PendingRequestTracker::new(16);
tracker.register(1).unwrap();
tracker.register(2).unwrap();
tracker.register(3).unwrap();
assert_eq!(tracker.len(), 3);
assert!(tracker.cancel(2));
assert_eq!(tracker.len(), 2);
assert!(!tracker.cancel(99)); }
#[test]
fn test_tracker_at_capacity() {
let mut tracker = PendingRequestTracker::new(2);
tracker.register(1).unwrap();
tracker.register(2).unwrap();
let result = tracker.register(3);
assert!(result.is_err());
match result.err().unwrap() {
PendingRequestError::AtCapacity { current, max } => {
assert_eq!(current, 2);
assert_eq!(max, 2);
}
other => panic!("Expected AtCapacity error, got: {:?}", other),
}
}
#[test]
fn test_tracker_duplicate_id_rejected() {
let mut tracker = PendingRequestTracker::new(16);
tracker.register(5).unwrap();
let result = tracker.register(5);
assert!(result.is_err());
match result.err().unwrap() {
PendingRequestError::DuplicateId(id) => assert_eq!(id, 5),
other => panic!("Expected DuplicateId error, got: {:?}", other),
}
}
#[test]
fn test_tracker_cancel_all() {
let mut tracker = PendingRequestTracker::new(16);
for i in 1..=10u32 {
tracker.register(i).unwrap();
}
assert_eq!(tracker.len(), 10);
let cancelled = tracker.cancel_all();
assert_eq!(cancelled, 10);
assert!(tracker.is_empty());
}
#[tokio::test]
async fn test_tracker_dropped_receiver_cleanup() {
let mut tracker = PendingRequestTracker::new(16);
{
let _rx = tracker.register(77).unwrap();
}
let response = SftpPacket::Status {
request_id: 77,
code: SSH_FX_OK,
message: "OK".to_string(),
language: "".to_string(),
};
let delivered = tracker.deliver_response(response).unwrap();
assert!(delivered); assert!(tracker.is_empty()); }
#[test]
fn test_status_code_descriptions_accessible() {
assert_eq!(status_code_description(SSH_FX_OK), "Operation succeeded");
assert_eq!(status_code_description(1), "End of file"); }
#[test]
fn test_file_attrs_in_session_context() {
let attrs = SftpFileAttrs::full(4096, 1000, 1000, 0o040755, 1700000000, 1700000100);
assert!(attrs.is_directory());
assert_eq!(attrs.size, Some(4096));
let flags = attrs.flags();
assert!(flags != 0);
}
#[test]
fn test_max_pending_requests_constant() {
assert_eq!(MAX_PENDING_REQUESTS, 256);
}
#[test]
fn test_pending_request_error_display() {
let err = PendingRequestError::TimedOut {
request_id: 42,
timeout_secs: 30,
};
let msg = format!("{}", err);
assert!(msg.contains("timed out"));
assert!(msg.contains("30")); }
}