use std::io::Read;
use std::path::{Path, PathBuf};
use std::sync::Arc;
use std::sync::atomic::AtomicU64;
use color_eyre::Result;
use object_store::{ObjectStore, ObjectStoreExt};
use crate::cloud::download::TempDownload;
use crate::error_display::{FileError, file_message, store_message};
use crate::formats::ipc_stream::{Merge, Part};
use crate::loading::unfinished::Writer;
use crate::{App, FileFormat, OpenOptions};
#[derive(Debug, Clone, PartialEq, Eq)]
pub(crate) struct Object {
pub url: String,
pub size: u64,
pub stream: bool,
}
type Listing = Vec<(String, u64)>;
fn failed(url: &str, what: impl Into<String>) -> color_eyre::Report {
FileError::new(Path::new(url), what).into()
}
const PEEKS: usize = 8;
pub(crate) fn list(
url: &str,
options: &OpenOptions,
cloud: &crate::config::CloudConfig,
runtime: &tokio::runtime::Handle,
) -> Result<(Vec<Object>, OpenOptions)> {
let (full, _, store) = App::cloud_store_for(Path::new(url), cloud, runtime)?;
let (_, key) = App::cloud_bucket_and_key(&full)?;
let key = key.trim_matches('/').to_string();
let base = full.trim_end_matches('/').to_string();
let (objects, options) = if url.ends_with('/') || key.is_empty() {
let names = names_under(&store, &key, url, runtime)?;
let (base, names, options) = if names
.iter()
.any(|(name, _)| name == crate::formats::hf_splits::DATASET_DICT)
{
dict_split(&store, &key, &base, url, options, runtime)?
} else {
let (names, options) = one_split(names, options, url)?;
(format!("{base}/"), names, options)
};
let objects = names
.into_iter()
.map(|(name, size)| (format!("{base}{name}"), size))
.collect();
(objects, options)
} else {
let path = crate::cloud::cloud_browse::object_path(&key);
let store = store.clone();
let meta = crate::wait_on_runtime(runtime, async move { store.head(&path).await })
.ok_or_else(|| failed(url, "looking at it was cancelled"))?
.map_err(|e| failed(url, store_message(&e)))?;
(vec![(full.clone(), meta.size)], options.clone())
};
if let (Some(table), None) = (&options.table, &options.splits) {
return Err(failed(
url,
format!(
"it holds one table, not splits. --table {table} picks a Hugging Face dataset's split"
),
));
}
Ok((peek(&store, objects, runtime)?, options))
}
fn dict_split(
store: &Arc<dyn ObjectStore>,
key: &str,
base: &str,
url: &str,
options: &OpenOptions,
runtime: &tokio::runtime::Handle,
) -> Result<(String, Listing, OpenOptions)> {
let text = get_text(
store,
&join(key, crate::formats::hf_splits::DATASET_DICT),
&format!("{base}/{}", crate::formats::hf_splits::DATASET_DICT),
runtime,
)?;
let splits = crate::formats::hf_splits::dict_splits(&text)
.ok_or_else(|| failed(url, "its dataset_dict.json names no splits"))?;
let listed: Vec<&str> = splits.iter().map(String::as_str).collect();
let picked = crate::formats::hf_splits::pick(&listed, options.table.as_deref())
.map_err(|e| failed(url, e))?;
let split = picked.split.clone().unwrap_or_default();
let names = names_under(store, &join(key, &split), url, runtime)?;
let inner = OpenOptions {
table: None,
..options.clone()
};
let (names, inner) = one_split(names, &inner, url)?;
let options = OpenOptions {
splits: Some(Arc::new(crate::formats::hf_splits::Splits {
caches: inner.splits.as_ref().map_or(0, |s| s.caches),
..picked
})),
..options.clone()
};
Ok((format!("{base}/{split}/"), names, options))
}
fn join(key: &str, name: &str) -> String {
if key.is_empty() {
name.to_string()
} else {
format!("{key}/{name}")
}
}
fn names_under(
store: &Arc<dyn ObjectStore>,
key: &str,
url: &str,
runtime: &tokio::runtime::Handle,
) -> Result<Vec<(String, u64)>> {
let store = store.clone();
let prefix = (!key.is_empty()).then(|| crate::cloud::cloud_browse::object_path(key));
let listed = crate::wait_on_runtime(runtime, async move {
store.list_with_delimiter(prefix.as_ref()).await
})
.ok_or_else(|| failed(url, "listing it was cancelled"))?
.map_err(|e| failed(url, store_message(&e)))?;
let mut names: Vec<(String, u64)> = listed
.objects
.iter()
.filter_map(|meta| Some((meta.location.filename()?.to_string(), meta.size)))
.collect();
names.sort();
Ok(names)
}
fn one_split(
mut names: Vec<(String, u64)>,
options: &OpenOptions,
url: &str,
) -> Result<(Vec<(String, u64)>, OpenOptions)> {
let hugging_face = names
.iter()
.any(|(name, _)| crate::home::discover::is_hugging_face_metadata(name));
names.retain(|(name, _)| {
!crate::home::discover::is_bookkeeping(name)
&& crate::home::discover::data_format(Path::new(name)) == Some(FileFormat::Arrow)
});
if names.is_empty() {
return Err(failed(url, "it holds no Arrow files"));
}
let mut options = options.clone();
if hugging_face {
let listed: Vec<&str> = names.iter().map(|(name, _)| name.as_str()).collect();
let (chosen, splits) = crate::formats::hf_splits::choose(&listed, options.table.as_deref())
.map_err(|e| failed(url, e))?;
names = chosen.into_iter().map(|i| names[i].clone()).collect();
options.splits = Some(Arc::new(splits));
}
Ok((names, options))
}
fn get_text(
store: &Arc<dyn ObjectStore>,
key: &str,
url: &str,
runtime: &tokio::runtime::Handle,
) -> Result<String> {
let store = store.clone();
let path = crate::cloud::cloud_browse::object_path(key);
let bytes = crate::wait_on_runtime(
runtime,
async move { store.get(&path).await?.bytes().await },
)
.ok_or_else(|| failed(url, "reading it was cancelled"))?
.map_err(|e| failed(url, store_message(&e)))?;
Ok(String::from_utf8_lossy(&bytes).into_owned())
}
fn peek(
store: &Arc<dyn ObjectStore>,
objects: Vec<(String, u64)>,
runtime: &tokio::runtime::Handle,
) -> Result<Vec<Object>> {
use futures::StreamExt;
let store = store.clone();
let heads = crate::wait_on_runtime(runtime, async move {
futures::stream::iter(objects)
.map(|(url, size)| {
let store = store.clone();
async move {
let (_, key) = App::cloud_bucket_and_key(&url)?;
let path = crate::cloud::cloud_browse::object_path(&key);
let head = store
.get_range(&path, 0..size.min(6))
.await
.map_err(|e| failed(&url, store_message(&e)))?;
let stream = !crate::formats::ipc_stream::is_ipc_file_head(&head);
Ok::<_, color_eyre::Report>(Object { url, size, stream })
}
})
.buffered(PEEKS)
.collect::<Vec<_>>()
.await
})
.ok_or_else(|| color_eyre::eyre::eyre!("Looking at the Arrow files was cancelled."))?;
heads.into_iter().collect()
}
pub(crate) fn stream_bytes(objects: &[Object]) -> u64 {
objects.iter().filter(|o| o.stream).map(|o| o.size).sum()
}
pub(crate) fn in_place(objects: &[Object]) -> Option<Vec<Part>> {
objects.iter().all(|o| !o.stream).then(|| {
objects
.iter()
.map(|o| Part::InPlace(PathBuf::from(&o.url)))
.collect()
})
}
pub(crate) fn download(
objects: &[Object],
options: &OpenOptions,
cloud: &crate::config::CloudConfig,
runtime: &tokio::runtime::Handle,
writer: &Writer,
) -> Result<(TempDownload, Vec<Part>)> {
crate::formats::ipc_stream::has_room(stream_bytes(objects), options.temp_dir.as_deref())?;
let mut merge = Merge::create(options.temp_dir.as_deref(), writer)?;
let read = AtomicU64::new(0);
let mut parts = Vec::with_capacity(objects.len());
for object in objects {
if !object.stream {
parts.push(Part::InPlace(PathBuf::from(&object.url)));
continue;
}
let body = Body::open(&object.url, cloud, runtime, writer)?;
parts.push(merge.append(body, Path::new(&object.url), 0, &read)?);
}
Ok((merge.finish()?, parts))
}
struct Body {
chunks: std::sync::mpsc::Receiver<std::result::Result<Vec<u8>, String>>,
chunk: Vec<u8>,
at: usize,
}
impl Body {
fn open(
url: &str,
cloud: &crate::config::CloudConfig,
runtime: &tokio::runtime::Handle,
writer: &Writer,
) -> Result<Body> {
use crate::cloud::download::StreamError;
let (_, key) = App::cloud_bucket_and_key(url)?;
let (_, _, store) = App::cloud_store_for(Path::new(url), cloud, runtime)?;
let path = crate::cloud::cloud_browse::object_path(&key);
let open = async move {
let got = store.get(&path).await.map_err(|e| store_message(&e))?;
let len = got.range.end - got.range.start;
Ok((got.into_stream(), Some(len)))
};
let (tx, chunks) = std::sync::mpsc::sync_channel(crate::cloud::download::QUEUED_CHUNKS);
let runtime = runtime.clone();
let stop = {
let writer = writer.clone();
move || writer.stopped()
};
let url = url.to_string();
std::thread::Builder::new()
.name("datui-arrow-download".to_string())
.spawn(move || {
let sent = crate::cloud::download::stream_into(&runtime, open, stop, |chunk| {
tx.send(Ok(chunk.to_vec()))
.map_err(|_| color_eyre::eyre::eyre!("the conversion stopped"))
});
let named = Path::new(&url);
let message = match sent {
Ok(_) => return,
Err(StreamError::Open(e)) => file_message(named, &e),
Err(StreamError::Read(e)) => {
file_message(named, &format!("the download stopped: {e}"))
}
Err(StreamError::Short { expected, got }) => {
file_message(named, &format!("it ended after {got} of {expected} bytes"))
}
Err(StreamError::Write(e)) => e.to_string(),
Err(StreamError::Cut) => file_message(named, "the download was cancelled"),
};
let _ = tx.send(Err(message));
})?;
Ok(Body {
chunks,
chunk: Vec::new(),
at: 0,
})
}
}
impl Read for Body {
fn read(&mut self, buf: &mut [u8]) -> std::io::Result<usize> {
while self.at == self.chunk.len() {
match self.chunks.recv() {
Ok(Ok(chunk)) => {
self.chunk = chunk;
self.at = 0;
}
Ok(Err(message)) => return Err(std::io::Error::other(message)),
Err(_) => return Ok(0),
}
}
let n = buf.len().min(self.chunk.len() - self.at);
buf[..n].copy_from_slice(&self.chunk[self.at..self.at + n]);
self.at += n;
Ok(n)
}
}
#[cfg(test)]
mod tests {
use super::*;
use object_store::{PutPayload, memory::InMemory, path::Path as Key};
#[test]
fn errors_name_the_file() {
let runtime = tokio::runtime::Runtime::new().unwrap();
let handle = runtime.handle().clone();
let memory = InMemory::new();
runtime.block_on(async {
for (key, bytes) in [
("csv/a.csv", &b"a\n1\n"[..]),
("dd/dataset_dict.json", b"{\"splits\": 3}"),
("hf/data-train.arrow", b"x"),
("hf/dataset_info.json", b"{}"),
] {
memory
.put(&Key::from(key), PutPayload::from_static(bytes))
.await
.unwrap();
}
});
let store: Arc<dyn ObjectStore> = Arc::new(memory);
let options = OpenOptions::default();
let said = |e: color_eyre::Report| crate::error_display::user_message_from_report(&e, None);
let check = |message: String, url: &str, says: &str| {
eprintln!("{message}");
crate::formats::readers::bad_input::assert_shape(&message, Path::new(url));
assert!(message.contains(says), "{url}: {message}");
};
let names = names_under(&store, "csv", "s3://b/csv/", &handle).unwrap();
let none = one_split(names, &options, "s3://b/csv/").err().unwrap();
check(said(none), "s3://b/csv/", "no Arrow files");
let dict = dict_split(&store, "dd", "s3://b/dd", "s3://b/dd/", &options, &handle);
check(said(dict.err().unwrap()), "s3://b/dd/", "names no splits");
let names = names_under(&store, "hf", "s3://b/hf/", &handle).unwrap();
let asked = OpenOptions {
table: Some("nope".to_string()),
..OpenOptions::default()
};
let split = one_split(names, &asked, "s3://b/hf/").err().unwrap();
check(said(split), "s3://b/hf/", "nope");
let gone = get_text(&store, "gone.json", "s3://b/gone.json", &handle).unwrap_err();
check(said(gone), "s3://b/gone.json", "No object there");
let peeked = peek(&store, vec![("s3://b/gone.arrow".to_string(), 8)], &handle);
check(
said(peeked.unwrap_err()),
"s3://b/gone.arrow",
"No object there",
);
}
}