1use std::io::{IsTerminal, Read, Write};
9use std::path::{Path, PathBuf};
10use std::sync::atomic::{AtomicU64, Ordering};
11
12const STAMP: &str = ".verified";
14
15#[derive(Debug)]
17pub struct ModelFile {
18 pub name: &'static str,
19 pub size: u64,
20 pub sha256: &'static str,
21}
22
23#[derive(Debug, Clone, Copy, PartialEq, Eq)]
26pub enum Architecture {
27 JinaBert,
29 NomicBert,
31}
32
33#[derive(Debug)]
35pub struct LocalModel {
36 pub id: &'static str,
38 pub revision: &'static str,
39 pub files: &'static [ModelFile],
40 pub max_tokens: usize,
42 pub architecture: Architecture,
43}
44
45pub const CODERANKEMBED: LocalModel = LocalModel {
46 id: "nomic-ai/CodeRankEmbed",
47 revision: "3c4b60807d71f79b43f3c4363786d9493691f8b1",
48 files: &[
49 ModelFile {
50 name: "config.json",
51 size: 1_525,
52 sha256: "5ff856a41d0f53ef2d74520627d464bd75c2efd8f26f381bd528654895c29b6c",
53 },
54 ModelFile {
55 name: "tokenizer.json",
56 size: 711_649,
57 sha256: "91f1def9b9391fdabe028cd3f3fcc4efd34e5d1f08c3bf2de513ebb5911a1854",
58 },
59 ModelFile {
60 name: "model.safetensors",
61 size: 546_938_168,
62 sha256: "827529bcd58aef0d9082e66eeff7e7d53a02f62bd005f841a26b3d3e2fb17ebe",
63 },
64 ],
65 max_tokens: 1024,
66 architecture: Architecture::NomicBert,
67};
68
69pub const JINA_V2_BASE_CODE: LocalModel = LocalModel {
70 id: "jinaai/jina-embeddings-v2-base-code",
71 revision: "516f4baf13dec4ddddda8631e019b5737c8bc250",
72 files: &[
73 ModelFile {
74 name: "config.json",
75 size: 1_216,
76 sha256: "e426aa684c7f9a95c5f020aa855faf93a24f065f5fad0c9e17b124670cabdea6",
77 },
78 ModelFile {
79 name: "tokenizer.json",
80 size: 2_561_316,
81 sha256: "b01c78a902aa4facb2f47f95449f48e2f7bbfea5d2472ee2f6ce92323c6f86e5",
82 },
83 ModelFile {
84 name: "model.safetensors",
85 size: 321_767_312,
86 sha256: "8b53bfd4ae2cd586004a6ca4a16551b630a2a1b1d655ff1ee9be1286a1781c5b",
87 },
88 ],
89 max_tokens: 1024,
90 architecture: Architecture::JinaBert,
91};
92
93impl LocalModel {
94 pub fn dir(&self, cache_root: &Path) -> PathBuf {
96 cache_root
97 .join("models")
98 .join(self.id.replace('/', "--"))
99 .join(self.revision)
100 }
101
102 pub fn size(&self) -> u64 {
104 self.files.iter().map(|f| f.size).sum()
105 }
106
107 pub fn is_downloaded(&self, dir: &Path) -> bool {
114 if self.is_stamped(dir) {
115 return true;
116 }
117 let verified = self
118 .files
119 .iter()
120 .all(|f| file_matches(&dir.join(f.name), f));
121 if verified {
122 let _ = write_atomically(&dir.join(STAMP), self.stamp().as_bytes());
123 }
124 verified
125 }
126
127 pub fn is_stamped(&self, dir: &Path) -> bool {
131 self.files
132 .iter()
133 .all(|f| has_size(&dir.join(f.name), f.size))
134 && std::fs::read_to_string(dir.join(STAMP)).is_ok_and(|stamp| stamp == self.stamp())
135 }
136
137 fn stamp(&self) -> String {
140 let mut stamp = format!("{} {}\n", self.id, self.revision);
141 for file in self.files {
142 stamp.push_str(&format!("{} {}\n", file.sha256, file.name));
143 }
144 stamp
145 }
146
147 pub fn download(&self, dir: &Path, agent: &ureq::Agent, quiet: bool) -> Result<(), String> {
149 std::fs::create_dir_all(dir).map_err(|e| format!("{}: {e}", dir.display()))?;
150 let endpoint = std::env::var("HF_ENDPOINT")
151 .ok()
152 .filter(|e| !e.is_empty())
153 .unwrap_or_else(|| "https://huggingface.co".to_string());
154 let endpoint = endpoint.trim_end_matches('/');
155 if !quiet {
156 eprintln!(
157 "Downloading {} ({}) from {endpoint} into {}",
158 self.id,
159 megabytes(self.size()),
160 dir.display()
161 );
162 }
163 for file in self.files {
164 let target = dir.join(file.name);
165 if file_matches(&target, file) {
166 continue;
167 }
168 let url = format!(
169 "{endpoint}/{}/resolve/{}/{}",
170 self.id, self.revision, file.name
171 );
172 fetch(agent, &url, file, &target, quiet)?;
173 }
174 write_atomically(&dir.join(STAMP), self.stamp().as_bytes())
175 }
176}
177
178fn has_size(path: &Path, size: u64) -> bool {
179 std::fs::metadata(path).is_ok_and(|m| m.len() == size)
180}
181
182fn file_matches(path: &Path, file: &ModelFile) -> bool {
184 has_size(path, file.size) && sha256_of(path).is_ok_and(|digest| digest == file.sha256)
185}
186
187fn sha256_of(path: &Path) -> std::io::Result<String> {
188 use sha2::Digest;
189 let mut input = std::fs::File::open(path)?;
190 let mut hasher = sha2::Sha256::new();
191 let mut buffer = vec![0u8; 1 << 20];
192 loop {
193 let n = input.read(&mut buffer)?;
194 if n == 0 {
195 break;
196 }
197 hasher.update(&buffer[..n]);
198 }
199 Ok(hex(&hasher.finalize()))
200}
201
202fn hex(bytes: &[u8]) -> String {
203 bytes.iter().map(|b| format!("{b:02x}")).collect()
204}
205
206fn scratch_path(target: &Path) -> PathBuf {
210 static NEXT: AtomicU64 = AtomicU64::new(0);
211 let nanos = std::time::SystemTime::now()
212 .duration_since(std::time::UNIX_EPOCH)
213 .map_or(0, |d| d.subsec_nanos());
214 let name = target
215 .file_name()
216 .map(|n| n.to_string_lossy().into_owned())
217 .unwrap_or_default();
218 target.with_file_name(format!(
219 "{name}.{}-{nanos}-{}.partial",
220 std::process::id(),
221 NEXT.fetch_add(1, Ordering::Relaxed)
222 ))
223}
224
225struct Scratch {
227 path: PathBuf,
228 kept: bool,
229}
230
231impl Drop for Scratch {
232 fn drop(&mut self) {
233 if !self.kept {
234 let _ = std::fs::remove_file(&self.path);
235 }
236 }
237}
238
239fn write_atomically(target: &Path, bytes: &[u8]) -> Result<(), String> {
242 let mut scratch = Scratch {
243 path: scratch_path(target),
244 kept: false,
245 };
246 std::fs::write(&scratch.path, bytes).map_err(|e| format!("{}: {e}", scratch.path.display()))?;
247 std::fs::rename(&scratch.path, target).map_err(|e| format!("{}: {e}", target.display()))?;
248 scratch.kept = true;
249 Ok(())
250}
251
252fn megabytes(bytes: u64) -> String {
253 match bytes {
254 0..1_000_000 => format!("{:.0} KB", bytes as f64 / 1_000.0),
255 _ => format!("{:.0} MB", bytes as f64 / 1_000_000.0),
256 }
257}
258
259fn fetch(
263 agent: &ureq::Agent,
264 url: &str,
265 file: &ModelFile,
266 target: &Path,
267 quiet: bool,
268) -> Result<(), String> {
269 use sha2::Digest;
270 let mut response = agent.get(url).call().map_err(|e| format!("{url}: {e}"))?;
271 let status = response.status().as_u16();
272 if status != 200 {
273 return Err(format!("{url}: HTTP {status}"));
274 }
275 let mut scratch = Scratch {
276 path: scratch_path(target),
277 kept: false,
278 };
279 let partial = scratch.path.clone();
280 let mut out = std::fs::OpenOptions::new()
281 .write(true)
282 .create_new(true)
283 .open(&partial)
284 .map_err(|e| format!("{}: {e}", partial.display()))?;
285 let mut reader = response
286 .body_mut()
287 .with_config()
288 .limit(file.size + 1)
289 .reader();
290 let mut hasher = sha2::Sha256::new();
291 let mut buffer = vec![0u8; 1 << 20];
292 let (mut done, mut shown) = (0u64, 0u64);
293 let live = !quiet && std::io::stderr().is_terminal();
294 loop {
295 let n = reader
296 .read(&mut buffer)
297 .map_err(|e| format!("{url}: {e}"))?;
298 if n == 0 {
299 break;
300 }
301 hasher.update(&buffer[..n]);
302 out.write_all(&buffer[..n])
303 .map_err(|e| format!("{}: {e}", partial.display()))?;
304 done += n as u64;
305 if live && (done - shown > file.size / 100 || done == file.size) {
306 shown = done;
307 eprint!("\r {} {:>3}%", file.name, done * 100 / file.size.max(1));
308 }
309 }
310 if live {
311 eprintln!();
312 }
313 out.sync_all()
314 .map_err(|e| format!("{}: {e}", partial.display()))?;
315 drop(out);
316 let digest = hex(&hasher.finalize());
317 if done != file.size || digest != file.sha256 {
318 return Err(format!(
319 "{url}: got {done} bytes with SHA-256 {digest}, expected {} bytes with {}",
320 file.size, file.sha256
321 ));
322 }
323 match std::fs::rename(&partial, target) {
324 Ok(()) => scratch.kept = true,
325 Err(_) if file_matches(target, file) => {}
328 Err(e) => return Err(format!("{}: {e}", target.display())),
329 }
330 if !quiet {
331 eprintln!(" {} {} verified", file.name, megabytes(file.size));
332 }
333 Ok(())
334}
335
336#[cfg(test)]
337mod tests {
338 use super::*;
339
340 #[test]
341 fn a_model_lives_in_a_folder_of_its_revision() {
342 let model = &CODERANKEMBED;
343 assert_eq!(model.size(), 1_525 + 711_649 + 546_938_168);
344 let dir = model.dir(Path::new("/cache"));
345 assert_eq!(
346 dir,
347 Path::new("/cache/models/nomic-ai--CodeRankEmbed")
348 .join("3c4b60807d71f79b43f3c4363786d9493691f8b1")
349 );
350 assert!(!model.is_downloaded(&dir));
351 }
352
353 static HELLO: LocalModel = LocalModel {
354 id: "test/hello",
355 revision: "r1",
356 files: &[ModelFile {
357 name: "hello.txt",
358 size: 5,
359 sha256: "2cf24dba5fb0a30e26e83b2ac5b9e29e1b161e5c1fa7425e73043362938b9824",
360 }],
361 max_tokens: 8,
362 architecture: Architecture::JinaBert,
363 };
364
365 use crate::embed::test_dir;
366
367 #[test]
368 fn a_file_counts_once_its_checksum_matched_not_its_size() {
369 let dir = test_dir("models-verify");
370 std::fs::write(dir.join("hello.txt"), "HELLO").unwrap();
371 assert!(!HELLO.is_stamped(&dir), "the right size, no stamp");
372 assert!(
373 !HELLO.is_downloaded(&dir),
374 "the right size with the wrong bytes"
375 );
376 assert!(!dir.join(STAMP).exists());
377
378 std::fs::write(dir.join("hello.txt"), "hello").unwrap();
379 assert!(!HELLO.is_stamped(&dir), "not hashed yet");
380 assert!(!dir.join(STAMP).exists(), "and nothing written");
381 assert!(HELLO.is_downloaded(&dir), "hashed once");
382 assert!(HELLO.is_stamped(&dir));
383 assert_eq!(
384 std::fs::read_to_string(dir.join(STAMP)).unwrap(),
385 HELLO.stamp(),
386 "and stamped, so later runs read the stamp instead of hashing"
387 );
388 std::fs::remove_dir_all(&dir).unwrap();
389 }
390
391 #[test]
392 fn scratch_files_are_unique_and_removed_unless_kept() {
393 let dir = test_dir("models-scratch");
394 let target = dir.join("model.safetensors");
395 let (a, b) = (scratch_path(&target), scratch_path(&target));
396 assert_ne!(a, b, "two downloads never share a file");
397 {
398 let scratch = Scratch {
399 path: a.clone(),
400 kept: false,
401 };
402 std::fs::write(&scratch.path, "part").unwrap();
403 }
404 assert!(!a.exists(), "an abandoned scratch file is removed");
405 write_atomically(&target, b"whole").unwrap();
406 assert_eq!(std::fs::read(&target).unwrap(), b"whole");
407 assert_eq!(std::fs::read_dir(&dir).unwrap().count(), 1, "no leftovers");
408 std::fs::remove_dir_all(&dir).unwrap();
409 }
410}