Skip to main content

dataforge/
lib.rs

1//! Fast dataset discovery and archive helpers for `thesa` and other Rust tools.
2//!
3//! Made by Trevor Knott for Knott Dynamics.
4//!
5//! DataForge normalizes dataset targets, queries public dataset catalogs, and
6//! archives dataset files with bounded concurrent downloads. The first-class
7//! provider is Hugging Face Datasets; Zenodo is included for DOI/record-backed
8//! research datasets.
9//!
10//! ```
11//! assert_eq!(
12//!     dataforge::parse_dataset_target_for_provider(
13//!         "https://huggingface.co/datasets/squad",
14//!         dataforge::DatasetProvider::Hf,
15//!     )
16//!     .unwrap(),
17//!     dataforge::DatasetTarget::Dataset("squad".to_string())
18//! );
19//! assert_eq!(
20//!     dataforge::DatasetProvider::from_slug("hugging-face"),
21//!     Some(dataforge::DatasetProvider::Hf)
22//! );
23//! ```
24
25use std::collections::{BTreeSet, VecDeque};
26use std::fmt;
27use std::fs::{self, File};
28use std::io;
29use std::path::{Path, PathBuf};
30use std::sync::{Arc, Mutex};
31use std::thread;
32
33use reqwest::StatusCode;
34use reqwest::blocking::{Client, Response};
35use serde::{Deserialize, Serialize};
36
37const HF_BASE: &str = "https://huggingface.co";
38const ZENODO_API: &str = "https://zenodo.org/api";
39const DEFAULT_PAGE_SIZE: usize = 100;
40const DEFAULT_CONCURRENCY: usize = 4;
41const DEFAULT_USER_AGENT: &str = "dataforge/0.1";
42
43/// Result type used by DataForge APIs.
44pub type Result<T> = std::result::Result<T, DataForgeError>;
45
46/// Supported dataset providers.
47#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
48pub enum DatasetProvider {
49    /// Hugging Face Datasets on the Hugging Face Hub.
50    Hf,
51    /// Zenodo records and files.
52    Zenodo,
53}
54
55impl DatasetProvider {
56    /// Stable provider list in CLI display order.
57    pub const ALL: [Self; 2] = [Self::Hf, Self::Zenodo];
58
59    /// Return all supported providers in stable CLI display order.
60    pub fn all() -> &'static [Self] {
61        &Self::ALL
62    }
63
64    /// Human-readable provider label.
65    pub fn label(self) -> &'static str {
66        match self {
67            Self::Hf => "Hugging Face Datasets",
68            Self::Zenodo => "Zenodo",
69        }
70    }
71
72    /// Compact UI label.
73    pub fn short_label(self) -> &'static str {
74        match self {
75            Self::Hf => "HF",
76            Self::Zenodo => "ZEN",
77        }
78    }
79
80    /// Stable lowercase provider slug.
81    pub fn slug(self) -> &'static str {
82        match self {
83            Self::Hf => "hf",
84            Self::Zenodo => "zenodo",
85        }
86    }
87
88    /// Parse common provider aliases accepted by thesa-style CLIs.
89    pub fn from_slug(value: &str) -> Option<Self> {
90        let normalized = normalized_slug(value);
91        match normalized.as_str() {
92            "hf" | "huggingface" | "huggingfacedatasets" | "hfdatasets" => Some(Self::Hf),
93            "zen" | "zenodo" => Some(Self::Zenodo),
94            _ => None,
95        }
96    }
97}
98
99impl fmt::Display for DatasetProvider {
100    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
101        write!(f, "{}", self.slug())
102    }
103}
104
105/// Normalized dataset target requested by a user or CLI.
106#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
107pub enum DatasetTarget {
108    /// Provider namespace, owner, or organization.
109    Namespace(String),
110    /// Concrete provider dataset id or record id.
111    Dataset(String),
112    /// Provider search query.
113    Search(String),
114    /// Provider-specific top/trending/popular preset.
115    Top,
116    /// Provider-specific latest/newest preset.
117    Latest,
118}
119
120impl DatasetTarget {
121    /// Stable target kind label.
122    pub fn kind(&self) -> &'static str {
123        match self {
124            Self::Namespace(_) => "namespace",
125            Self::Dataset(_) => "dataset",
126            Self::Search(_) => "search",
127            Self::Top => "top",
128            Self::Latest => "latest",
129        }
130    }
131
132    /// Borrow the contained value for namespace, dataset, or search targets.
133    pub fn value(&self) -> Option<&str> {
134        match self {
135            Self::Namespace(value) | Self::Dataset(value) | Self::Search(value) => Some(value),
136            Self::Top | Self::Latest => None,
137        }
138    }
139
140    /// Return true when this target points at one concrete dataset/record.
141    pub fn is_dataset(&self) -> bool {
142        matches!(self, Self::Dataset(_))
143    }
144
145    /// Return true when this target may expand into many datasets.
146    pub fn is_collection(&self) -> bool {
147        matches!(
148            self,
149            Self::Namespace(_) | Self::Search(_) | Self::Top | Self::Latest
150        )
151    }
152}
153
154/// Lightweight dataset metadata returned by discovery calls.
155#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
156pub struct DatasetSummary {
157    /// Provider that produced the dataset.
158    pub provider: DatasetProvider,
159    /// Stable provider id, such as `owner/name` or a Zenodo record id.
160    pub id: String,
161    /// Human-readable title when the provider exposes one.
162    pub title: Option<String>,
163    /// Download count or similar popularity signal.
164    pub downloads: Option<u64>,
165    /// Like/favorite count when available.
166    pub likes: Option<u64>,
167    /// Last modification timestamp as returned by the provider.
168    pub last_modified: Option<String>,
169    /// Short description when available.
170    pub description: Option<String>,
171    /// Provider tags, keywords, or topics.
172    pub tags: Vec<String>,
173    /// Number of downloadable files if known.
174    pub file_count: Option<usize>,
175    /// Total known file size in bytes.
176    pub size: Option<u64>,
177    /// Canonical web URL for humans.
178    pub url: Option<String>,
179}
180
181/// One downloadable file in a dataset archive.
182#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
183pub struct DatasetFile {
184    /// Provider-relative file path.
185    pub path: String,
186    /// Direct download URL.
187    pub download_url: String,
188    /// File size in bytes when known.
189    pub size: Option<u64>,
190    /// Provider checksum string when known.
191    pub checksum: Option<String>,
192}
193
194/// Full dataset metadata plus files to archive.
195#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
196pub struct DatasetDetails {
197    /// Dataset summary.
198    pub summary: DatasetSummary,
199    /// Downloadable files.
200    pub files: Vec<DatasetFile>,
201}
202
203/// Options controlling archive writes.
204#[derive(Debug, Clone)]
205pub struct ArchiveOptions {
206    /// Directory where dataset subdirectories are written.
207    pub output: PathBuf,
208    /// Maximum number of files downloaded at once.
209    pub concurrency: usize,
210    /// Skip a dataset directory when it already exists.
211    pub skip_existing: bool,
212    /// Optional substring filter applied to discovered dataset ids and titles.
213    pub filter: Option<String>,
214}
215
216impl ArchiveOptions {
217    /// Create archive options for an output directory.
218    pub fn new(output: impl Into<PathBuf>) -> Self {
219        Self {
220            output: output.into(),
221            ..Self::default()
222        }
223    }
224}
225
226impl Default for ArchiveOptions {
227    fn default() -> Self {
228        Self {
229            output: PathBuf::from("datasets"),
230            concurrency: DEFAULT_CONCURRENCY,
231            skip_existing: false,
232            filter: None,
233        }
234    }
235}
236
237/// One archived or skipped file.
238#[derive(Debug, Clone, PartialEq, Eq)]
239pub struct ArchivedFile {
240    /// Dataset id that owns the file.
241    pub dataset_id: String,
242    /// Provider-relative file path.
243    pub source_path: String,
244    /// Local path written or skipped.
245    pub local_path: PathBuf,
246    /// Number of bytes written during this run.
247    pub bytes_written: u64,
248    /// Whether an existing file was skipped.
249    pub skipped: bool,
250}
251
252/// Failed dataset or file archive entry.
253#[derive(Debug, Clone, PartialEq, Eq)]
254pub struct ArchiveFailure {
255    /// Dataset id that failed.
256    pub dataset_id: String,
257    /// Optional provider-relative file path.
258    pub file: Option<String>,
259    /// Human-readable failure message.
260    pub message: String,
261}
262
263/// Summary returned after an archive run.
264#[derive(Debug, Clone, Default, PartialEq, Eq)]
265pub struct ArchiveSummary {
266    /// Number of datasets discovered before filtering and archiving.
267    pub discovered: usize,
268    /// Number of datasets selected after filter application.
269    pub selected: usize,
270    /// Number of dataset directories completed without file failures.
271    pub archived: usize,
272    /// Number of dataset directories skipped because they already existed.
273    pub skipped: usize,
274    /// Number of files written.
275    pub files_archived: usize,
276    /// Bytes written during this run.
277    pub bytes_written: u64,
278    /// Per-file successful writes or skips.
279    pub files: Vec<ArchivedFile>,
280    /// Dataset or file failures.
281    pub failed: Vec<ArchiveFailure>,
282}
283
284impl ArchiveSummary {
285    /// Return true when no archive failures occurred.
286    pub fn is_success(&self) -> bool {
287        self.failed.is_empty()
288    }
289
290    /// Total attempted datasets after filtering.
291    pub fn attempted(&self) -> usize {
292        self.archived + self.skipped + self.failed_dataset_count()
293    }
294
295    /// Number of unique datasets with at least one failure.
296    pub fn failed_dataset_count(&self) -> usize {
297        self.failed
298            .iter()
299            .map(|failure| failure.dataset_id.as_str())
300            .collect::<BTreeSet<_>>()
301            .len()
302    }
303}
304
305/// Errors returned by DataForge.
306#[derive(Debug, thiserror::Error)]
307pub enum DataForgeError {
308    /// The provided target string could not be normalized.
309    #[error("invalid dataset target '{0}'")]
310    InvalidTarget(String),
311
312    /// The requested provider target is not supported by the current operation.
313    #[error("unsupported dataset provider target: {0}")]
314    UnsupportedTarget(String),
315
316    /// Provider API returned a non-success status.
317    #[error("{provider} API returned {status}: {body}")]
318    ProviderApi {
319        /// Provider slug.
320        provider: &'static str,
321        /// HTTP status code.
322        status: u16,
323        /// Response body or fallback message.
324        body: String,
325    },
326
327    /// Provider did not find the requested dataset or record.
328    #[error("{provider} dataset '{target}' was not found")]
329    NotFound {
330        /// Provider slug.
331        provider: &'static str,
332        /// Requested target.
333        target: String,
334    },
335
336    /// Download failed for one dataset file.
337    #[error("download failed for '{dataset}' file '{file}': {message}")]
338    DownloadFailed {
339        /// Dataset id.
340        dataset: String,
341        /// Provider-relative file path.
342        file: String,
343        /// Human-readable failure message.
344        message: String,
345    },
346
347    /// Output path exists but is not a directory.
348    #[error("output path '{}' is not a directory", .0.display())]
349    OutputNotDirectory(PathBuf),
350
351    /// Filesystem I/O failure.
352    #[error("I/O error: {0}")]
353    Io(#[from] io::Error),
354
355    /// HTTP request failure.
356    #[error("request error: {0}")]
357    Request(#[from] reqwest::Error),
358
359    /// JSON serialization failure.
360    #[error("JSON error: {0}")]
361    Json(#[from] serde_json::Error),
362}
363
364/// Blocking dataset discovery and archive client.
365#[derive(Clone)]
366pub struct DataForge {
367    provider: DatasetProvider,
368    client: Client,
369    token: Option<String>,
370    page_size: usize,
371}
372
373impl DataForge {
374    /// Build a new client for the selected provider.
375    pub fn new(provider: DatasetProvider) -> Result<Self> {
376        Self::with_user_agent(provider, DEFAULT_USER_AGENT)
377    }
378
379    /// Build a new client with a custom User-Agent.
380    pub fn with_user_agent(provider: DatasetProvider, user_agent: &str) -> Result<Self> {
381        let client = Client::builder().user_agent(user_agent).build()?;
382        Ok(Self {
383            provider,
384            client,
385            token: None,
386            page_size: DEFAULT_PAGE_SIZE,
387        })
388    }
389
390    /// Attach a provider token. Hugging Face uses `Authorization: Bearer <token>`.
391    pub fn with_token(mut self, token: impl Into<String>) -> Self {
392        let token = token.into();
393        if !token.trim().is_empty() {
394            self.token = Some(token);
395        }
396        self
397    }
398
399    /// Set the API page size used for list requests.
400    pub fn with_page_size(mut self, page_size: usize) -> Self {
401        self.page_size = page_size.max(1);
402        self
403    }
404
405    /// Return the provider this client targets.
406    pub fn provider(&self) -> DatasetProvider {
407        self.provider
408    }
409
410    /// Discover datasets for a normalized target.
411    pub fn discover(&self, target: &DatasetTarget) -> Result<Vec<DatasetSummary>> {
412        match self.provider {
413            DatasetProvider::Hf => {
414                discover_hf(&self.client, target, self.token.as_deref(), self.page_size)
415            }
416            DatasetProvider::Zenodo => discover_zenodo(&self.client, target, self.page_size),
417        }
418    }
419
420    /// Fetch full details and file URLs for one dataset id.
421    pub fn inspect(&self, dataset_id: &str) -> Result<DatasetDetails> {
422        match self.provider {
423            DatasetProvider::Hf => inspect_hf(&self.client, dataset_id, self.token.as_deref()),
424            DatasetProvider::Zenodo => inspect_zenodo(&self.client, dataset_id),
425        }
426    }
427
428    /// Discover and archive datasets for a target.
429    pub fn archive_target(
430        &self,
431        target: &DatasetTarget,
432        options: &ArchiveOptions,
433    ) -> Result<ArchiveSummary> {
434        let datasets = self.discover(target)?;
435        self.archive_datasets(&datasets, options)
436    }
437
438    /// Archive discovered dataset summaries with bounded concurrent file downloads.
439    pub fn archive_datasets(
440        &self,
441        datasets: &[DatasetSummary],
442        options: &ArchiveOptions,
443    ) -> Result<ArchiveSummary> {
444        prepare_output(&options.output)?;
445
446        let mut summary = ArchiveSummary {
447            discovered: datasets.len(),
448            ..ArchiveSummary::default()
449        };
450        let selected = datasets
451            .iter()
452            .filter(|dataset| dataset_matches_filter(dataset, options.filter.as_deref()))
453            .cloned()
454            .collect::<Vec<_>>();
455        summary.selected = selected.len();
456
457        let mut jobs = VecDeque::new();
458        let mut scheduled_datasets = BTreeSet::new();
459
460        for dataset in selected {
461            let dataset_dir = options.output.join(dataset_dir_name(&dataset.id));
462            if dataset_dir.exists() {
463                if options.skip_existing {
464                    summary.skipped += 1;
465                    continue;
466                }
467                if dataset_dir.is_dir() {
468                    fs::remove_dir_all(&dataset_dir)?;
469                } else {
470                    fs::remove_file(&dataset_dir)?;
471                }
472            }
473            fs::create_dir_all(&dataset_dir)?;
474
475            let details = match self.inspect(&dataset.id) {
476                Ok(details) => details,
477                Err(error) => {
478                    summary.failed.push(ArchiveFailure {
479                        dataset_id: dataset.id,
480                        file: None,
481                        message: error.to_string(),
482                    });
483                    continue;
484                }
485            };
486
487            write_dataset_manifest(&dataset_dir, &details)?;
488            scheduled_datasets.insert(details.summary.id.clone());
489
490            for file in details.files {
491                let local_path = dataset_dir.join(safe_relative_path(&file.path));
492                jobs.push_back(DownloadJob {
493                    dataset_id: details.summary.id.clone(),
494                    source_path: file.path,
495                    url: file.download_url,
496                    local_path,
497                });
498            }
499        }
500
501        if jobs.is_empty() {
502            summary.archived += scheduled_datasets.len();
503            return Ok(summary);
504        }
505
506        let results = run_download_jobs(
507            self.client.clone(),
508            self.token.clone(),
509            options.concurrency,
510            jobs,
511        );
512        let mut failed_datasets = BTreeSet::new();
513
514        for result in results {
515            match result {
516                Ok(file) => {
517                    summary.bytes_written += file.bytes_written;
518                    if !file.skipped {
519                        summary.files_archived += 1;
520                    }
521                    summary.files.push(file);
522                }
523                Err(failure) => {
524                    failed_datasets.insert(failure.dataset_id.clone());
525                    summary.failed.push(failure);
526                }
527            }
528        }
529
530        summary.archived += scheduled_datasets.difference(&failed_datasets).count();
531
532        Ok(summary)
533    }
534}
535
536/// Parse a Hugging Face-style dataset target.
537pub fn parse_dataset_target(input: &str) -> Result<DatasetTarget> {
538    parse_dataset_target_for_provider(input, DatasetProvider::Hf)
539}
540
541/// Parse a dataset target using provider-specific rules.
542pub fn parse_dataset_target_for_provider(
543    input: &str,
544    provider: DatasetProvider,
545) -> Result<DatasetTarget> {
546    let normalized = input.trim();
547    if normalized.is_empty() {
548        return Err(DataForgeError::InvalidTarget(
549            "empty dataset target".to_string(),
550        ));
551    }
552
553    let canonical = normalized.to_lowercase();
554    match canonical.as_str() {
555        "top" | "trending" | "popular" => return Ok(DatasetTarget::Top),
556        "latest" | "newest" | "new" => return Ok(DatasetTarget::Latest),
557        _ => {}
558    }
559
560    if let Some(value) = normalized.strip_prefix("dataset:") {
561        return non_empty_target(value, DatasetTarget::Dataset);
562    }
563    if let Some(value) = normalized.strip_prefix("search:") {
564        return non_empty_target(value, DatasetTarget::Search);
565    }
566    if let Some(value) = normalized
567        .strip_prefix("org:")
568        .or_else(|| normalized.strip_prefix("owner:"))
569        .or_else(|| normalized.strip_prefix("namespace:"))
570    {
571        return non_empty_target(value, DatasetTarget::Namespace);
572    }
573
574    match provider {
575        DatasetProvider::Hf => parse_hf_target(normalized),
576        DatasetProvider::Zenodo => parse_zenodo_target(normalized),
577    }
578}
579
580fn discover_hf(
581    client: &Client,
582    target: &DatasetTarget,
583    token: Option<&str>,
584    page_size: usize,
585) -> Result<Vec<DatasetSummary>> {
586    match target {
587        DatasetTarget::Dataset(id) => Ok(vec![inspect_hf(client, id, token)?.summary]),
588        DatasetTarget::Namespace(namespace) => list_hf_datasets(
589            client,
590            Some(("author", namespace)),
591            "downloads",
592            -1,
593            token,
594            page_size,
595        ),
596        DatasetTarget::Search(query) => list_hf_datasets(
597            client,
598            Some(("search", query)),
599            "downloads",
600            -1,
601            token,
602            page_size,
603        ),
604        DatasetTarget::Top => list_hf_datasets(client, None, "downloads", -1, token, page_size),
605        DatasetTarget::Latest => {
606            list_hf_datasets(client, None, "lastModified", -1, token, page_size)
607        }
608    }
609}
610
611fn list_hf_datasets(
612    client: &Client,
613    filter: Option<(&str, &str)>,
614    sort: &str,
615    direction: i8,
616    token: Option<&str>,
617    page_size: usize,
618) -> Result<Vec<DatasetSummary>> {
619    let mut page = 1usize;
620    let mut datasets = Vec::new();
621
622    loop {
623        let mut url = format!(
624            "{HF_BASE}/api/datasets?sort={}&direction={direction}&limit={}&page={page}",
625            urlencoding::encode(sort),
626            page_size
627        );
628        if let Some((key, value)) = filter {
629            url.push('&');
630            url.push_str(key);
631            url.push('=');
632            url.push_str(&urlencoding::encode(value));
633        }
634
635        let response = send_hf(client.get(&url), token)?;
636        let status = response.status();
637        if !status.is_success() {
638            return Err(api_error(DatasetProvider::Hf, response));
639        }
640
641        let batch: Vec<HfDatasetResponse> = response.json()?;
642        let count = batch.len();
643        datasets.extend(batch.into_iter().map(hf_summary));
644        if count < page_size {
645            break;
646        }
647        page = page.saturating_add(1);
648    }
649
650    Ok(datasets)
651}
652
653fn inspect_hf(client: &Client, dataset_id: &str, token: Option<&str>) -> Result<DatasetDetails> {
654    let url = format!("{HF_BASE}/api/datasets/{}", encode_repo_id(dataset_id));
655    let response = send_hf(client.get(&url), token)?;
656    let status = response.status();
657
658    if status == StatusCode::NOT_FOUND {
659        return Err(DataForgeError::NotFound {
660            provider: DatasetProvider::Hf.slug(),
661            target: dataset_id.to_string(),
662        });
663    }
664    if !status.is_success() {
665        return Err(api_error(DatasetProvider::Hf, response));
666    }
667
668    let data: HfDatasetResponse = response.json()?;
669    let summary = hf_summary(data.clone());
670    let files = data
671        .siblings
672        .into_iter()
673        .filter(|file| !file.rfilename.trim().is_empty())
674        .map(|file| DatasetFile {
675            download_url: hf_file_url(&summary.id, &file.rfilename),
676            path: file.rfilename,
677            size: file.size,
678            checksum: file.blob_id,
679        })
680        .collect();
681
682    Ok(DatasetDetails { summary, files })
683}
684
685fn discover_zenodo(
686    client: &Client,
687    target: &DatasetTarget,
688    page_size: usize,
689) -> Result<Vec<DatasetSummary>> {
690    match target {
691        DatasetTarget::Dataset(id) => Ok(vec![inspect_zenodo(client, id)?.summary]),
692        DatasetTarget::Namespace(value) | DatasetTarget::Search(value) => {
693            list_zenodo(client, Some(value), "mostviewed", page_size)
694        }
695        DatasetTarget::Top => list_zenodo(client, None, "mostviewed", page_size),
696        DatasetTarget::Latest => list_zenodo(client, None, "newest", page_size),
697    }
698}
699
700fn list_zenodo(
701    client: &Client,
702    query: Option<&str>,
703    sort: &str,
704    page_size: usize,
705) -> Result<Vec<DatasetSummary>> {
706    let mut page = 1usize;
707    let mut datasets = Vec::new();
708
709    loop {
710        let mut url = format!(
711            "{ZENODO_API}/records?sort={}&size={}&page={page}",
712            urlencoding::encode(sort),
713            page_size
714        );
715        if let Some(query) = query {
716            url.push_str("&q=");
717            url.push_str(&urlencoding::encode(query));
718        }
719
720        let response = client.get(&url).send()?;
721        let status = response.status();
722        if !status.is_success() {
723            return Err(api_error(DatasetProvider::Zenodo, response));
724        }
725
726        let batch: ZenodoSearchResponse = response.json()?;
727        let count = batch.hits.hits.len();
728        datasets.extend(batch.hits.hits.into_iter().map(zenodo_summary));
729        if count < page_size {
730            break;
731        }
732        page = page.saturating_add(1);
733    }
734
735    Ok(datasets)
736}
737
738fn inspect_zenodo(client: &Client, record_id: &str) -> Result<DatasetDetails> {
739    let url = format!("{ZENODO_API}/records/{}", urlencoding::encode(record_id));
740    let response = client.get(&url).send()?;
741    let status = response.status();
742
743    if status == StatusCode::NOT_FOUND {
744        return Err(DataForgeError::NotFound {
745            provider: DatasetProvider::Zenodo.slug(),
746            target: record_id.to_string(),
747        });
748    }
749    if !status.is_success() {
750        return Err(api_error(DatasetProvider::Zenodo, response));
751    }
752
753    let data: ZenodoRecord = response.json()?;
754    let summary = zenodo_summary(data.clone());
755    let files = data
756        .files
757        .into_iter()
758        .filter_map(|file| {
759            let url = file.links.self_url.or(file.links.download)?;
760            Some(DatasetFile {
761                path: file.key,
762                download_url: url,
763                size: file.size,
764                checksum: file.checksum,
765            })
766        })
767        .collect();
768
769    Ok(DatasetDetails { summary, files })
770}
771
772fn run_download_jobs(
773    client: Client,
774    token: Option<String>,
775    concurrency: usize,
776    jobs: VecDeque<DownloadJob>,
777) -> Vec<std::result::Result<ArchivedFile, ArchiveFailure>> {
778    let worker_count = concurrency.max(1).min(jobs.len().max(1));
779    let jobs = Arc::new(Mutex::new(jobs));
780    let results = Arc::new(Mutex::new(Vec::new()));
781    let mut workers = Vec::with_capacity(worker_count);
782
783    for _ in 0..worker_count {
784        let jobs = Arc::clone(&jobs);
785        let results = Arc::clone(&results);
786        let client = client.clone();
787        let token = token.clone();
788        workers.push(thread::spawn(move || {
789            loop {
790                let job = {
791                    let mut queue = jobs.lock().expect("download queue lock poisoned");
792                    queue.pop_front()
793                };
794                let Some(job) = job else {
795                    break;
796                };
797                let result = download_job(&client, token.as_deref(), job);
798                results
799                    .lock()
800                    .expect("download results lock poisoned")
801                    .push(result);
802            }
803        }));
804    }
805
806    for worker in workers {
807        let _ = worker.join();
808    }
809
810    Arc::try_unwrap(results)
811        .unwrap_or_else(|arc| {
812            Mutex::new(arc.lock().expect("download results lock poisoned").clone())
813        })
814        .into_inner()
815        .expect("download results lock poisoned")
816}
817
818fn download_job(
819    client: &Client,
820    token: Option<&str>,
821    job: DownloadJob,
822) -> std::result::Result<ArchivedFile, ArchiveFailure> {
823    if let Some(parent) = job.local_path.parent() {
824        if let Err(error) = fs::create_dir_all(parent) {
825            return Err(job.failure(None, error));
826        }
827    }
828
829    let mut request = client.get(&job.url);
830    if let Some(token) = token {
831        request = request.header("Authorization", format!("Bearer {token}"));
832    }
833
834    let mut response = match request.send() {
835        Ok(response) => response,
836        Err(error) => return Err(job.failure(None, error)),
837    };
838    let status = response.status();
839    if !status.is_success() {
840        let body = response
841            .text()
842            .unwrap_or_else(|_| "<unable to read body>".to_string());
843        return Err(job.failure(None, format!("HTTP {status}: {body}")));
844    }
845
846    let tmp_path = partial_path(&job.local_path);
847    let mut file = match File::create(&tmp_path) {
848        Ok(file) => file,
849        Err(error) => return Err(job.failure(None, error)),
850    };
851    let bytes_written = match io::copy(&mut response, &mut file) {
852        Ok(bytes) => bytes,
853        Err(error) => {
854            let _ = fs::remove_file(&tmp_path);
855            return Err(job.failure(None, error));
856        }
857    };
858    if let Err(error) = fs::rename(&tmp_path, &job.local_path) {
859        let _ = fs::remove_file(&tmp_path);
860        return Err(job.failure(None, error));
861    }
862
863    Ok(ArchivedFile {
864        dataset_id: job.dataset_id,
865        source_path: job.source_path,
866        local_path: job.local_path,
867        bytes_written,
868        skipped: false,
869    })
870}
871
872#[derive(Debug, Clone)]
873struct DownloadJob {
874    dataset_id: String,
875    source_path: String,
876    url: String,
877    local_path: PathBuf,
878}
879
880impl DownloadJob {
881    fn failure(&self, file: Option<String>, error: impl fmt::Display) -> ArchiveFailure {
882        ArchiveFailure {
883            dataset_id: self.dataset_id.clone(),
884            file: file.or_else(|| Some(self.source_path.clone())),
885            message: error.to_string(),
886        }
887    }
888}
889
890#[derive(Debug, Clone, Deserialize)]
891struct HfDatasetResponse {
892    #[serde(default)]
893    id: String,
894    #[serde(default)]
895    downloads: u64,
896    #[serde(default)]
897    likes: u64,
898    #[serde(rename = "lastModified", default)]
899    last_modified: Option<String>,
900    #[serde(default)]
901    tags: Vec<String>,
902    #[serde(default)]
903    siblings: Vec<HfSibling>,
904    #[serde(default)]
905    description: Option<String>,
906}
907
908#[derive(Debug, Clone, Deserialize)]
909struct HfSibling {
910    #[serde(default)]
911    rfilename: String,
912    #[serde(default)]
913    size: Option<u64>,
914    #[serde(rename = "blobId", default)]
915    blob_id: Option<String>,
916}
917
918#[derive(Debug, Clone, Deserialize)]
919struct ZenodoSearchResponse {
920    hits: ZenodoHits,
921}
922
923#[derive(Debug, Clone, Deserialize)]
924struct ZenodoHits {
925    hits: Vec<ZenodoRecord>,
926}
927
928#[derive(Debug, Clone, Deserialize)]
929struct ZenodoRecord {
930    id: u64,
931    #[serde(default)]
932    metadata: ZenodoMetadata,
933    #[serde(default)]
934    updated: Option<String>,
935    #[serde(default)]
936    created: Option<String>,
937    #[serde(default)]
938    files: Vec<ZenodoFile>,
939    #[serde(default)]
940    stats: ZenodoStats,
941    #[serde(default)]
942    links: ZenodoRecordLinks,
943}
944
945#[derive(Debug, Clone, Default, Deserialize)]
946struct ZenodoMetadata {
947    #[serde(default)]
948    title: Option<String>,
949    #[serde(default)]
950    description: Option<String>,
951    #[serde(default)]
952    keywords: Vec<String>,
953}
954
955#[derive(Debug, Clone, Default, Deserialize)]
956struct ZenodoStats {
957    #[serde(default)]
958    downloads: Option<u64>,
959}
960
961#[derive(Debug, Clone, Default, Deserialize)]
962struct ZenodoRecordLinks {
963    #[serde(default)]
964    html: Option<String>,
965}
966
967#[derive(Debug, Clone, Deserialize)]
968struct ZenodoFile {
969    key: String,
970    #[serde(default)]
971    size: Option<u64>,
972    #[serde(default)]
973    checksum: Option<String>,
974    #[serde(default)]
975    links: ZenodoFileLinks,
976}
977
978#[derive(Debug, Clone, Default, Deserialize)]
979struct ZenodoFileLinks {
980    #[serde(rename = "self", default)]
981    self_url: Option<String>,
982    #[serde(default)]
983    download: Option<String>,
984}
985
986fn hf_summary(data: HfDatasetResponse) -> DatasetSummary {
987    let file_count = if data.siblings.is_empty() {
988        None
989    } else {
990        Some(data.siblings.len())
991    };
992    let size = data
993        .siblings
994        .iter()
995        .try_fold(0u64, |acc, file| file.size.map(|size| acc + size));
996    DatasetSummary {
997        provider: DatasetProvider::Hf,
998        id: data.id.clone(),
999        title: Some(data.id.clone()),
1000        downloads: Some(data.downloads),
1001        likes: Some(data.likes),
1002        last_modified: data.last_modified,
1003        description: data.description,
1004        tags: data.tags,
1005        file_count,
1006        size,
1007        url: Some(format!("{HF_BASE}/datasets/{}", encode_repo_id(&data.id))),
1008    }
1009}
1010
1011fn zenodo_summary(data: ZenodoRecord) -> DatasetSummary {
1012    let size = data
1013        .files
1014        .iter()
1015        .try_fold(0u64, |acc, file| file.size.map(|size| acc + size));
1016    let file_count = if data.files.is_empty() {
1017        None
1018    } else {
1019        Some(data.files.len())
1020    };
1021    DatasetSummary {
1022        provider: DatasetProvider::Zenodo,
1023        id: data.id.to_string(),
1024        title: data.metadata.title,
1025        downloads: data.stats.downloads,
1026        likes: None,
1027        last_modified: data.updated.or(data.created),
1028        description: data.metadata.description,
1029        tags: data.metadata.keywords,
1030        file_count,
1031        size,
1032        url: data.links.html,
1033    }
1034}
1035
1036fn parse_hf_target(value: &str) -> Result<DatasetTarget> {
1037    if let Some(path) = strip_any_prefix_ignore_ascii_case(
1038        value,
1039        &[
1040            "https://huggingface.co/datasets/",
1041            "http://huggingface.co/datasets/",
1042            "https://www.huggingface.co/datasets/",
1043            "http://www.huggingface.co/datasets/",
1044            "https://hf.co/datasets/",
1045            "http://hf.co/datasets/",
1046            "https://www.hf.co/datasets/",
1047            "http://www.hf.co/datasets/",
1048        ],
1049    ) {
1050        return parse_dataset_path(path, true);
1051    }
1052
1053    if value.contains(char::is_whitespace) {
1054        return Ok(DatasetTarget::Search(value.to_string()));
1055    }
1056
1057    parse_dataset_path(value, false)
1058}
1059
1060fn parse_zenodo_target(value: &str) -> Result<DatasetTarget> {
1061    if let Some(path) = strip_any_prefix_ignore_ascii_case(
1062        value,
1063        &[
1064            "https://zenodo.org/records/",
1065            "http://zenodo.org/records/",
1066            "https://www.zenodo.org/records/",
1067            "http://www.zenodo.org/records/",
1068            "https://zenodo.org/record/",
1069            "http://zenodo.org/record/",
1070            "https://www.zenodo.org/record/",
1071            "http://www.zenodo.org/record/",
1072        ],
1073    ) {
1074        let id = path.split(['?', '#', '/']).next().unwrap_or_default();
1075        return non_empty_target(id, DatasetTarget::Dataset);
1076    }
1077
1078    if value.chars().all(|ch| ch.is_ascii_digit()) {
1079        return Ok(DatasetTarget::Dataset(value.to_string()));
1080    }
1081    Ok(DatasetTarget::Search(value.to_string()))
1082}
1083
1084fn parse_dataset_path(path: &str, explicit_dataset_url: bool) -> Result<DatasetTarget> {
1085    let parts = path
1086        .split(['?', '#'])
1087        .next()
1088        .unwrap_or_default()
1089        .split('/')
1090        .filter(|part| !part.trim().is_empty())
1091        .collect::<Vec<_>>();
1092    match parts.len() {
1093        0 => Err(DataForgeError::InvalidTarget(path.to_string())),
1094        1 if explicit_dataset_url => Ok(DatasetTarget::Dataset(parts[0].to_string())),
1095        1 => Ok(DatasetTarget::Namespace(parts[0].to_string())),
1096        _ => Ok(DatasetTarget::Dataset(format!("{}/{}", parts[0], parts[1]))),
1097    }
1098}
1099
1100fn non_empty_target(
1101    value: &str,
1102    build: impl FnOnce(String) -> DatasetTarget,
1103) -> Result<DatasetTarget> {
1104    let value = value.trim();
1105    if value.is_empty() {
1106        Err(DataForgeError::InvalidTarget(value.to_string()))
1107    } else {
1108        Ok(build(value.to_string()))
1109    }
1110}
1111
1112fn send_hf(
1113    mut request: reqwest::blocking::RequestBuilder,
1114    token: Option<&str>,
1115) -> reqwest::Result<Response> {
1116    if let Some(token) = token {
1117        request = request.header("Authorization", format!("Bearer {token}"));
1118    }
1119    request.send()
1120}
1121
1122fn api_error(provider: DatasetProvider, response: Response) -> DataForgeError {
1123    let status = response.status();
1124    let body = response
1125        .text()
1126        .unwrap_or_else(|_| "<unable to read body>".to_string());
1127    DataForgeError::ProviderApi {
1128        provider: provider.slug(),
1129        status: status.as_u16(),
1130        body,
1131    }
1132}
1133
1134fn prepare_output(output: &Path) -> Result<()> {
1135    if output.exists() && !output.is_dir() {
1136        return Err(DataForgeError::OutputNotDirectory(output.to_path_buf()));
1137    }
1138    fs::create_dir_all(output)?;
1139    Ok(())
1140}
1141
1142fn write_dataset_manifest(dataset_dir: &Path, details: &DatasetDetails) -> Result<()> {
1143    let manifest = serde_json::to_vec_pretty(details)?;
1144    fs::write(dataset_dir.join("dataforge-dataset.json"), manifest)?;
1145    Ok(())
1146}
1147
1148fn dataset_matches_filter(dataset: &DatasetSummary, filter: Option<&str>) -> bool {
1149    let Some(filter) = filter.map(str::trim).filter(|filter| !filter.is_empty()) else {
1150        return true;
1151    };
1152    let filter = filter.to_lowercase();
1153    dataset.id.to_lowercase().contains(&filter)
1154        || dataset
1155            .title
1156            .as_deref()
1157            .is_some_and(|title| title.to_lowercase().contains(&filter))
1158}
1159
1160fn dataset_dir_name(dataset_id: &str) -> String {
1161    sanitize_component(dataset_id).replace('/', "_")
1162}
1163
1164fn safe_relative_path(path: &str) -> PathBuf {
1165    let mut clean = PathBuf::new();
1166    for part in path.split('/') {
1167        let part = sanitize_component(part);
1168        if part.is_empty() || part == "." || part == ".." {
1169            continue;
1170        }
1171        clean.push(part);
1172    }
1173    if clean.as_os_str().is_empty() {
1174        PathBuf::from("dataset.bin")
1175    } else {
1176        clean
1177    }
1178}
1179
1180fn sanitize_component(value: &str) -> String {
1181    value
1182        .trim()
1183        .chars()
1184        .map(|ch| match ch {
1185            'a'..='z' | 'A'..='Z' | '0'..='9' | '-' | '_' | '.' | '/' => ch,
1186            _ => '_',
1187        })
1188        .collect::<String>()
1189        .trim_matches('_')
1190        .to_string()
1191}
1192
1193fn partial_path(path: &Path) -> PathBuf {
1194    let mut partial = path.to_path_buf();
1195    let extension = path.extension().and_then(|ext| ext.to_str()).map_or_else(
1196        || "dataforge-part".to_string(),
1197        |ext| format!("{ext}.dataforge-part"),
1198    );
1199    partial.set_extension(extension);
1200    partial
1201}
1202
1203fn hf_file_url(dataset_id: &str, path: &str) -> String {
1204    format!(
1205        "{HF_BASE}/datasets/{}/resolve/main/{}",
1206        encode_repo_id(dataset_id),
1207        encode_path(path)
1208    )
1209}
1210
1211fn encode_repo_id(value: &str) -> String {
1212    encode_path(value)
1213}
1214
1215fn encode_path(value: &str) -> String {
1216    value
1217        .split('/')
1218        .map(urlencoding::encode)
1219        .collect::<Vec<_>>()
1220        .join("/")
1221}
1222
1223fn normalized_slug(value: &str) -> String {
1224    value
1225        .trim()
1226        .chars()
1227        .filter(|ch| !matches!(ch, '-' | '_' | ' ' | '\t' | '\n' | '\r'))
1228        .flat_map(char::to_lowercase)
1229        .collect()
1230}
1231
1232fn strip_any_prefix_ignore_ascii_case<'a>(value: &'a str, prefixes: &[&str]) -> Option<&'a str> {
1233    prefixes.iter().find_map(|prefix| {
1234        value
1235            .get(..prefix.len())
1236            .is_some_and(|head| head.eq_ignore_ascii_case(prefix))
1237            .then(|| &value[prefix.len()..])
1238    })
1239}
1240
1241#[cfg(test)]
1242mod tests {
1243    use super::*;
1244
1245    #[test]
1246    fn parses_hugging_face_dataset_urls() {
1247        assert_eq!(
1248            parse_dataset_target_for_provider(
1249                "https://huggingface.co/datasets/openai/gsm8k",
1250                DatasetProvider::Hf,
1251            )
1252            .unwrap(),
1253            DatasetTarget::Dataset("openai/gsm8k".to_string())
1254        );
1255        assert_eq!(
1256            parse_dataset_target_for_provider(
1257                "https://huggingface.co/datasets/squad",
1258                DatasetProvider::Hf,
1259            )
1260            .unwrap(),
1261            DatasetTarget::Dataset("squad".to_string())
1262        );
1263    }
1264
1265    #[test]
1266    fn parses_hugging_face_collections_and_prefixes() {
1267        assert_eq!(
1268            parse_dataset_target_for_provider("openai", DatasetProvider::Hf).unwrap(),
1269            DatasetTarget::Namespace("openai".to_string())
1270        );
1271        assert_eq!(
1272            parse_dataset_target_for_provider("search:code data", DatasetProvider::Hf).unwrap(),
1273            DatasetTarget::Search("code data".to_string())
1274        );
1275        assert_eq!(
1276            parse_dataset_target_for_provider("dataset:squad", DatasetProvider::Hf).unwrap(),
1277            DatasetTarget::Dataset("squad".to_string())
1278        );
1279    }
1280
1281    #[test]
1282    fn parses_zenodo_records() {
1283        assert_eq!(
1284            parse_dataset_target_for_provider(
1285                "https://zenodo.org/records/12345",
1286                DatasetProvider::Zenodo
1287            )
1288            .unwrap(),
1289            DatasetTarget::Dataset("12345".to_string())
1290        );
1291        assert_eq!(
1292            parse_dataset_target_for_provider("climate data", DatasetProvider::Zenodo).unwrap(),
1293            DatasetTarget::Search("climate data".to_string())
1294        );
1295    }
1296
1297    #[test]
1298    fn parses_provider_aliases() {
1299        assert_eq!(
1300            DatasetProvider::from_slug("hugging-face"),
1301            Some(DatasetProvider::Hf)
1302        );
1303        assert_eq!(
1304            DatasetProvider::from_slug("zen"),
1305            Some(DatasetProvider::Zenodo)
1306        );
1307    }
1308
1309    #[test]
1310    fn safe_paths_drop_traversal() {
1311        assert_eq!(safe_relative_path("../a/./b.csv"), PathBuf::from("a/b.csv"));
1312        assert_eq!(
1313            dataset_dir_name("owner/name with spaces"),
1314            "owner_name_with_spaces"
1315        );
1316    }
1317}