1use core::sync::atomic::{AtomicBool, AtomicUsize, Ordering};
4
5use crate::braceexpand::braceexpand;
6use crate::error::{Error, Result};
7use crate::prelude::*;
8
9pub fn base_plus_ext(path: &str) -> Option<(&str, &str)> {
23 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
34pub fn is_shard_metadata(name: &str) -> bool {
39 if !name.contains('/') && name.starts_with("__") && name.ends_with("__") {
40 return true;
41 }
42 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#[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 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#[cfg(not(feature = "std"))]
92pub fn envsubst(text: &str) -> Result<String> {
93 Ok(String::from(text))
94}
95
96pub fn expand_urls(spec: &str) -> Result<Vec<String>> {
108 let mut result = Vec::new();
109 for part in spec.split("::") {
110 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
125pub 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
134fn 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
143pub 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#[derive(Debug, Clone, Copy, PartialEq, Eq)]
158pub struct WorkerInfo {
159 pub rank: usize,
161 pub world_size: usize,
163 pub worker: usize,
165 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 pub fn seed(&self) -> u64 {
178 (self.rank * 1000 + self.worker) as u64
179 }
180}
181
182pub 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
218pub 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
234pub fn set_enforce_security(on: bool) {
236 SECURE_INIT.store(true, Ordering::Relaxed);
237 SECURE.store(on, Ordering::Relaxed);
238}
239
240pub 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
250pub fn next_unique() -> usize {
252 COUNTER.fetch_add(1, Ordering::Relaxed)
253}
254
255pub 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}