Skip to main content

webdataset_core/
utils.rs

1//! Shared helpers: filename splitting, worker identity, seeds, and secure mode.
2
3use core::sync::atomic::{AtomicBool, AtomicUsize, Ordering};
4
5use crate::braceexpand::braceexpand;
6use crate::error::{Error, Result};
7use crate::prelude::*;
8
9/// Split a path into its basename and its full extension.
10///
11/// The basename is everything up to the first dot after the last slash; the
12/// extension is everything after it. Files in a WebDataset share a basename
13/// exactly when they belong to the same sample.
14///
15/// ```
16/// use webdataset_core::utils::base_plus_ext;
17///
18/// assert_eq!(base_plus_ext("a/b/c.png"), Some(("a/b/c", "png")));
19/// assert_eq!(base_plus_ext("a/b/c.seg.png"), Some(("a/b/c", "seg.png")));
20/// assert_eq!(base_plus_ext("noextension"), None);
21/// ```
22pub fn base_plus_ext(path: &str) -> Option<(&str, &str)> {
23    // Everything after the last slash is the file name; the basename runs up to
24    // its first dot, and the extension is the rest. A name that starts with a
25    // dot has no basename, so it belongs to no sample.
26    let start = path.rfind('/').map(|at| at + 1).unwrap_or(0);
27    let dot = start + path[start..].find('.')?;
28    if dot == start {
29        return None;
30    }
31    Some((&path[..dot], &path[dot + 1..]))
32}
33
34/// Whether a tar member name is shard-level metadata that should be skipped.
35///
36/// These are top-level entries wrapped in double underscores, such as
37/// `__index__` or the `__dup__/` directory.
38pub fn is_shard_metadata(name: &str) -> bool {
39    if !name.contains('/') && name.starts_with("__") && name.ends_with("__") {
40        return true;
41    }
42    // The Python default `skip_meta` regexp: `__[^/]*__($|/)`.
43    let Some(rest) = name.strip_prefix("__") else {
44        return false;
45    };
46    let Some(end) = rest.find("__") else {
47        return false;
48    };
49    if rest[..end].contains('/') {
50        return false;
51    }
52    let after = &rest[end + 2..];
53    after.is_empty() || after.starts_with('/')
54}
55
56/// Substitute `${NAME}` with the environment variable `WDS_NAME`.
57///
58/// Fails when a referenced variable is not set, matching the reference
59/// implementation's assertion.
60#[cfg(feature = "std")]
61pub fn envsubst(text: &str) -> Result<String> {
62    let mut out = String::with_capacity(text.len());
63    let mut rest = text;
64    while let Some(at) = rest.find("${") {
65        out.push_str(&rest[..at]);
66        let body = &rest[at + 2..];
67        let Some(end) = body.find('}') else {
68            // An unterminated `${` is literal text, as in the shell.
69            out.push_str(&rest[at..]);
70            return Ok(out);
71        };
72        let name = &body[..end];
73        if name.is_empty() || !name.chars().all(|c| c.is_alphanumeric() || c == '_') {
74            out.push_str(&rest[at..at + 2 + end + 1]);
75            rest = &body[end + 1..];
76            continue;
77        }
78        let variable = format!("WDS_{name}");
79        match std::env::var(&variable) {
80            Ok(value) => out.push_str(&value),
81            Err(_) => return Err(Error::value(format!("missing environment variable {variable}"))),
82        }
83        rest = &body[end + 1..];
84    }
85    out.push_str(rest);
86    Ok(out)
87}
88
89/// Without `std` there is no environment to substitute from, so the text is
90/// returned unchanged.
91#[cfg(not(feature = "std"))]
92pub fn envsubst(text: &str) -> Result<String> {
93    Ok(String::from(text))
94}
95
96/// Expand a URL specification into a concrete list of shard URLs.
97///
98/// The specification is split on `::`, each part has `${VAR}` substituted from
99/// `WDS_*` environment variables, and brace expressions are expanded.
100///
101/// ```
102/// use webdataset_core::utils::expand_urls;
103///
104/// let urls = expand_urls("a-{0..1}.tar::b.tar").unwrap();
105/// assert_eq!(urls, ["a-0.tar", "a-1.tar", "b.tar"]);
106/// ```
107pub fn expand_urls(spec: &str) -> Result<Vec<String>> {
108    let mut result = Vec::new();
109    for part in spec.split("::") {
110        // Substitution can itself introduce `${...}`, so iterate to a fixpoint
111        // with the same bound the reference implementation uses.
112        let mut url = part.to_string();
113        for _ in 0..10 {
114            let next = envsubst(&url)?;
115            if next == url {
116                break;
117            }
118            url = next;
119        }
120        result.extend(braceexpand(&url)?);
121    }
122    Ok(result)
123}
124
125/// Combine several values into a 31-bit seed, mirroring `utils.make_seed`.
126pub fn make_seed(parts: &[u64]) -> u64 {
127    let mut seed: u64 = 0;
128    for part in parts {
129        seed = seed.wrapping_mul(31).wrapping_add(mix(*part)) & 0x7FFF_FFFF;
130    }
131    seed
132}
133
134/// A cheap avalanche used by [`make_seed`] so that adjacent inputs diverge.
135fn mix(mut x: u64) -> u64 {
136    x ^= x >> 33;
137    x = x.wrapping_mul(0xff51_afd7_ed55_8ccd);
138    x ^= x >> 33;
139    x = x.wrapping_mul(0xc4ce_b9fe_1a85_ec53);
140    x ^ (x >> 33)
141}
142
143/// Hash a string into a seed component.
144pub fn seed_from_str(text: &str) -> u64 {
145    let mut hash: u64 = 0xcbf2_9ce4_8422_2325;
146    for byte in text.as_bytes() {
147        hash ^= *byte as u64;
148        hash = hash.wrapping_mul(0x0000_0100_0000_01b3);
149    }
150    hash
151}
152
153/// Which process and worker this code is running in.
154///
155/// Mirrors `pytorch_worker_info`: distributed rank and world size identify the
156/// node, worker and count identify the loader worker within it.
157#[derive(Debug, Clone, Copy, PartialEq, Eq)]
158pub struct WorkerInfo {
159    /// This process's index among all distributed processes.
160    pub rank: usize,
161    /// The total number of distributed processes.
162    pub world_size: usize,
163    /// This thread's index among the loader workers of this process.
164    pub worker: usize,
165    /// The total number of loader workers in this process.
166    pub num_workers: usize,
167}
168
169impl Default for WorkerInfo {
170    fn default() -> Self {
171        WorkerInfo { rank: 0, world_size: 1, worker: 0, num_workers: 1 }
172    }
173}
174
175impl WorkerInfo {
176    /// A deterministic per-worker seed, as `pytorch_worker_seed` computes it.
177    pub fn seed(&self) -> u64 {
178        (self.rank * 1000 + self.worker) as u64
179    }
180}
181
182/// Report the current worker identity.
183///
184/// The per-thread binding set by [`with_worker`] wins; otherwise
185/// `RANK`/`WORLD_SIZE` and `WORKER`/`NUM_WORKERS` are read from the
186/// environment. See [`crate::workers`] for how the binding is stored on targets
187/// without thread-local storage.
188pub fn worker_info() -> WorkerInfo {
189    let mut info = WorkerInfo::default();
190    #[cfg(feature = "std")]
191    if let (Ok(rank), Ok(world)) = (env_usize("RANK"), env_usize("WORLD_SIZE")) {
192        info.rank = rank;
193        info.world_size = world.max(1);
194    }
195    if let Some((worker, num_workers)) = crate::workers::current() {
196        info.worker = worker;
197        info.num_workers = num_workers.max(1);
198        return info;
199    }
200    #[cfg(feature = "std")]
201    if let (Ok(worker), Ok(num)) = (env_usize("WORKER"), env_usize("NUM_WORKERS")) {
202        info.worker = worker;
203        info.num_workers = num.max(1);
204    }
205    info
206}
207
208pub use crate::workers::{clear_worker, set_thread_id_hook, set_worker, with_worker};
209
210#[cfg(feature = "std")]
211fn env_usize(name: &str) -> core::result::Result<usize, ()> {
212    std::env::var(name).map_err(|_| ())?.parse().map_err(|_| ())
213}
214
215static SECURE: AtomicBool = AtomicBool::new(false);
216static SECURE_INIT: AtomicBool = AtomicBool::new(false);
217
218/// Whether secure mode is on.
219///
220/// Secure mode disables the `pipe:` and `file:` URL schemes, URL rewriting from
221/// the environment, and decoders that execute embedded code. It is enabled by
222/// setting `WDS_SECURE=1` or by calling [`set_enforce_security`].
223pub fn enforce_security() -> bool {
224    if !SECURE_INIT.swap(true, Ordering::Relaxed) {
225        #[cfg(feature = "std")]
226        {
227            let on = std::env::var("WDS_SECURE").map(|v| v != "0" && !v.is_empty()).unwrap_or(false);
228            SECURE.store(on, Ordering::Relaxed);
229        }
230    }
231    SECURE.load(Ordering::Relaxed)
232}
233
234/// Turn secure mode on or off explicitly.
235pub fn set_enforce_security(on: bool) {
236    SECURE_INIT.store(true, Ordering::Relaxed);
237    SECURE.store(on, Ordering::Relaxed);
238}
239
240/// Fail when secure mode forbids `what`.
241pub fn check_security(what: &str) -> Result<()> {
242    if enforce_security() {
243        return Err(Error::Security(what.to_string()));
244    }
245    Ok(())
246}
247
248static COUNTER: AtomicUsize = AtomicUsize::new(0);
249
250/// A process-unique counter, used to keep temporary file names from colliding.
251pub fn next_unique() -> usize {
252    COUNTER.fetch_add(1, Ordering::Relaxed)
253}
254
255/// Format `value` into a `printf`-style integer pattern such as `%06d`.
256///
257/// Shard writers take their output patterns in this form for compatibility
258/// with the Python API, where the pattern is used with the `%` operator.
259///
260/// ```
261/// use webdataset_core::utils::format_shard_pattern;
262///
263/// assert_eq!(format_shard_pattern("out-%06d.tar", 12).unwrap(), "out-000012.tar");
264/// assert_eq!(format_shard_pattern("out-%d.tar", 12).unwrap(), "out-12.tar");
265/// ```
266pub fn format_shard_pattern(pattern: &str, value: usize) -> Result<String> {
267    let mut out = String::with_capacity(pattern.len() + 8);
268    let mut chars = pattern.chars().peekable();
269    let mut substituted = false;
270
271    while let Some(c) = chars.next() {
272        if c != '%' {
273            out.push(c);
274            continue;
275        }
276        if chars.peek() == Some(&'%') {
277            chars.next();
278            out.push('%');
279            continue;
280        }
281        let mut width = String::new();
282        while let Some(&d) = chars.peek() {
283            if d.is_ascii_digit() {
284                width.push(d);
285                chars.next();
286            } else {
287                break;
288            }
289        }
290        match chars.next() {
291            Some('d') | Some('i') => {
292                let zero_padded = width.starts_with('0');
293                let w: usize = width.parse().unwrap_or(0);
294                if zero_padded {
295                    out.push_str(&format!("{value:0>w$}"));
296                } else {
297                    out.push_str(&format!("{value:>w$}"));
298                }
299                substituted = true;
300            }
301            other => {
302                return Err(Error::value(format!(
303                    "shard pattern {pattern:?} has an unsupported conversion %{}",
304                    other.unwrap_or(' ')
305                )));
306            }
307        }
308    }
309
310    if !substituted {
311        return Err(Error::value(format!("shard pattern {pattern:?} has no %d conversion")));
312    }
313    Ok(out)
314}
315
316#[cfg(test)]
317mod tests {
318    use super::*;
319
320    #[test]
321    fn splits_names_the_way_the_format_requires() {
322        assert_eq!(base_plus_ext("a.png"), Some(("a", "png")));
323        assert_eq!(base_plus_ext("dir/a.seg.png"), Some(("dir/a", "seg.png")));
324        assert_eq!(base_plus_ext("a/b.c/d.png"), Some(("a/b.c/d", "png")));
325        assert_eq!(base_plus_ext("plain"), None);
326    }
327
328    #[test]
329    fn recognises_shard_metadata() {
330        assert!(is_shard_metadata("__index__"));
331        assert!(is_shard_metadata("__dup__/a.png"));
332        assert!(!is_shard_metadata("dir/__index__"));
333        assert!(!is_shard_metadata("normal.png"));
334    }
335
336    #[test]
337    fn formats_shard_patterns() {
338        assert_eq!(format_shard_pattern("s-%06d.tar", 7).unwrap(), "s-000007.tar");
339        assert_eq!(format_shard_pattern("s-%d.tar", 7).unwrap(), "s-7.tar");
340        assert_eq!(format_shard_pattern("100%%-%d", 1).unwrap(), "100%-1");
341        assert!(format_shard_pattern("no-conversion.tar", 1).is_err());
342        assert!(format_shard_pattern("%s.tar", 1).is_err());
343    }
344
345    #[test]
346    fn overrides_worker_identity_per_thread() {
347        let info = with_worker(2, 4, worker_info);
348        assert_eq!(info.worker, 2);
349        assert_eq!(info.num_workers, 4);
350        assert_eq!(worker_info().worker, 0, "override is scoped to the closure");
351    }
352
353    #[test]
354    fn derives_stable_distinct_seeds() {
355        assert_eq!(make_seed(&[1, 2, 3]), make_seed(&[1, 2, 3]));
356        assert_ne!(make_seed(&[1, 2, 3]), make_seed(&[1, 2, 4]));
357        assert!(make_seed(&[u64::MAX, 7]) <= 0x7FFF_FFFF);
358    }
359}