use std::{
collections::HashMap,
io::{Read, Seek, SeekFrom},
path::{Path, PathBuf},
sync::{
Arc, Condvar, Mutex, MutexGuard, OnceLock,
atomic::{AtomicU64, Ordering},
},
time::Instant,
};
use anyhow::{Context, Result, anyhow};
use reqwest::Url;
use ssh2::{Session, Sftp};
use super::sftp_utils;
use crate::{Blob, ByteRange};
fn next_connection_id() -> u64 {
static NEXT: AtomicU64 = AtomicU64::new(0);
NEXT.fetch_add(1, Ordering::Relaxed)
}
const DEFAULT_CONNECTIONS_PER_SERVER: usize = 8;
fn connections_per_server() -> usize {
static CAP: OnceLock<usize> = OnceLock::new();
*CAP.get_or_init(|| {
std::env::var("VERSATILES_SFTP_MAX_CONNECTIONS")
.ok()
.and_then(|value| value.parse::<usize>().ok())
.filter(|&n| n > 0)
.unwrap_or(DEFAULT_CONNECTIONS_PER_SERVER)
})
}
const MAX_CONCURRENT_OPENS: usize = 8;
struct OpenThrottle {
permits: Mutex<usize>,
released: Condvar,
}
impl OpenThrottle {
fn acquire(&self) -> OpenPermit<'_> {
let mut permits = self.permits.lock().unwrap_or_else(std::sync::PoisonError::into_inner);
while *permits == 0 {
permits = self
.released
.wait(permits)
.unwrap_or_else(std::sync::PoisonError::into_inner);
}
*permits -= 1;
OpenPermit { throttle: self }
}
}
struct OpenPermit<'a> {
throttle: &'a OpenThrottle,
}
impl Drop for OpenPermit<'_> {
fn drop(&mut self) {
let mut permits = self
.throttle
.permits
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
*permits += 1;
self.throttle.released.notify_one();
}
}
fn open_throttle() -> &'static OpenThrottle {
static THROTTLE: OnceLock<OpenThrottle> = OnceLock::new();
THROTTLE.get_or_init(|| OpenThrottle {
permits: Mutex::new(MAX_CONCURRENT_OPENS),
released: Condvar::new(),
})
}
#[derive(Clone, PartialEq, Eq, Hash, Debug)]
struct ServerKey {
host: String,
port: u16,
username: String,
}
impl ServerKey {
fn from_url(url: &Url) -> Result<Self> {
Ok(ServerKey {
host: url.host_str().context("SFTP URL has no host")?.to_string(),
port: url.port().unwrap_or(22),
username: if url.username().is_empty() {
"root"
} else {
url.username()
}
.to_string(),
})
}
}
struct LiveConnection {
_session: Session,
sftp: Sftp,
files: HashMap<u64, ssh2::File>,
}
struct ConnectionInner {
live: Option<LiveConnection>,
registered: HashMap<u64, PathBuf>,
generation: u64,
next_file_id: u64,
last_used: Instant,
}
pub struct Connection {
id: u64,
inner: Mutex<ConnectionInner>,
url: Url,
identity_file: Option<PathBuf>,
}
impl Connection {
fn open(url: &Url, identity_file: Option<&Path>) -> Result<Arc<Connection>> {
let id = next_connection_id();
let display = sftp_utils::display_name(url);
log::debug!("[sftp conn {id}] opening SSH+SFTP session to '{display}'");
let started = Instant::now();
let session = sftp_utils::open_session(url, identity_file)?;
let sftp = session.sftp()?;
log::debug!(
"[sftp conn {id}] session ready in {:.2}s",
started.elapsed().as_secs_f32()
);
Ok(Arc::new(Connection {
id,
inner: Mutex::new(ConnectionInner {
live: Some(LiveConnection {
_session: session,
sftp,
files: HashMap::new(),
}),
registered: HashMap::new(),
generation: 0,
next_file_id: 0,
last_used: Instant::now(),
}),
url: url.clone(),
identity_file: identity_file.map(Path::to_path_buf),
}))
}
fn lock(&self) -> Result<MutexGuard<'_, ConnectionInner>> {
self
.inner
.lock()
.map_err(|e| anyhow!("SFTP connection lock poisoned: {e}"))
}
pub fn register(&self, path: &Path) -> Result<(u64, u64)> {
let mut inner = self.lock()?;
let id = inner.next_file_id;
let live = inner
.live
.as_mut()
.ok_or_else(|| anyhow!("SFTP connection is not established"))?;
let size = live
.sftp
.stat(path)
.with_context(|| format!("failed to stat remote file {path:?}"))?
.size
.unwrap_or(0);
let file = live
.sftp
.open(path)
.with_context(|| format!("failed to open remote file {path:?}"))?;
live.files.insert(id, file);
inner.registered.insert(id, path.to_path_buf());
inner.next_file_id += 1;
Ok((id, size))
}
pub fn unregister(&self, id: u64) {
if let Ok(mut inner) = self.lock() {
inner.registered.remove(&id);
if let Some(live) = inner.live.as_mut() {
live.files.remove(&id);
}
}
}
pub fn generation(&self) -> Result<u64> {
Ok(self.lock()?.generation)
}
pub fn read_range(&self, id: u64, range: &ByteRange) -> Result<Blob> {
let mut inner = self.lock()?;
let idle = inner.last_used.elapsed();
let read_result: Result<Blob> = (|| {
let file = inner
.live
.as_mut()
.ok_or_else(|| anyhow!("SFTP connection is not established"))?
.files
.get_mut(&id)
.ok_or_else(|| anyhow!("SFTP file id {id} is not registered"))?;
file.seek(SeekFrom::Start(range.offset))?;
let mut buffer = vec![0u8; usize::try_from(range.length)?];
file.read_exact(&mut buffer)?;
Ok(Blob::from(buffer))
})();
match read_result {
Ok(blob) => {
inner.last_used = Instant::now();
Ok(blob)
}
Err(e) => {
log::debug!(
"[sftp conn {}] read {range} failed after {:.1}s idle: {e}",
self.id,
idle.as_secs_f32(),
);
Err(e)
}
}
}
pub fn reconnect(&self, seen_generation: u64) -> Result<()> {
let mut inner = self.lock()?;
if inner.generation != seen_generation {
return Ok(());
}
log::info!(
"[sftp conn {}] reconnecting (idle for {:.1}s) to '{}'",
self.id,
inner.last_used.elapsed().as_secs_f32(),
sftp_utils::display_name(&self.url)
);
inner.live = None;
let session = sftp_utils::open_session(&self.url, self.identity_file.as_deref())?;
let sftp = session.sftp()?;
let mut files = HashMap::with_capacity(inner.registered.len());
for (id, path) in &inner.registered {
let file = sftp
.open(path)
.with_context(|| format!("failed to reopen remote file {path:?}"))?;
files.insert(*id, file);
}
inner.live = Some(LiveConnection {
_session: session,
sftp,
files,
});
inner.generation += 1;
inner.last_used = Instant::now();
log::debug!(
"[sftp conn {}] reconnect complete (generation {})",
self.id,
inner.generation
);
Ok(())
}
}
#[derive(Default)]
struct ServerPool {
connections: Vec<Arc<Connection>>,
opening: usize,
next: usize,
}
static POOL: OnceLock<Mutex<HashMap<ServerKey, ServerPool>>> = OnceLock::new();
static POOL_READY: OnceLock<Condvar> = OnceLock::new();
fn pool() -> &'static Mutex<HashMap<ServerKey, ServerPool>> {
POOL.get_or_init(|| Mutex::new(HashMap::new()))
}
fn pool_ready() -> &'static Condvar {
POOL_READY.get_or_init(Condvar::new)
}
enum Decision {
Open,
Reuse(Arc<Connection>),
Wait,
}
pub fn acquire(url: &Url, identity_file: Option<&Path>) -> Result<Arc<Connection>> {
let key = ServerKey::from_url(url)?;
let mut guard = pool().lock().map_err(|e| anyhow!("SFTP pool lock poisoned: {e}"))?;
loop {
let decision = {
let server = guard.entry(key.clone()).or_default();
if server.connections.len() + server.opening < connections_per_server() {
server.opening += 1;
Decision::Open
} else if server.connections.is_empty() {
Decision::Wait
} else {
let connection = Arc::clone(&server.connections[server.next]);
server.next = (server.next + 1) % server.connections.len();
Decision::Reuse(connection)
}
};
match decision {
Decision::Reuse(connection) => {
log::debug!("[sftp pool] {}: reusing conn {}", key.host, connection.id);
return Ok(connection);
}
Decision::Wait => {
log::debug!(
"[sftp pool] {}: at cap ({}), waiting for a connection",
key.host,
connections_per_server()
);
guard = pool_ready()
.wait(guard)
.map_err(|e| anyhow!("SFTP pool lock poisoned: {e}"))?;
}
Decision::Open => {
drop(guard);
let result = {
let _permit = open_throttle().acquire();
Connection::open(url, identity_file)
};
let outcome = {
let mut g = pool().lock().map_err(|e| anyhow!("SFTP pool lock poisoned: {e}"))?;
let server = g.entry(key.clone()).or_default();
server.opening -= 1;
match result {
Ok(connection) => {
server.connections.push(Arc::clone(&connection));
log::debug!(
"[sftp pool] {}: opened conn {}, pool now {} connection(s) (cap {})",
key.host,
connection.id,
server.connections.len(),
connections_per_server()
);
Ok(connection)
}
Err(e) => {
log::debug!("[sftp pool] {}: open failed: {e}", key.host);
Err(e)
}
}
};
pool_ready().notify_all();
return outcome;
}
}
}
}
#[cfg(test)]
#[allow(dead_code)]
fn connection_count(url: &Url) -> Result<usize> {
let key = ServerKey::from_url(url)?;
let guard = pool().lock().map_err(|e| anyhow!("SFTP pool lock poisoned: {e}"))?;
Ok(guard.get(&key).map_or(0, |server| server.connections.len()))
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_server_key_from_url_defaults() {
let key = ServerKey::from_url(&Url::parse("sftp://host/path").unwrap()).unwrap();
assert_eq!(key.host, "host");
assert_eq!(key.port, 22);
assert_eq!(key.username, "root");
}
#[test]
fn test_server_key_from_url_explicit() {
let key = ServerKey::from_url(&Url::parse("sftp://alice@host:2222/path").unwrap()).unwrap();
assert_eq!(key.port, 2222);
assert_eq!(key.username, "alice");
}
#[test]
fn test_server_key_ignores_path() {
let a = ServerKey::from_url(&Url::parse("sftp://host/one.bin").unwrap()).unwrap();
let b = ServerKey::from_url(&Url::parse("sftp://host/two.bin").unwrap()).unwrap();
assert_eq!(a, b);
}
#[test]
fn test_server_key_missing_host() {
assert!(ServerKey::from_url(&Url::parse("sftp:///path").unwrap()).is_err());
}
#[cfg(all(feature = "ssh2", unix))]
mod sftp_server_tests {
use super::*;
use crate::io::test_sftp_server::TestSftpServer;
#[tokio::test(flavor = "multi_thread")]
#[serial_test::serial]
async fn caps_connections_per_server() {
let server = TestSftpServer::start().await;
server.write_file("/a.bin", b"hello").await;
let url = server.url("/a.bin");
let cap = connections_per_server();
let total = cap + 4;
let mut handles = Vec::with_capacity(total);
for _ in 0..total {
let url = url.clone();
handles.push(tokio::task::spawn_blocking(move || acquire(&url, None)));
}
let mut acquired = Vec::with_capacity(total);
for handle in handles {
acquired.push(handle.await.unwrap().unwrap());
}
assert_eq!(acquired.len(), total);
assert_eq!(connection_count(&url).unwrap(), cap);
}
#[tokio::test(flavor = "multi_thread")]
#[serial_test::serial]
async fn many_files_share_one_connection() {
let server = TestSftpServer::start().await;
server.write_file("/a.bin", b"aaaa").await;
server.write_file("/b.bin", b"bbbbbb").await;
let url = server.url("/a.bin");
tokio::task::spawn_blocking(move || -> Result<()> {
let connection = acquire(&url, None)?;
let (id_a, size_a) = connection.register(Path::new("/a.bin"))?;
let (id_b, size_b) = connection.register(Path::new("/b.bin"))?;
assert_eq!(size_a, 4);
assert_eq!(size_b, 6);
assert_eq!(connection.read_range(id_a, &ByteRange::new(0, 4))?.as_slice(), b"aaaa");
assert_eq!(connection.read_range(id_b, &ByteRange::new(2, 4))?.as_slice(), b"bbbb");
connection.unregister(id_a);
assert!(connection.read_range(id_a, &ByteRange::new(0, 4)).is_err());
Ok(())
})
.await
.unwrap()
.unwrap();
}
#[tokio::test(flavor = "multi_thread")]
#[serial_test::serial]
async fn reconnect_reopens_registered_files() {
let server = TestSftpServer::start().await;
server.write_file("/a.bin", b"hello").await;
let url = server.url("/a.bin");
tokio::task::spawn_blocking(move || -> Result<()> {
let connection = acquire(&url, None)?;
let (id, _) = connection.register(Path::new("/a.bin"))?;
let generation = connection.generation()?;
connection.reconnect(generation)?;
assert_eq!(connection.generation()?, generation + 1);
assert_eq!(connection.read_range(id, &ByteRange::new(0, 5))?.as_slice(), b"hello");
connection.reconnect(generation)?;
assert_eq!(connection.generation()?, generation + 1);
Ok(())
})
.await
.unwrap()
.unwrap();
}
}
}