1use 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
43pub type Result<T> = std::result::Result<T, DataForgeError>;
45
46#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
48pub enum DatasetProvider {
49 Hf,
51 Zenodo,
53}
54
55impl DatasetProvider {
56 pub const ALL: [Self; 2] = [Self::Hf, Self::Zenodo];
58
59 pub fn all() -> &'static [Self] {
61 &Self::ALL
62 }
63
64 pub fn label(self) -> &'static str {
66 match self {
67 Self::Hf => "Hugging Face Datasets",
68 Self::Zenodo => "Zenodo",
69 }
70 }
71
72 pub fn short_label(self) -> &'static str {
74 match self {
75 Self::Hf => "HF",
76 Self::Zenodo => "ZEN",
77 }
78 }
79
80 pub fn slug(self) -> &'static str {
82 match self {
83 Self::Hf => "hf",
84 Self::Zenodo => "zenodo",
85 }
86 }
87
88 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#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
107pub enum DatasetTarget {
108 Namespace(String),
110 Dataset(String),
112 Search(String),
114 Top,
116 Latest,
118}
119
120impl DatasetTarget {
121 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 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 pub fn is_dataset(&self) -> bool {
142 matches!(self, Self::Dataset(_))
143 }
144
145 pub fn is_collection(&self) -> bool {
147 matches!(
148 self,
149 Self::Namespace(_) | Self::Search(_) | Self::Top | Self::Latest
150 )
151 }
152}
153
154#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
156pub struct DatasetSummary {
157 pub provider: DatasetProvider,
159 pub id: String,
161 pub title: Option<String>,
163 pub downloads: Option<u64>,
165 pub likes: Option<u64>,
167 pub last_modified: Option<String>,
169 pub description: Option<String>,
171 pub tags: Vec<String>,
173 pub file_count: Option<usize>,
175 pub size: Option<u64>,
177 pub url: Option<String>,
179}
180
181#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
183pub struct DatasetFile {
184 pub path: String,
186 pub download_url: String,
188 pub size: Option<u64>,
190 pub checksum: Option<String>,
192}
193
194#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
196pub struct DatasetDetails {
197 pub summary: DatasetSummary,
199 pub files: Vec<DatasetFile>,
201}
202
203#[derive(Debug, Clone)]
205pub struct ArchiveOptions {
206 pub output: PathBuf,
208 pub concurrency: usize,
210 pub skip_existing: bool,
212 pub filter: Option<String>,
214}
215
216impl ArchiveOptions {
217 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#[derive(Debug, Clone, PartialEq, Eq)]
239pub struct ArchivedFile {
240 pub dataset_id: String,
242 pub source_path: String,
244 pub local_path: PathBuf,
246 pub bytes_written: u64,
248 pub skipped: bool,
250}
251
252#[derive(Debug, Clone, PartialEq, Eq)]
254pub struct ArchiveFailure {
255 pub dataset_id: String,
257 pub file: Option<String>,
259 pub message: String,
261}
262
263#[derive(Debug, Clone, Default, PartialEq, Eq)]
265pub struct ArchiveSummary {
266 pub discovered: usize,
268 pub selected: usize,
270 pub archived: usize,
272 pub skipped: usize,
274 pub files_archived: usize,
276 pub bytes_written: u64,
278 pub files: Vec<ArchivedFile>,
280 pub failed: Vec<ArchiveFailure>,
282}
283
284impl ArchiveSummary {
285 pub fn is_success(&self) -> bool {
287 self.failed.is_empty()
288 }
289
290 pub fn attempted(&self) -> usize {
292 self.archived + self.skipped + self.failed_dataset_count()
293 }
294
295 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#[derive(Debug, thiserror::Error)]
307pub enum DataForgeError {
308 #[error("invalid dataset target '{0}'")]
310 InvalidTarget(String),
311
312 #[error("unsupported dataset provider target: {0}")]
314 UnsupportedTarget(String),
315
316 #[error("{provider} API returned {status}: {body}")]
318 ProviderApi {
319 provider: &'static str,
321 status: u16,
323 body: String,
325 },
326
327 #[error("{provider} dataset '{target}' was not found")]
329 NotFound {
330 provider: &'static str,
332 target: String,
334 },
335
336 #[error("download failed for '{dataset}' file '{file}': {message}")]
338 DownloadFailed {
339 dataset: String,
341 file: String,
343 message: String,
345 },
346
347 #[error("output path '{}' is not a directory", .0.display())]
349 OutputNotDirectory(PathBuf),
350
351 #[error("I/O error: {0}")]
353 Io(#[from] io::Error),
354
355 #[error("request error: {0}")]
357 Request(#[from] reqwest::Error),
358
359 #[error("JSON error: {0}")]
361 Json(#[from] serde_json::Error),
362}
363
364#[derive(Clone)]
366pub struct DataForge {
367 provider: DatasetProvider,
368 client: Client,
369 token: Option<String>,
370 page_size: usize,
371}
372
373impl DataForge {
374 pub fn new(provider: DatasetProvider) -> Result<Self> {
376 Self::with_user_agent(provider, DEFAULT_USER_AGENT)
377 }
378
379 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 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 pub fn with_page_size(mut self, page_size: usize) -> Self {
401 self.page_size = page_size.max(1);
402 self
403 }
404
405 pub fn provider(&self) -> DatasetProvider {
407 self.provider
408 }
409
410 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 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 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 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
536pub fn parse_dataset_target(input: &str) -> Result<DatasetTarget> {
538 parse_dataset_target_for_provider(input, DatasetProvider::Hf)
539}
540
541pub 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}