use std::{
collections::HashMap,
path::{Path, PathBuf},
sync::{
Arc, OnceLock,
atomic::{AtomicU64, Ordering},
},
time::Instant,
};
use anyhow::{Context, Result, anyhow};
use reqwest::Url;
use tokio::{
io::{AsyncReadExt, AsyncSeekExt},
sync::{Mutex, Notify, Semaphore},
};
use super::sftp_utils::{self, Sftp, SshHandle};
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;
fn open_throttle() -> &'static Semaphore {
static THROTTLE: OnceLock<Semaphore> = OnceLock::new();
THROTTLE.get_or_init(|| Semaphore::new(MAX_CONCURRENT_OPENS))
}
#[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: SshHandle,
sftp: Sftp,
files: HashMap<u64, russh_sftp::client::fs::File>,
}
struct ConnectionInner {
live: Option<LiveConnection>,
registered: HashMap<u64, String>,
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 {
async 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).await?;
let sftp = sftp_utils::open_sftp(&session).await?;
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),
}))
}
pub async fn register(&self, path: &str) -> Result<(u64, u64)> {
let mut inner = self.inner.lock().await;
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
.metadata(path.to_owned())
.await
.with_context(|| format!("failed to stat remote file {path:?}"))?
.size
.unwrap_or(0);
let file = live
.sftp
.open(path.to_owned())
.await
.with_context(|| format!("failed to open remote file {path:?}"))?;
live.files.insert(id, file);
inner.registered.insert(id, path.to_owned());
inner.next_file_id += 1;
Ok((id, size))
}
pub async fn unregister(&self, id: u64) {
let mut inner = self.inner.lock().await;
inner.registered.remove(&id);
if let Some(live) = inner.live.as_mut() {
live.files.remove(&id);
}
}
pub async fn generation(&self) -> u64 {
self.inner.lock().await.generation
}
pub async fn read_range(&self, id: u64, range: &ByteRange) -> Result<Blob> {
let mut inner = self.inner.lock().await;
let idle = inner.last_used.elapsed();
let read_result: Result<Blob> = async {
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(std::io::SeekFrom::Start(range.offset)).await?;
let mut buffer = vec![0u8; usize::try_from(range.length)?];
file.read_exact(&mut buffer).await?;
Ok(Blob::from(buffer))
}
.await;
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 async fn reconnect(&self, seen_generation: u64) -> Result<()> {
let mut inner = self.inner.lock().await;
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()).await?;
let sftp = sftp_utils::open_sftp(&session).await?;
let mut files = HashMap::with_capacity(inner.registered.len());
for (id, path) in &inner.registered {
let file = sftp
.open(path.clone())
.await
.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<Notify> = OnceLock::new();
fn pool() -> &'static Mutex<HashMap<ServerKey, ServerPool>> {
POOL.get_or_init(|| Mutex::new(HashMap::new()))
}
fn pool_ready() -> &'static Notify {
POOL_READY.get_or_init(Notify::new)
}
enum Decision {
Open,
Reuse(Arc<Connection>),
Wait,
}
pub async fn acquire(url: &Url, identity_file: Option<&Path>) -> Result<Arc<Connection>> {
let key = ServerKey::from_url(url)?;
let ready = pool_ready().notified();
tokio::pin!(ready);
loop {
ready.as_mut().enable();
let decision = {
let mut guard = pool().lock().await;
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()
);
ready.as_mut().await;
ready.set(pool_ready().notified());
}
Decision::Open => {
let result = {
let _permit = open_throttle()
.acquire()
.await
.map_err(|e| anyhow!("the SFTP handshake throttle was closed: {e}"))?;
Connection::open(url, identity_file).await
};
let outcome = {
let mut guard = pool().lock().await;
let server = guard.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_waiters();
return outcome;
}
}
}
}
#[cfg(test)]
async fn connection_count(url: &Url) -> Result<usize> {
let key = ServerKey::from_url(url)?;
let guard = pool().lock().await;
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 url = Url::parse("sftp://host/path").unwrap();
let key = ServerKey::from_url(&url).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 url = Url::parse("sftp://user@host:2222/path").unwrap();
let key = ServerKey::from_url(&url).unwrap();
assert_eq!(key.host, "host");
assert_eq!(key.port, 2222);
assert_eq!(key.username, "user");
}
#[test]
fn test_server_key_ignores_path() {
let a = ServerKey::from_url(&Url::parse("sftp://host/a").unwrap()).unwrap();
let b = ServerKey::from_url(&Url::parse("sftp://host/b").unwrap()).unwrap();
assert_eq!(a, b);
}
#[test]
fn test_server_key_missing_host() {
let url = Url::parse("sftp:///path").unwrap();
assert!(ServerKey::from_url(&url).is_err());
}
#[cfg(feature = "sftp")]
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::spawn(async move { acquire(&url, None).await }));
}
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).await.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");
let connection = acquire(&url, None).await.unwrap();
let (id_a, size_a) = connection.register("/a.bin").await.unwrap();
let (id_b, size_b) = connection.register("/b.bin").await.unwrap();
assert_eq!(size_a, 4);
assert_eq!(size_b, 6);
assert_eq!(
connection
.read_range(id_a, &ByteRange::new(0, 4))
.await
.unwrap()
.as_slice(),
b"aaaa"
);
assert_eq!(
connection
.read_range(id_b, &ByteRange::new(2, 4))
.await
.unwrap()
.as_slice(),
b"bbbb"
);
connection.unregister(id_a).await;
assert!(connection.read_range(id_a, &ByteRange::new(0, 4)).await.is_err());
}
#[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");
let connection = acquire(&url, None).await.unwrap();
let (id, _) = connection.register("/a.bin").await.unwrap();
let generation = connection.generation().await;
connection.reconnect(generation).await.unwrap();
assert_eq!(connection.generation().await, generation + 1);
assert_eq!(
connection
.read_range(id, &ByteRange::new(0, 5))
.await
.unwrap()
.as_slice(),
b"hello"
);
connection.reconnect(generation).await.unwrap();
assert_eq!(connection.generation().await, generation + 1);
}
}
}