Skip to main content

cherenkov/model/index/hub/
headers.rs

1//! Read each shard's prefix and header together, with bounded parallel requests.
2
3use super::*;
4use std::sync::{
5    atomic::{AtomicBool, AtomicUsize, Ordering},
6    mpsc::{SyncSender, sync_channel},
7};
8
9const CONCURRENCY: usize = 8;
10type Shard = (RemoteObject, Inventory);
11
12struct Request<'a> {
13    id: ObjectId,
14    file: &'a str,
15    url: reqwest::Url,
16}
17
18impl Hub {
19    /// Workers only fetch data. The caller emits events and restores file order.
20    pub(super) fn headers(
21        &self,
22        root: &[&str],
23        files: &[String],
24        events: &mut dyn FnMut(ModelEvent),
25    ) -> Result<Vec<Shard>> {
26        let requests = files
27            .iter()
28            .enumerate()
29            .map(|(id, file)| {
30                ensure!(
31                    crate::storage::safe_component(file),
32                    "unsupported shard path {file:?}"
33                );
34
35                let mut components = root.to_vec();
36
37                components.push(file);
38
39                Ok(Request {
40                    id: ObjectId(id),
41                    file,
42                    url: self.url(&components)?,
43                })
44            })
45            .collect::<Result<Vec<_>>>()?;
46        let next = AtomicUsize::new(0);
47        let failed = AtomicBool::new(false);
48        let workers = CONCURRENCY.min(requests.len());
49
50        events(ModelEvent::Headers {
51            completed: 0,
52            total: files.len(),
53            file: None,
54        });
55
56        std::thread::scope(|scope| {
57            let (sender, receiver) = sync_channel(workers);
58
59            for _ in 0..workers {
60                let sender = sender.clone();
61                let requests = &requests;
62                let next = &next;
63                let failed = &failed;
64
65                std::thread::Builder::new()
66                    .name("hf-header".into())
67                    .spawn_scoped(scope, move || {
68                        self.fetch_headers(requests, next, failed, sender)
69                    })
70                    .context("starting HF header worker")?;
71            }
72
73            drop(sender);
74
75            let mut shards = Vec::with_capacity(files.len());
76            let mut error = None;
77
78            for (id, result) in receiver {
79                match result {
80                    Ok(shard) => {
81                        shards.push((id, shard));
82                        events(ModelEvent::Headers {
83                            completed: shards.len(),
84                            total: files.len(),
85                            file: Some(files[id].clone()),
86                        });
87                    }
88                    Err(failure) => {
89                        error.get_or_insert(failure);
90                    }
91                }
92            }
93
94            if let Some(error) = error {
95                return Err(error);
96            }
97
98            shards.sort_unstable_by_key(|(id, _)| *id);
99
100            Ok(shards.into_iter().map(|(_, shard)| shard).collect())
101        })
102    }
103
104    /// Stop claiming shards after an error; finish and drain shards already claimed.
105    fn fetch_headers(
106        &self,
107        requests: &[Request<'_>],
108        next: &AtomicUsize,
109        failed: &AtomicBool,
110        sender: SyncSender<(usize, Result<Shard>)>,
111    ) {
112        while !failed.load(Ordering::Relaxed) {
113            let id = next.fetch_add(1, Ordering::Relaxed);
114            let Some(request) = requests.get(id) else {
115                break;
116            };
117            let result = self
118                .header(request)
119                .with_context(|| format!("reading shard {}", request.file));
120
121            if result.is_err() {
122                failed.store(true, Ordering::Relaxed);
123            }
124
125            if sender.send((id, result)).is_err() {
126                break;
127            }
128        }
129    }
130
131    fn header(&self, request: &Request<'_>) -> Result<Shard> {
132        let (prefix, size) = self.range(request.url.clone(), 0, 8)?;
133        let object = RemoteObject {
134            id: request.id,
135            url: request.url.clone(),
136            prefix,
137            size,
138        };
139        let reader = RemoteReader {
140            hub: self,
141            objects: std::slice::from_ref(&object),
142        };
143        let shard = read_safetensors(&reader, object.id)?;
144
145        Ok((object, shard))
146    }
147}