Skip to main content

cpd_semantic/embed/
models.rs

1// models.rs — embedding models jscpd can run itself, and their download.
2//
3// Every file is pinned by repository revision, size and SHA-256, so what
4// runs is exactly what was calibrated. Files come from Hugging Face (or the
5// mirror in `HF_ENDPOINT`) into the jscpd cache directory, once, on
6// `--semantic-download`; a scan never downloads anything on its own.
7
8use std::io::{IsTerminal, Read, Write};
9use std::path::{Path, PathBuf};
10use std::sync::atomic::{AtomicU64, Ordering};
11
12/// Written into the model directory once every file's SHA-256 matched.
13const STAMP: &str = ".verified";
14
15/// A file of a model repository.
16#[derive(Debug)]
17pub struct ModelFile {
18    pub name: &'static str,
19    pub size: u64,
20    pub sha256: &'static str,
21}
22
23/// The network a model's weights belong to, which decides the code that
24/// runs it.
25#[derive(Debug, Clone, Copy, PartialEq, Eq)]
26pub enum Architecture {
27    /// JinaBERT v2: ALiBi attention, the mean of the token states.
28    JinaBert,
29    /// NomicBERT: rotary attention, the first token's state.
30    NomicBert,
31}
32
33/// A model the local provider can run.
34#[derive(Debug)]
35pub struct LocalModel {
36    /// Hugging Face repository id.
37    pub id: &'static str,
38    pub revision: &'static str,
39    pub files: &'static [ModelFile],
40    /// Longest input in tokens; a longer function is embedded by its head.
41    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    /// Where the files live: `<cache>/models/<owner>--<name>/<revision>`.
95    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    /// Total download size in bytes.
103    pub fn size(&self) -> u64 {
104        self.files.iter().map(|f| f.size).sum()
105    }
106
107    /// Whether every file is in `dir`, verified. A download stamps the
108    /// directory once every file's SHA-256 matched; files that have their
109    /// pinned sizes but no stamp (copied in by hand, or left by an earlier
110    /// build) are hashed once here, and stamped when they match. A size
111    /// alone proves nothing: an interrupted or overlapping download can
112    /// leave a file of the right size with the wrong bytes.
113    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    /// Whether every file in `dir` has its pinned size and the directory
128    /// carries this model's stamp. Unlike [`Self::is_downloaded`], it never
129    /// hashes a file or writes a stamp, so it is cheap enough for a listing.
130    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    /// What the stamp says: the model, its revision and every checksum, so a
138    /// stamp from another revision does not count.
139    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    /// Download the files missing from `dir`, checking size and checksum.
148    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
182/// Whether the file at `path` is `file`: its pinned size and SHA-256.
183fn 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
206/// A temporary file next to `target` that is this process's alone: two
207/// jscpd processes downloading into one cache (parallel CI jobs, containers
208/// sharing a volume) never write into each other's file.
209fn 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
225/// A scratch file, removed when dropped unless it was moved into place.
226struct 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
239/// Write `bytes` to `target` through a scratch file, so a reader sees the
240/// old content or the new, never a part.
241fn 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
259/// Stream `url` into `target`, through a scratch file of this process that
260/// becomes the target only once its size and SHA-256 match, and is removed
261/// on any error.
262fn 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        // Another jscpd put the same file there first and still holds it
326        // open (Windows refuses to replace an open file): theirs will do.
327        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}