use super::sftp_utils::{self, SftpKeepalive, SharedSession};
use super::{DataWriterTrait, network_writer::NetworkWriter};
use crate::{Blob, ByteRange};
use anyhow::{Context, Result};
use reqwest::Url;
use ssh2::{OpenFlags, OpenType};
use std::{
io::{Seek, SeekFrom, Write},
path::{Path, PathBuf},
sync::{Arc, Mutex},
time::{Duration, Instant},
};
const BUFFER_CAPACITY: usize = 16 * 1024 * 1024;
fn format_bytes(n: u64) -> String {
const UNITS: [&str; 5] = ["B", "KiB", "MiB", "GiB", "TiB"];
let mut value = n as f64;
let mut unit = 0;
while value >= 1024.0 && unit < UNITS.len() - 1 {
value /= 1024.0;
unit += 1;
}
format!("{value:.1} {}", UNITS[unit])
}
pub struct DataWriterSftp {
file: ssh2::File,
position: u64,
buffer: Vec<u8>,
url: Url,
identity_file: Option<PathBuf>,
name: String,
session: SharedSession,
_keepalive: SftpKeepalive,
connected_at: Instant,
bytes_on_connection: u64,
last_write_end: Instant,
last_attempt_idle: Duration,
}
impl DataWriterSftp {
pub fn from_url(url: &Url, identity_file: Option<&Path>) -> Result<Self> {
let session = sftp_utils::open_session(url, identity_file)?;
let path = sftp_utils::remote_path(url);
let sftp = session.sftp()?;
let file = sftp
.create(&path)
.with_context(|| format!("failed to create remote file {path:?}"))?;
let name = sftp_utils::display_name(url);
let session: SharedSession = Arc::new(Mutex::new(session));
let keepalive = SftpKeepalive::start(Arc::clone(&session), name.clone());
let now = Instant::now();
Ok(DataWriterSftp {
file,
position: 0,
buffer: Vec::with_capacity(BUFFER_CAPACITY),
url: url.clone(),
identity_file: identity_file.map(Path::to_path_buf),
name,
session,
_keepalive: keepalive,
connected_at: now,
bytes_on_connection: 0,
last_write_end: now,
last_attempt_idle: Duration::ZERO,
})
}
#[must_use]
pub fn path_from_url(url: &Url) -> PathBuf {
sftp_utils::remote_path(url)
}
fn flush_buffer(&mut self) -> Result<()> {
if self.buffer.is_empty() {
return Ok(());
}
let blob = Blob::from(self.buffer.as_slice());
self.network_append(&blob)?;
self.buffer.clear();
Ok(())
}
}
impl NetworkWriter for DataWriterSftp {
fn try_append(&mut self, blob: &Blob) -> Result<ByteRange> {
self.last_attempt_idle = self.last_write_end.elapsed();
let pos = self.position;
self.file.write_all(blob.as_slice())?;
self.position += blob.len();
self.bytes_on_connection += blob.len();
self.last_write_end = Instant::now();
Ok(ByteRange::new(pos, blob.len()))
}
fn try_write_at(&mut self, offset: u64, blob: &Blob, restore_pos: u64) -> Result<()> {
self
.file
.seek(SeekFrom::Start(offset))
.with_context(|| format!("failed to seek to offset {offset} in '{}'", self.name))?;
self.file.write_all(blob.as_slice()).with_context(|| {
format!(
"failed to write {} bytes at offset {offset} in '{}'",
blob.len(),
self.name
)
})?;
self
.file
.seek(SeekFrom::Start(restore_pos))
.with_context(|| format!("failed to seek back to position {restore_pos} in '{}'", self.name))?;
Ok(())
}
fn try_seek(&mut self, position: u64) -> Result<()> {
self
.file
.seek(SeekFrom::Start(position))
.with_context(|| format!("failed to seek to position {position} in '{}'", self.name))?;
self.position = position;
Ok(())
}
fn reconnect(&mut self) -> Result<()> {
let path = sftp_utils::remote_path(&self.url);
log::info!(
"reconnecting SFTP writer to '{}' (previous connection: alive {:.1}s, wrote {}, idle {:.1}s before failure)",
self.name,
self.connected_at.elapsed().as_secs_f64(),
format_bytes(self.bytes_on_connection),
self.last_attempt_idle.as_secs_f64(),
);
let session = sftp_utils::open_session(&self.url, self.identity_file.as_deref())?;
let sftp = session.sftp()?;
let mut file = sftp
.open_mode(&path, OpenFlags::WRITE, 0o644, OpenType::File)
.with_context(|| format!("failed to reopen remote file {path:?} for writing"))?;
file
.seek(SeekFrom::Start(self.position))
.with_context(|| format!("failed to seek to position {} in {path:?}", self.position))?;
self.file = file;
*self.session.lock().expect("session mutex poisoned") = session;
let now = Instant::now();
self.connected_at = now;
self.bytes_on_connection = 0;
self.last_write_end = now;
Ok(())
}
fn writer_name(&self) -> &str {
&self.name
}
fn tracked_position(&self) -> u64 {
self.position
}
fn failure_context(&self) -> String {
format!(
" [conn alive {:.1}s, {} written, idle {:.1}s before this write]",
self.connected_at.elapsed().as_secs_f64(),
format_bytes(self.bytes_on_connection),
self.last_attempt_idle.as_secs_f64(),
)
}
}
impl DataWriterTrait for DataWriterSftp {
fn append(&mut self, blob: &Blob) -> Result<ByteRange> {
let offset = self.position + self.buffer.len() as u64;
self.buffer.extend_from_slice(blob.as_slice());
if self.buffer.len() >= BUFFER_CAPACITY {
self.flush_buffer()?;
}
Ok(ByteRange::new(offset, blob.len()))
}
fn write_start(&mut self, blob: &Blob) -> Result<()> {
self.flush_buffer()?;
self.network_write_start(blob)
}
fn position(&mut self) -> Result<u64> {
Ok(self.position + self.buffer.len() as u64)
}
fn set_position(&mut self, position: u64) -> Result<()> {
self.flush_buffer()?;
self.network_set_position(position)
}
fn finalize(&mut self) -> Result<()> {
self.flush_buffer()
}
}
impl Drop for DataWriterSftp {
fn drop(&mut self) {
if !self.buffer.is_empty() {
log::warn!(
"SFTP writer for '{}' dropped with {} unflushed; call finalize() before dropping",
self.name,
format_bytes(self.buffer.len() as u64),
);
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_path_from_url() {
let url = Url::parse("sftp://host/data/out.versatiles").unwrap();
assert_eq!(
DataWriterSftp::path_from_url(&url),
PathBuf::from("/data/out.versatiles")
);
}
#[test]
fn test_path_from_url_with_credentials() {
let url = Url::parse("sftp://user:pass@host:2222/output/tiles.tar").unwrap();
assert_eq!(DataWriterSftp::path_from_url(&url), PathBuf::from("/output/tiles.tar"));
}
#[test]
fn test_path_from_url_root() {
let url = Url::parse("sftp://host/file.versatiles").unwrap();
assert_eq!(DataWriterSftp::path_from_url(&url), PathBuf::from("/file.versatiles"));
}
#[test]
fn test_path_from_url_nested() {
let url = Url::parse("sftp://host/a/b/c/d.tar").unwrap();
assert_eq!(DataWriterSftp::path_from_url(&url), PathBuf::from("/a/b/c/d.tar"));
}
#[test]
fn test_from_url_unreachable_host() {
let url = Url::parse("sftp://192.0.2.1:22222/path/file.versatiles").unwrap();
let result = DataWriterSftp::from_url(&url, None);
assert!(result.is_err());
}
#[cfg(all(feature = "ssh2", unix))]
mod sftp_server_tests {
use super::*;
use crate::{Blob, io::test_sftp_server::TestSftpServer};
#[tokio::test(flavor = "multi_thread")]
async fn append_writes_bytes() {
let server = TestSftpServer::start().await;
let url = server.url("/out.bin");
tokio::task::spawn_blocking(move || -> anyhow::Result<()> {
let mut w = DataWriterSftp::from_url(&url, None)?;
w.append(&Blob::from(b"hello"))?;
w.append(&Blob::from(b"world"))?;
w.finalize()?;
Ok(())
})
.await
.unwrap()
.unwrap();
assert_eq!(server.read_file("/out.bin").await, b"helloworld");
}
#[tokio::test(flavor = "multi_thread")]
async fn write_start_overwrites_beginning() {
let server = TestSftpServer::start().await;
let url = server.url("/out.bin");
tokio::task::spawn_blocking(move || -> anyhow::Result<()> {
let mut w = DataWriterSftp::from_url(&url, None)?;
w.append(&Blob::from(b"AAAAABBBBB"))?;
w.write_start(&Blob::from(b"12345"))?;
w.finalize()?;
Ok(())
})
.await
.unwrap()
.unwrap();
assert_eq!(server.read_file("/out.bin").await, b"12345BBBBB");
}
#[tokio::test(flavor = "multi_thread")]
async fn position_tracking() {
let server = TestSftpServer::start().await;
let url = server.url("/out.bin");
tokio::task::spawn_blocking(move || -> anyhow::Result<()> {
let mut w = DataWriterSftp::from_url(&url, None)?;
assert_eq!(w.position()?, 0);
w.append(&Blob::from(b"abc"))?;
assert_eq!(w.position()?, 3);
w.append(&Blob::from(b"de"))?;
assert_eq!(w.position()?, 5);
Ok(())
})
.await
.unwrap()
.unwrap();
}
#[tokio::test(flavor = "multi_thread")]
async fn set_position_then_append() {
let server = TestSftpServer::start().await;
let url = server.url("/out.bin");
tokio::task::spawn_blocking(move || -> anyhow::Result<()> {
let mut w = DataWriterSftp::from_url(&url, None)?;
w.append(&Blob::from(vec![0u8; 10]))?;
w.set_position(5)?;
w.append(&Blob::from(vec![1u8; 5]))?;
w.finalize()?;
Ok(())
})
.await
.unwrap()
.unwrap();
assert_eq!(server.read_file("/out.bin").await, [0, 0, 0, 0, 0, 1, 1, 1, 1, 1]);
}
#[tokio::test(flavor = "multi_thread")]
async fn write_retry_after_disconnect() {
let server = TestSftpServer::start().await;
let url = server.url("/out.bin");
let mut writer = tokio::task::spawn_blocking(move || DataWriterSftp::from_url(&url, None))
.await
.unwrap()
.unwrap();
server.schedule_disconnect();
tokio::task::spawn_blocking(move || -> anyhow::Result<()> {
writer.append(&Blob::from(b"hello"))?;
writer.finalize()?;
Ok(())
})
.await
.unwrap()
.unwrap();
assert_eq!(server.read_file("/out.bin").await, b"hello");
}
#[tokio::test(flavor = "multi_thread")]
async fn many_small_appends_are_coalesced_and_flushed_on_finalize() {
let server = TestSftpServer::start().await;
let url = server.url("/out.bin");
tokio::task::spawn_blocking(move || -> anyhow::Result<()> {
let mut w = DataWriterSftp::from_url(&url, None)?;
for i in 0..1000u32 {
let range = w.append(&Blob::from(i.to_le_bytes().to_vec()))?;
assert_eq!(range.offset, u64::from(i) * 4);
}
assert_eq!(w.position()?, 4000);
w.finalize()?;
Ok(())
})
.await
.unwrap()
.unwrap();
let bytes = server.read_file("/out.bin").await;
assert_eq!(bytes.len(), 4000);
let mut expected = Vec::with_capacity(4000);
for i in 0..1000u32 {
expected.extend_from_slice(&i.to_le_bytes());
}
assert_eq!(bytes, expected);
}
#[tokio::test(flavor = "multi_thread")]
async fn append_larger_than_buffer_capacity_flushes() {
let server = TestSftpServer::start().await;
let url = server.url("/out.bin");
let big = vec![7u8; BUFFER_CAPACITY + 1024];
let expected = big.clone();
tokio::task::spawn_blocking(move || -> anyhow::Result<()> {
let mut w = DataWriterSftp::from_url(&url, None)?;
w.append(&Blob::from(big))?;
w.finalize()?;
Ok(())
})
.await
.unwrap()
.unwrap();
assert_eq!(server.read_file("/out.bin").await, expected);
}
#[tokio::test(flavor = "multi_thread")]
async fn buffered_writer_offsets_resolve_with_real_reader() {
use crate::io::{DataReaderSftp, DataReaderTrait};
let server = TestSftpServer::start().await;
let url = server.url("/round_trip.bin");
let read_url = url.clone();
let (header_range, chunks) = tokio::task::spawn_blocking(move || -> anyhow::Result<_> {
let mut w = DataWriterSftp::from_url(&url, None)?;
let header_range = w.append(&Blob::from(vec![0u8; 16]))?;
assert_eq!(header_range.offset, 0);
let mut chunks = Vec::new();
for i in 0..300u32 {
let payload = format!("tile-{i:05}-payload").into_bytes();
let range = w.append(&Blob::from(payload.clone()))?;
chunks.push((range, payload));
}
w.write_start(&Blob::from(b"VERSATILES\0\0\0\0\0\0".to_vec()))?;
w.finalize()?;
Ok((header_range, chunks))
})
.await
.unwrap()
.unwrap();
let reader = tokio::task::spawn_blocking(move || DataReaderSftp::open(&read_url, None))
.await
.unwrap()
.unwrap();
let header = reader.read_range(&header_range).await.unwrap();
assert_eq!(header.as_slice(), b"VERSATILES\0\0\0\0\0\0");
for (range, payload) in chunks {
let got = reader.read_range(&range).await.unwrap();
assert_eq!(
got.as_slice(),
payload.as_slice(),
"mismatch at offset {}",
range.offset
);
}
}
}
}