cherenkov/model/index/hub/
headers.rs1use 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 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 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}