Skip to main content

inillucent_cli/
setup.rs

1//! `inillucent setup-embeddings`: install the embedder, on any of three
2//! platforms, in one command.
3//!
4//! Invariant: **a component this reports as installed is verified, complete and
5//! reachable without an environment variable.** Every archive and every weights
6//! file is checked against a digest pinned in this file before anything is
7//! written where the engine looks, the manifest is sealed from the files that
8//! actually landed rather than from what was expected, and the directories are
9//! the ones `inillucent_core::install` computes - so `embed(TEXT)` works
10//! immediately afterwards on a shell that has never exported anything.
11//!
12//! ## What is pinned, and why by digest rather than only by version
13//!
14//! A version pins what was asked for. A digest pins what arrived. The two
15//! differ whenever a release asset is replaced, a content network serves a
16//! truncated body, or something between here and the origin rewrites it - and
17//! the failure mode of the first two is a shared library that loads and
18//! misbehaves rather than one that refuses. So the digests are here, in the
19//! source, and a mismatch deletes what it fetched and names both digests.
20//!
21//! A version this build has no digest for is still installable, because pinning
22//! every future release is not possible and refusing them would make
23//! `--onnxruntime-version` useless. It is fetched, it is reported as
24//! unverified in the output *and* in the install state, and the state file is
25//! where somebody later finds out that this one was taken on trust.
26//!
27//! ## Why 1.22.0
28//!
29//! It is the version this repository's embedding numbers were taken on, and it
30//! is the last release Microsoft publishes a `universal2` macOS archive for.
31//! After it, macOS is two archives and an installer that picks the wrong one is
32//! a support question rather than an error.
33
34use std::io::{IsTerminal, Write};
35use std::path::Path;
36
37use inillucent_core::install::{self, InstalledModel, InstalledRuntime};
38use inillucent_core::model::ModelManifest;
39use inillucent_core::residency::Residency;
40use inillucent_remote::archive::{self, Member};
41use inillucent_remote::http::{self, Progress};
42
43use crate::command::{Arguments, Context, Failed, Outcome};
44use crate::json::{self, Json};
45
46/// The ONNX Runtime version installed when nothing else is asked for.
47pub const DEFAULT_RUNTIME: &str = "1.22.0";
48
49/// Where Microsoft publishes the runtime.
50const RUNTIME_BASE: &str = "https://github.com/microsoft/onnxruntime/releases/download";
51
52/// Where the weights come from.
53const MODEL_BASE: &str = "https://huggingface.co/nomic-ai/nomic-embed-text-v1.5/resolve/main";
54
55/// One platform's ONNX Runtime archive.
56struct RuntimeArchive {
57    /// The Rust target operating system this is for.
58    os: &'static str,
59    /// The Rust target architecture this is for.
60    arch: &'static str,
61    /// Whether this is the build carrying the CUDA execution provider.
62    gpu: bool,
63    /// The asset name, with `{version}` where the version goes.
64    asset: &'static str,
65    /// SHA-256 of the asset at [`DEFAULT_RUNTIME`], lowercase hex.
66    ///
67    /// Only the pinned version's digest is here, because a digest is a fact
68    /// about one file and a table of them per version would be a table nobody
69    /// updates.
70    sha256: &'static str,
71}
72
73/// Every archive this build knows how to install.
74///
75/// `universal2` covers both macOS architectures, which is why there is one row
76/// for the platform rather than two.
77const RUNTIMES: &[RuntimeArchive] = &[
78    RuntimeArchive {
79        os: "windows",
80        arch: "x86_64",
81        gpu: false,
82        asset: "onnxruntime-win-x64-{version}.zip",
83        sha256: "174c616efc0271194488642a72f1a514e01487da4dfe84c49296d66e40ebe0da",
84    },
85    RuntimeArchive {
86        os: "windows",
87        arch: "aarch64",
88        gpu: false,
89        asset: "onnxruntime-win-arm64-{version}.zip",
90        sha256: "7008f7ff82f8e7de563a22f2b590e08e706a1289eba606b93de2b56edfb1e04b",
91    },
92    RuntimeArchive {
93        os: "windows",
94        arch: "x86_64",
95        gpu: true,
96        asset: "onnxruntime-win-x64-gpu-{version}.zip",
97        sha256: "5b5241716b2628c1ab5e79ee620be767531021149ee68f30fc46c16263fb94dd",
98    },
99    RuntimeArchive {
100        os: "macos",
101        arch: "x86_64",
102        gpu: false,
103        asset: "onnxruntime-osx-universal2-{version}.tgz",
104        sha256: "cfa6f6584d87555ed9f6e7e8a000d3947554d589efe3723b8bfa358cd263d03c",
105    },
106    RuntimeArchive {
107        os: "macos",
108        arch: "aarch64",
109        gpu: false,
110        asset: "onnxruntime-osx-universal2-{version}.tgz",
111        sha256: "cfa6f6584d87555ed9f6e7e8a000d3947554d589efe3723b8bfa358cd263d03c",
112    },
113    RuntimeArchive {
114        os: "linux",
115        arch: "x86_64",
116        gpu: false,
117        asset: "onnxruntime-linux-x64-{version}.tgz",
118        sha256: "8344d55f93d5bc5021ce342db50f62079daf39aaafb5d311a451846228be49b3",
119    },
120    RuntimeArchive {
121        os: "linux",
122        arch: "aarch64",
123        gpu: false,
124        asset: "onnxruntime-linux-aarch64-{version}.tgz",
125        sha256: "bb76395092d150b52c7092dc6b8f2fe4d80f0f3bf0416d2f269193e347e24702",
126    },
127    RuntimeArchive {
128        os: "linux",
129        arch: "x86_64",
130        gpu: true,
131        asset: "onnxruntime-linux-x64-gpu-{version}.tgz",
132        sha256: "2a19dbfa403672ec27378c3d40a68f793ac7a6327712cd0e8240a86be2b10c55",
133    },
134];
135
136/// One file of the model, and the digest it must have.
137struct ModelFile {
138    /// The path under the Hugging Face repository.
139    remote: &'static str,
140    /// The name it is installed under.
141    local: &'static str,
142    /// SHA-256, lowercase hex.
143    sha256: &'static str,
144    /// Its size, so the progress bar has a total before a byte arrives.
145    bytes: u64,
146}
147
148/// The five files `nomic-embed-text-v1.5` needs to run.
149///
150/// The ONNX export is under `onnx/`, and it is the **fp32** one. The repository
151/// also publishes fp16, int8, uint4 and three other quantized exports, and
152/// `docs/embeddings.md` prints what two of them cost: int8 is faster and its
153/// query vector agrees with this one at 0.9727 cosine, which is a retrieval
154/// change rather than a speed change. Installing one of those by accident is
155/// exactly the kind of quiet wrongness a pinned list prevents.
156const NOMIC_FILES: &[ModelFile] = &[
157    ModelFile {
158        remote: "onnx/model.onnx",
159        local: "model.onnx",
160        sha256: "147d5aa88c2101237358e17796cf3a227cead1ec304ec34b465bb08e9d952965",
161        bytes: 547_310_275,
162    },
163    ModelFile {
164        remote: "tokenizer.json",
165        local: "tokenizer.json",
166        sha256: "d241a60d5e8f04cc1b2b3e9ef7a4921b27bf526d9f6050ab90f9267a1f9e5c66",
167        bytes: 711_396,
168    },
169    ModelFile {
170        remote: "tokenizer_config.json",
171        local: "tokenizer_config.json",
172        sha256: "d7e0000bcc80134debd2222220427e6bf5fa20a669f40a0d0d1409cc18e0a9bc",
173        bytes: 1_191,
174    },
175    ModelFile {
176        remote: "special_tokens_map.json",
177        local: "special_tokens_map.json",
178        sha256: "5d5b662e421ea9fac075174bb0688ee0d9431699900b90662acd44b2a350503a",
179        bytes: 695,
180    },
181    ModelFile {
182        remote: "config.json",
183        local: "config.json",
184        sha256: "9ab00bd92cee80a569f708140b7b6c1661a65891ff3765b1519e181ba2f2c92b",
185        bytes: 2_538,
186    },
187];
188
189/// Where the reranker's files come from.
190///
191/// **A commit and not `main`**, so a later push to the repository cannot change what a fresh
192/// install downloads, and the digests below stay true. The commit, `f7481e60...`, is the
193/// repository's head on 29 September 2026, from
194/// `https://huggingface.co/api/models/Alibaba-NLP/gte-reranker-modernbert-base`. The model's
195/// licence, read from its README at that commit, is Apache 2.0.
196const RERANKER_BASE: &str = "https://huggingface.co/Alibaba-NLP/gte-reranker-modernbert-base/resolve/f7481e6055501a30fb19d090657df9ec1f79ab2c";
197
198/// The five files `gte-reranker-modernbert-base` needs to run, each pinned by digest.
199///
200/// `onnx/model.onnx` is the fp32 export, 599 MB. The repository also publishes fp16, int8 and
201/// four other quantized exports. The study measured the fp32 one, and installing another by
202/// accident would change the scores without any message.
203const RERANKER_FILES: &[ModelFile] = &[
204    ModelFile {
205        remote: "onnx/model.onnx",
206        local: "model.onnx",
207        sha256: "c6d3226502addbcd4d2cf273802957ebf8a2a6bf94037dcb9b1d95bfc01e5d93",
208        bytes: 598_803_940,
209    },
210    ModelFile {
211        remote: "tokenizer.json",
212        local: "tokenizer.json",
213        sha256: "2aea6ff4701d063e7e029b6be695a1659f2caaa2ae4fb0e8b18285818271becd",
214        bytes: 3_583_499,
215    },
216    ModelFile {
217        remote: "tokenizer_config.json",
218        local: "tokenizer_config.json",
219        sha256: "626c86908d7c711f93b0feffd8657b782cc2727b391f9be190240e8cafb626d5",
220        bytes: 21_031,
221    },
222    ModelFile {
223        remote: "special_tokens_map.json",
224        local: "special_tokens_map.json",
225        sha256: "ea97ecdbcc73713039d8d64dbb05e3689495c96657fbd9a18f5bed381be81049",
226        bytes: 694,
227    },
228    ModelFile {
229        remote: "config.json",
230        local: "config.json",
231        sha256: "c9316ff715158502dad782f35454eee18de984160618dc30afe7508feb46b7ce",
232        bytes: 1_333,
233    },
234];
235
236/// One model this command can install: where it comes from, which files it needs, and the manifest
237/// it is sealed with.
238struct ModelSpec {
239    /// The model id, which is its directory name.
240    id: &'static str,
241    /// The URL every file is fetched from, with the file's remote path added.
242    base: &'static str,
243    /// The model's page, recorded in the manifest and the install state.
244    source: &'static str,
245    /// Every file, with its digest.
246    files: &'static [ModelFile],
247    /// The manifest that describes the model.
248    manifest: fn() -> ModelManifest,
249}
250
251/// The embedding model: `nomic-embed-text-v1.5`.
252const EMBEDDER: ModelSpec = ModelSpec {
253    id: install::DEFAULT_MODEL,
254    base: MODEL_BASE,
255    source: "https://huggingface.co/nomic-ai/nomic-embed-text-v1.5",
256    files: NOMIC_FILES,
257    manifest: ModelManifest::nomic_v1_5,
258};
259
260/// The reranker: `gte-reranker-modernbert-base`, from a pinned commit.
261const RERANKER: ModelSpec = ModelSpec {
262    id: install::RERANKER_MODEL,
263    base: RERANKER_BASE,
264    source: "https://huggingface.co/Alibaba-NLP/gte-reranker-modernbert-base",
265    files: RERANKER_FILES,
266    manifest: ModelManifest::gte_reranker_modernbert_base,
267};
268
269/// Which parts of the install to do.
270#[derive(Debug, Clone, Copy, PartialEq, Eq)]
271enum Component {
272    /// The runtime and the embedding model. Not the reranker, so nobody gets an unexpected 600 MB.
273    All,
274    /// The ONNX Runtime shared library.
275    Runtime,
276    /// The weights.
277    Model,
278    /// The reranker, a cross encoder of about 600 MB.
279    Reranker,
280}
281
282impl Component {
283    /// Parses the positional argument.
284    ///
285    /// @param text - `all`, `runtime` or `model`
286    fn parse(text: &str) -> Result<Component, Failed> {
287        match text.trim().to_ascii_lowercase().as_str() {
288            "" | "all" | "both" => Ok(Component::All),
289            "runtime" | "onnxruntime" | "onnx" => Ok(Component::Runtime),
290            "model" | "weights" | "embeddings" => Ok(Component::Model),
291            "reranker" | "rerank" | "cross-encoder" => Ok(Component::Reranker),
292            other => Err(Failed::misuse(format!(
293                "'{other}' is not a component. Use all, runtime, model or reranker."
294            ))),
295        }
296    }
297}
298
299/// `setup-embeddings`: download and install the embedder.
300///
301/// @param context - the surface, which may be confined
302/// @param arguments - what was asked for
303pub fn setup_embeddings(context: &mut Context, arguments: &Arguments) -> Result<Outcome, Failed> {
304    let root = match arguments.text("dir") {
305        Some(named) => context.confine(named)?,
306        None => install::home(),
307    };
308
309    // A bare `inillucent setup-embeddings` reports and installs nothing.
310    //
311    // The command downloads about 620 MB, and a verb typed to find out what it
312    // does should not start that. It was not a hypothetical: the command table's
313    // own parity suite calls every command with no arguments to check that each
314    // one either answers or refuses with a usable message, and the first version
315    // of this command answered by fetching the whole model - so `cargo test`
316    // pulled 535 MB into the machine's install directory. The bare call now
317    // prints the status and the one line that starts the install.
318    let named = arguments.text("component").map(str::trim).unwrap_or("");
319    // The profile is read before anything is fetched, so a typo in it costs a
320    // message rather than 620 MB and then a message.
321    let residency = match arguments.text("residency") {
322        Some(text) => {
323            Some(Residency::parse(text).map_err(|reason| Failed::misuse(format!("{reason:#}")))?)
324        }
325        None => None,
326    };
327
328    let machine = machine_settings(arguments)?;
329
330    if arguments.flag("status") || named.is_empty() {
331        // Changing when the model is in memory, how many threads it uses or which
332        // processor it runs on is not a reason to fetch it again, so a setting
333        // given without a component is recorded here and nothing is downloaded.
334        if residency.is_some() || machine.any() {
335            let mut state = install::read_state(&root).unwrap_or_default();
336            if let Some(residency) = residency {
337                state.residency = Some(residency.label());
338            }
339            machine.record(&mut state);
340            install::write_state(&root, &state).map_err(|error| {
341                Failed::misuse(format!("the install state could not be written: {error}"))
342            })?;
343        }
344        let mut outcome = status(&root);
345        if named.is_empty() && !arguments.flag("status") {
346            outcome.text.push_str(
347                "\n\nTo install what is missing, naming what you want so a 620 MB download \
348                 is never a surprise:\n  inillucent setup-embeddings all",
349            );
350        }
351        return Ok(outcome);
352    }
353
354    let component = Component::parse(named)?;
355    let version = arguments
356        .text("onnxruntime-version")
357        .unwrap_or(DEFAULT_RUNTIME)
358        .to_string();
359    let gpu = arguments.flag("gpu");
360    let force = arguments.flag("force");
361
362    let mut state = install::read_state(&root).unwrap_or_default();
363    let mut lines: Vec<String> = Vec::new();
364    let mut fields: Vec<(String, Json)> = Vec::new();
365
366    if matches!(component, Component::All | Component::Runtime) {
367        let installed = install_runtime(
368            &root,
369            &version,
370            if gpu {
371                Accelerator::Gpu
372            } else {
373                Accelerator::Cpu
374            },
375            if force {
376                Reinstall::Always
377            } else {
378                Reinstall::WhenMissing
379            },
380        )?;
381        lines.push(format!(
382            "ONNX Runtime {} -> {}{}",
383            installed.version,
384            installed.library,
385            if installed.verified {
386                ""
387            } else {
388                "  (digest not pinned in this build)"
389            }
390        ));
391        fields.push(("runtime".to_string(), runtime_json(&installed)));
392        state.runtime = Some(installed);
393    }
394
395    if matches!(component, Component::All | Component::Model) {
396        let installed = install_model(&root, &EMBEDDER, force)?;
397        lines.push(format!("{} -> {}", installed.id, installed.dir));
398        fields.push(("model".to_string(), model_json(&installed)));
399        state.put_model(installed);
400    }
401
402    if component == Component::Reranker {
403        let installed = install_model(&root, &RERANKER, force)?;
404        lines.push(format!("{} -> {}", installed.id, installed.dir));
405        fields.push(("reranker".to_string(), model_json(&installed)));
406        state.put_model(installed);
407    }
408
409    if let Some(residency) = residency {
410        state.residency = Some(residency.label());
411        lines.push(format!("residency profile: {}", residency.label()));
412    }
413    machine.record(&mut state);
414    lines.extend(machine.lines());
415    let effective = state
416        .residency
417        .as_deref()
418        .and_then(|text| Residency::parse(text).ok())
419        .unwrap_or_default();
420    fields.push(("residency".to_string(), json::text(effective.label())));
421
422    install::write_state(&root, &state).map_err(|error| {
423        Failed::misuse(format!("the install state could not be written: {error}"))
424    })?;
425    // The downloads directory holds only partial fetches, and every one of them
426    // has either been renamed into place or deleted by the time this runs. It is
427    // removed rather than left, because an empty directory nobody explains is a
428    // question somebody has to answer later.
429    let _ = std::fs::remove_dir(install::downloads_dir(&root));
430
431    lines.push(String::new());
432    lines.push(format!("Installed under {}.", root.display()));
433    lines.push(
434        "Nothing to export: the engine finds both of these on its own. Check it with".to_string(),
435    );
436    lines.push("  inillucent --db test.rdb query \"SELECT length(embed('hello'))\"".to_string());
437
438    let mut outcome = Outcome::said("setup-embeddings", lines.join("\n"));
439    outcome = outcome.with("root", json::text(root.display().to_string()));
440    for (name, value) in fields {
441        outcome = outcome.with(&name, value);
442    }
443    Ok(outcome)
444}
445
446/// The thread count and the device `setup-embeddings` was asked to record.
447struct MachineSettings {
448    threads: Option<usize>,
449    device: Option<String>,
450}
451
452impl MachineSettings {
453    /// Reports whether either setting was given.
454    fn any(&self) -> bool {
455        self.threads.is_some() || self.device.is_some()
456    }
457
458    /// Writes the settings that were given into the install state, and leaves the others as they were.
459    ///
460    /// @param state - the install state about to be written
461    fn record(&self, state: &mut install::State) {
462        if let Some(threads) = self.threads {
463            state.threads = Some(threads);
464        }
465        if let Some(device) = self.device.as_ref() {
466            state.device = Some(device.clone());
467        }
468    }
469
470    /// Returns the lines that tell the caller what was recorded.
471    fn lines(&self) -> Vec<String> {
472        let mut lines = Vec::new();
473        if let Some(threads) = self.threads {
474            lines.push(format!("threads: {threads}"));
475        }
476        if let Some(device) = self.device.as_ref() {
477            lines.push(format!("device: {device}"));
478        }
479        lines
480    }
481}
482
483/// Reads `--threads` and `--device`, refusing a value that cannot be used before anything is fetched.
484///
485/// @param arguments - what was asked for
486fn machine_settings(arguments: &Arguments) -> Result<MachineSettings, Failed> {
487    let threads = match arguments.integer("threads") {
488        Some(count) => Some(install::parse_threads(&count.to_string()).map_err(Failed::misuse)?),
489        None => None,
490    };
491    let device = match arguments.text("device") {
492        Some(text) => Some(install::parse_device(text).map_err(Failed::misuse)?),
493        None => None,
494    };
495    Ok(MachineSettings { threads, device })
496}
497
498/// Describes one machine setting for `--status`: its value, then where the value came from.
499///
500/// @param name - the setting's name
501/// @param value - the value in force
502/// @param source - where it came from
503/// @param variable - the environment variable that overrides it
504fn setting_line(name: &str, value: &str, source: install::SettingSource, variable: &str) -> String {
505    format!(
506        "{name}: {value} ({}). {variable} overrides it for one process",
507        source.label()
508    )
509}
510
511/// What `--status` prints.
512///
513/// It reports what is *there* rather than what the state file claims, because
514/// the two differ exactly when something has gone wrong - a directory moved, a
515/// disk cleaned - and that is the case a status command exists for.
516///
517/// @param root - the install root
518fn status(root: &Path) -> Outcome {
519    let state = install::read_state(root).unwrap_or_default();
520    let mut lines = vec![format!("Install root: {}", root.display())];
521
522    match install::runtime_library() {
523        Some(library) => {
524            let recorded = state.runtime.as_ref();
525            let version = recorded
526                .map(|r| r.version.as_str())
527                .unwrap_or("unknown version");
528            let unverified = recorded.is_some_and(|r| !r.verified);
529            lines.push(format!(
530                "ONNX Runtime: {} ({version}){}",
531                library.display(),
532                if unverified {
533                    "  (digest not pinned in this build)"
534                } else {
535                    ""
536                }
537            ));
538        }
539        None => lines.push(
540            "ONNX Runtime: not installed. Run: inillucent setup-embeddings runtime".to_string(),
541        ),
542    }
543
544    match install::model_dir(install::DEFAULT_MODEL) {
545        Some(dir) => lines.push(format!("{}: {}", install::DEFAULT_MODEL, dir.display())),
546        None => lines.push(format!(
547            "{}: not installed. Run: inillucent setup-embeddings model",
548            install::DEFAULT_MODEL
549        )),
550    }
551
552    let effective = Residency::configured();
553    lines.push(format!("Residency profile: {}", effective.label()));
554    lines.push(match effective {
555        Residency::Resident => {
556            "  loaded on first use and kept, which is about 1.9 GB held and 12 to 36 ms a query"
557                .to_string()
558        }
559        Residency::OnDemand => {
560            "  loaded per call and dropped, which is nothing held and about 0.8 s a query"
561                .to_string()
562        }
563        Residency::Idle(after) => format!(
564            "  loaded on use and dropped after {}s idle: the first query in a burst pays about \
565             0.8 s and the rest pay 12 to 36 ms",
566            after.as_secs()
567        ),
568    });
569
570    let threads = install::configured_threads();
571    let device = install::configured_device();
572    let threads_text = threads
573        .value
574        .map_or("ONNX Runtime's own choice".to_string(), |count| {
575            count.to_string()
576        });
577    lines.push(setting_line(
578        "Threads",
579        &threads_text,
580        threads.source,
581        install::THREADS_VAR,
582    ));
583    lines.push(setting_line(
584        "Device",
585        &device.value,
586        device.source,
587        install::DEVICE_VAR,
588    ));
589
590    match install::model_dir(install::RERANKER_MODEL) {
591        Some(dir) => lines.push(format!("{}: {}", install::RERANKER_MODEL, dir.display())),
592        None => lines.push(format!(
593            "{}: not installed. Run: inillucent setup-embeddings reranker (about 600 MB, needed \
594             only for rerank() and a search that names question)",
595            install::RERANKER_MODEL
596        )),
597    }
598
599    let ready = install::runtime_library().is_some()
600        && install::model_dir(install::DEFAULT_MODEL).is_some();
601    Outcome::said("setup-embeddings", lines.join("\n"))
602        .with("root", json::text(root.display().to_string()))
603        .with("ready", Json::Bool(ready))
604        .with("residency", json::text(effective.label()))
605        .with(
606            "threads",
607            setting_json(threads.value.map(|count| count.to_string()), threads.source),
608        )
609        .with(
610            "device",
611            setting_json(Some(device.value.clone()), device.source),
612        )
613        .with(
614            "runtime",
615            match state.runtime.as_ref() {
616                Some(runtime) => runtime_json(runtime),
617                None => Json::Null,
618            },
619        )
620        .with(
621            "model",
622            match state.model(install::DEFAULT_MODEL) {
623                Some(model) => model_json(model),
624                None => Json::Null,
625            },
626        )
627        .with(
628            "reranker_installed",
629            Json::Bool(install::model_dir(install::RERANKER_MODEL).is_some()),
630        )
631        .with(
632            "reranker",
633            match state.model(install::RERANKER_MODEL) {
634                Some(model) => model_json(model),
635                None => Json::Null,
636            },
637        )
638}
639
640/// One machine setting as JSON: its value, or null for the default, and where it came from.
641///
642/// @param value - the value in force, when there is one
643/// @param source - where it came from
644fn setting_json(value: Option<String>, source: install::SettingSource) -> Json {
645    json::object(vec![
646        ("value", value.map_or(Json::Null, json::text)),
647        ("source", json::text(source.label())),
648    ])
649}
650
651/// Downloads and installs the ONNX Runtime shared library.
652///
653/// @param root - the install root
654/// @param version - the ONNX Runtime version
655/// @param gpu - whether to take the build carrying the CUDA execution provider
656/// @param force - install again even when it is already there
657/// Which build of the runtime `setup-embeddings` installs.
658///
659/// **An enum rather than a `bool` beside another `bool` (task-1962, A9).**
660/// `install_runtime` took `gpu` and `force` adjacent and positional, and the
661/// call site read `install_runtime(&root, &version, gpu, force)` - two words
662/// that say which is which only because they happen to be named after the
663/// parameters.
664#[derive(Clone, Copy, Debug, Eq, PartialEq)]
665pub enum Accelerator {
666    /// The CPU build, which every machine can run.
667    Cpu,
668    /// The GPU build, which needs a supported card and its driver.
669    Gpu,
670}
671
672/// Whether an install replaces a runtime that is already there.
673#[derive(Clone, Copy, Debug, Eq, PartialEq)]
674pub enum Reinstall {
675    /// Download and unpack even when the version is already installed.
676    Always,
677    /// Leave an installed version alone.
678    WhenMissing,
679}
680
681fn install_runtime(
682    root: &Path,
683    version: &str,
684    accelerator: Accelerator,
685    force: Reinstall,
686) -> Result<InstalledRuntime, Failed> {
687    let gpu = accelerator == Accelerator::Gpu;
688    let force = force == Reinstall::Always;
689    let archive_spec = pick_runtime(gpu)?;
690    let asset = archive_spec.asset.replace("{version}", version);
691    let directory = install::runtime_dir(root, version);
692    let library = directory.join("lib").join(install::runtime_library_name());
693
694    if library.exists() && !force {
695        return Ok(InstalledRuntime {
696            version: version.to_string(),
697            archive: asset,
698            library: library.display().to_string(),
699            verified: true,
700            gpu,
701        });
702    }
703
704    // Only the pinned version's digest is a fact about the file being fetched.
705    // Asking for another version is allowed and is reported as unverified.
706    let expected = (version == DEFAULT_RUNTIME).then_some(archive_spec.sha256);
707    let url = format!("{RUNTIME_BASE}/v{version}/{asset}");
708    let downloaded = install::downloads_dir(root).join(&asset);
709    let mut progress = bar();
710    let fetched = http::download(&url, &downloaded, expected, &mut progress)
711        .map_err(|error| Failed::misuse(format!("{error}")))?;
712
713    let bytes = std::fs::read(&fetched.path).map_err(|error| {
714        Failed::misuse(format!(
715            "{} could not be read: {error}",
716            fetched.path.display()
717        ))
718    })?;
719    let members = archive::read(&bytes).map_err(|error| Failed::misuse(format!("{error}")))?;
720    let libraries = shared_libraries(&members);
721    if libraries.is_empty() {
722        return Err(Failed::misuse(format!(
723            "{asset} holds no {} - it is not an ONNX Runtime release, or its layout has changed",
724            install::runtime_library_name()
725        )));
726    }
727
728    let lib_dir = directory.join("lib");
729    std::fs::create_dir_all(&lib_dir).map_err(|error| {
730        Failed::misuse(format!(
731            "{} could not be created: {error}",
732            lib_dir.display()
733        ))
734    })?;
735    for (member, name) in &libraries {
736        archive::write_member(member, &lib_dir, Some(name))
737            .map_err(|error| Failed::misuse(format!("{error}")))?;
738    }
739    // The archive is a download, not an installed component, and it is between
740    // 7 MB and 300 MB. It goes once the library is out of it.
741    let _ = std::fs::remove_file(&fetched.path);
742
743    if !library.exists() {
744        return Err(Failed::misuse(format!(
745            "{asset} was extracted and {} is still not there",
746            library.display()
747        )));
748    }
749
750    Ok(InstalledRuntime {
751        version: version.to_string(),
752        archive: asset,
753        library: library.display().to_string(),
754        verified: expected.is_some(),
755        gpu,
756    })
757}
758
759/// The archive for the machine this is running on.
760///
761/// @param gpu - whether the CUDA build was asked for
762fn pick_runtime(gpu: bool) -> Result<&'static RuntimeArchive, Failed> {
763    let os = std::env::consts::OS;
764    let arch = std::env::consts::ARCH;
765    RUNTIMES
766        .iter()
767        .find(|entry| entry.os == os && entry.arch == arch && entry.gpu == gpu)
768        .ok_or_else(|| {
769            let plain = RUNTIMES
770                .iter()
771                .any(|e| e.os == os && e.arch == arch && !e.gpu);
772            if gpu && plain {
773                Failed::misuse(format!(
774                    "there is no CUDA build of ONNX Runtime for {os} on {arch}. Run the command \
775                     without --gpu."
776                ))
777            } else {
778                Failed::misuse(format!(
779                    "there is no ONNX Runtime release for {os} on {arch} that this command knows \
780                     how to install. Build it, and point ORT_DYLIB_PATH at the result."
781                ))
782            }
783        })
784}
785
786/// Picks the shared libraries out of an archive, with the names to install them
787/// under.
788///
789/// Two things make this less obvious than a name match. The real library on
790/// macOS and Linux is *versioned* - `libonnxruntime.so.1.22.0`,
791/// `libonnxruntime.1.22.0.dylib` - and the unversioned name beside it is a
792/// symbolic link, which this extractor does not carry across; so the versioned
793/// file is installed under the unversioned name that the loader will ask for.
794/// And the macOS archive carries a `.dSYM` bundle holding a 150 MB file with
795/// `.dylib` in its path, which is debug information rather than a library.
796///
797/// @param members - everything in the archive
798fn shared_libraries(members: &[Member]) -> Vec<(&Member, String)> {
799    let wanted = install::runtime_library_name();
800    let mut out = Vec::new();
801    for member in members {
802        let name = member.name.rsplit('/').next().unwrap_or(&member.name);
803        if member.name.contains(".dSYM/") {
804            continue;
805        }
806        if !is_shared_library(name) {
807            continue;
808        }
809        if is_main_library(name) {
810            out.push((member, wanted.to_string()));
811        } else if name.contains("onnxruntime_providers") {
812            // A provider library keeps its own name, because the main library
813            // loads it by that name at run time. The version suffix is dropped
814            // for the same reason it is on the main library.
815            out.push((member, unversioned(name)));
816        }
817    }
818    out
819}
820
821/// Whether a file name is a shared library on some platform.
822///
823/// @param name - the base name
824fn is_shared_library(name: &str) -> bool {
825    name.ends_with(".dll") || name.ends_with(".dylib") || name.contains(".so")
826}
827
828/// Whether a file name is ONNX Runtime itself rather than one of its providers.
829///
830/// @param name - the base name
831fn is_main_library(name: &str) -> bool {
832    if name.contains("providers") {
833        return false;
834    }
835    name == "onnxruntime.dll"
836        || name.starts_with("libonnxruntime.so")
837        || (name.starts_with("libonnxruntime.") && name.ends_with(".dylib"))
838        || name == "libonnxruntime.dylib"
839}
840
841/// A library's name with any version numbers taken out of it.
842///
843/// `libonnxruntime_providers_cuda.so.1.22.0` becomes
844/// `libonnxruntime_providers_cuda.so`, which is the name the main library asks
845/// the loader for.
846///
847/// @param name - the base name as it is in the archive
848fn unversioned(name: &str) -> String {
849    if let Some(at) = name.find(".so") {
850        return format!("{}.so", name.get(..at).unwrap_or(name));
851    }
852    if name.ends_with(".dylib") {
853        let stem = name.trim_end_matches(".dylib");
854        let base = stem.split('.').next().unwrap_or(stem);
855        return format!("{base}.dylib");
856    }
857    name.to_string()
858}
859
860/// Downloads and installs one model's weights, and seals a manifest over what landed.
861///
862/// @param root - the install root
863/// @param spec - which model, where it comes from and what its files must hash to
864/// @param force - fetch again even when the files are already there
865fn install_model(root: &Path, spec: &ModelSpec, force: bool) -> Result<InstalledModel, Failed> {
866    let directory = install::models_root(root).join(spec.id);
867    std::fs::create_dir_all(&directory).map_err(|error| {
868        Failed::misuse(format!(
869            "{} could not be created: {error}",
870            directory.display()
871        ))
872    })?;
873
874    for file in spec.files {
875        let destination = directory.join(file.local);
876        if destination.exists() && !force && already_correct(&destination, file) {
877            continue;
878        }
879        let url = format!("{}/{}", spec.base, file.remote);
880        let mut progress = bar();
881        http::download(&url, &destination, Some(file.sha256), &mut progress)
882            .map_err(|error| Failed::misuse(format!("{error}")))?;
883    }
884
885    // The manifest is this repository's contract rather than the model author's,
886    // so it is written here. Its digests come from the files that actually
887    // landed, which is what makes it impossible for a manifest to describe
888    // weights that are not there.
889    let mut manifest = (spec.manifest)();
890    manifest.weights_sha256 = spec
891        .files
892        .iter()
893        .find(|f| f.local == "model.onnx")
894        .map(|f| f.sha256.to_string())
895        .unwrap_or_default();
896    manifest.tokenizer_sha256 = spec
897        .files
898        .iter()
899        .find(|f| f.local == "tokenizer.json")
900        .map(|f| f.sha256.to_string())
901        .unwrap_or_default();
902    manifest.source = Some(spec.source.to_string());
903    manifest
904        .write(&directory)
905        .map_err(|reason| Failed::misuse(format!("the model manifest: {reason:#}")))?;
906
907    Ok(InstalledModel {
908        id: spec.id.to_string(),
909        dir: directory.display().to_string(),
910        source: spec.source.to_string(),
911        verified: true,
912    })
913}
914
915/// Whether a file on disk is already the one that would be downloaded.
916///
917/// Checked by size first and by digest only when the size matches, because
918/// digesting 547 MB costs a second and a size mismatch settles it for nothing.
919/// A digest rather than a size alone, because a truncated file that happens to
920/// be the right length is precisely what a resumed download can leave.
921///
922/// @param path - the file on disk
923/// @param file - what it is supposed to be
924fn already_correct(path: &Path, file: &ModelFile) -> bool {
925    let Ok(metadata) = std::fs::metadata(path) else {
926        return false;
927    };
928    if metadata.len() != file.bytes {
929        return false;
930    }
931    let Ok(mut handle) = std::fs::File::open(path) else {
932        return false;
933    };
934    let mut digest = inillucent_base::hash::Sha256::new();
935    let mut buffer = vec![0u8; 1 << 20];
936    loop {
937        use std::io::Read;
938        match handle.read(&mut buffer) {
939            Ok(0) => break,
940            Ok(read) => digest.update(buffer.get(..read).unwrap_or(&[])),
941            Err(_) => return false,
942        }
943    }
944    inillucent_base::hash::to_hex(&digest.finish()).eq_ignore_ascii_case(file.sha256)
945}
946
947/// The installed runtime, as JSON.
948///
949/// @param runtime - what was installed
950fn runtime_json(runtime: &InstalledRuntime) -> Json {
951    json::object(vec![
952        ("version", json::text(&runtime.version)),
953        ("archive", json::text(&runtime.archive)),
954        ("library", json::text(&runtime.library)),
955        ("verified", Json::Bool(runtime.verified)),
956        ("gpu", Json::Bool(runtime.gpu)),
957    ])
958}
959
960/// The installed model, as JSON.
961///
962/// @param model - what was installed
963fn model_json(model: &InstalledModel) -> Json {
964    json::object(vec![
965        ("id", json::text(&model.id)),
966        ("dir", json::text(&model.dir)),
967        ("source", json::text(&model.source)),
968        ("verified", Json::Bool(model.verified)),
969    ])
970}
971
972/// Builds the progress reporter for this terminal.
973fn bar() -> Bar {
974    Bar {
975        terminal: std::io::stderr().is_terminal(),
976        name: String::new(),
977        width: 0,
978        last_percent: -1,
979        started: std::time::Instant::now(),
980    }
981}
982
983/// A one-line progress bar, on standard error.
984///
985/// Standard error rather than standard output, so `--output json` stays
986/// parseable while the download runs. And **nothing at all when standard error
987/// is not a terminal** except a line per ten percent: a bar rewritten with
988/// carriage returns into a log file is one enormous line, and a log with no
989/// progress in it at all cannot be used to tell a slow download from a stalled
990/// one.
991struct Bar {
992    terminal: bool,
993    name: String,
994    /// The whole size, when the server said what it is.
995    width: u64,
996    /// The last decile reported, for the non-terminal case.
997    last_percent: i64,
998    started: std::time::Instant,
999}
1000
1001impl Bar {
1002    /// Draws the line for a given position.
1003    ///
1004    /// @param done - bytes so far
1005    fn draw(&self, done: u64) {
1006        let elapsed = self.started.elapsed().as_secs_f64().max(0.001);
1007        let rate = done as f64 / elapsed / (1024.0 * 1024.0);
1008        let mut line = if self.width > 0 {
1009            let share = (done as f64 / self.width as f64).clamp(0.0, 1.0);
1010            let filled = (share * 20.0).round() as usize;
1011            format!(
1012                "  {:<38} [{}{}] {:3.0}%  {:.1}/{:.1} MB  {rate:.1} MB/s",
1013                self.name,
1014                "#".repeat(filled.min(20)),
1015                ".".repeat(20usize.saturating_sub(filled)),
1016                share * 100.0,
1017                done as f64 / (1024.0 * 1024.0),
1018                self.width as f64 / (1024.0 * 1024.0),
1019            )
1020        } else {
1021            format!(
1022                "  {:<38} {:.1} MB  {rate:.1} MB/s",
1023                self.name,
1024                done as f64 / (1024.0 * 1024.0)
1025            )
1026        };
1027        line.push('\r');
1028        let mut stderr = std::io::stderr();
1029        let _ = stderr.write_all(line.as_bytes());
1030        let _ = stderr.flush();
1031    }
1032}
1033
1034impl Progress for Bar {
1035    /// Notes what is being fetched and how big it is.
1036    fn started(&mut self, name: &str, total: Option<u64>, resumed: u64) {
1037        self.name = name.to_string();
1038        self.width = total.unwrap_or(0);
1039        self.last_percent = -1;
1040        self.started = std::time::Instant::now();
1041        if resumed > 0 {
1042            eprintln!(
1043                "  {name}: resuming at {:.1} MB",
1044                resumed as f64 / (1024.0 * 1024.0)
1045            );
1046        }
1047    }
1048
1049    /// Redraws the bar, or reports another ten percent into a log.
1050    fn advanced(&mut self, done: u64, _total: Option<u64>) {
1051        if self.terminal {
1052            self.draw(done);
1053            return;
1054        }
1055        if self.width == 0 {
1056            return;
1057        }
1058        let percent = (done as i64).saturating_mul(100) / self.width.max(1) as i64;
1059        if percent / 10 > self.last_percent / 10 {
1060            self.last_percent = percent;
1061            eprintln!("  {}: {percent}%", self.name);
1062        }
1063    }
1064
1065    /// Ends the line, so whatever prints next starts on its own.
1066    fn finished(&mut self, done: u64) {
1067        if self.terminal {
1068            self.draw(done);
1069        }
1070        eprintln!(
1071            "  {:<38} {:.1} MB",
1072            self.name,
1073            done as f64 / (1024.0 * 1024.0)
1074        );
1075    }
1076}
1077
1078#[cfg(test)]
1079mod tests {
1080    use super::*;
1081
1082    /// Every platform this ships on has an archive, and every pinned digest is
1083    /// a SHA-256.
1084    ///
1085    /// The digest check is not decoration: a row with a truncated or
1086    /// pasted-wrong digest would refuse every download on that platform, and
1087    /// this repository builds on one platform at a time.
1088    #[test]
1089    fn every_platform_has_an_archive_with_a_pinned_digest() {
1090        for (os, arch) in [
1091            ("windows", "x86_64"),
1092            ("windows", "aarch64"),
1093            ("macos", "x86_64"),
1094            ("macos", "aarch64"),
1095            ("linux", "x86_64"),
1096            ("linux", "aarch64"),
1097        ] {
1098            let found = RUNTIMES
1099                .iter()
1100                .find(|entry| entry.os == os && entry.arch == arch && !entry.gpu);
1101            assert!(found.is_some(), "{os} on {arch} has no archive");
1102        }
1103        for entry in RUNTIMES {
1104            assert_eq!(entry.sha256.len(), 64, "{} has a short digest", entry.asset);
1105            assert!(
1106                entry.sha256.chars().all(|c| c.is_ascii_hexdigit()),
1107                "{} has a digest that is not hex",
1108                entry.asset
1109            );
1110            assert!(
1111                entry.asset.contains("{version}"),
1112                "{} has no version slot",
1113                entry.asset
1114            );
1115        }
1116    }
1117
1118    /// Every model file has a digest and a size, and the two files the manifest
1119    /// seals are among them.
1120    #[test]
1121    fn every_model_file_has_a_digest_and_a_size() {
1122        for file in NOMIC_FILES {
1123            assert_eq!(file.sha256.len(), 64, "{} has a short digest", file.local);
1124            assert!(file.bytes > 0, "{} has no size", file.local);
1125        }
1126        assert!(NOMIC_FILES.iter().any(|f| f.local == "model.onnx"));
1127        assert!(NOMIC_FILES.iter().any(|f| f.local == "tokenizer.json"));
1128    }
1129
1130    /// The fp32 export is what is installed, not one of the quantized ones.
1131    ///
1132    /// `onnx/` in that repository holds seven other exports whose names differ
1133    /// by a suffix, and one of them agrees with this one at 0.9727 cosine. A
1134    /// test that names the file is what stops a plausible-looking edit changing
1135    /// what every vector in every index means.
1136    #[test]
1137    fn the_installed_weights_are_the_full_precision_export() {
1138        let weights = NOMIC_FILES
1139            .iter()
1140            .find(|f| f.local == "model.onnx")
1141            .expect("the weights are in the list");
1142        assert_eq!(weights.remote, "onnx/model.onnx");
1143        assert_eq!(weights.bytes, 547_310_275, "the fp32 export's size");
1144    }
1145
1146    /// A component name parses from every form, and anything else is refused.
1147    #[test]
1148    fn a_component_parses_or_is_refused() {
1149        assert_eq!(Component::parse("all").unwrap(), Component::All);
1150        assert_eq!(Component::parse("").unwrap(), Component::All);
1151        assert_eq!(Component::parse(" RUNTIME ").unwrap(), Component::Runtime);
1152        assert_eq!(Component::parse("model").unwrap(), Component::Model);
1153        assert!(Component::parse("everything").is_err());
1154    }
1155
1156    /// The main library is told apart from its providers, and debug information
1157    /// is not mistaken for either.
1158    #[test]
1159    fn the_main_library_is_told_apart_from_its_providers() {
1160        assert!(is_main_library("onnxruntime.dll"));
1161        assert!(is_main_library("libonnxruntime.so.1.22.0"));
1162        assert!(is_main_library("libonnxruntime.1.22.0.dylib"));
1163        assert!(!is_main_library("onnxruntime_providers_cuda.dll"));
1164        assert!(!is_main_library("libonnxruntime_providers_shared.so"));
1165        assert!(!is_shared_library("onnxruntime.lib"));
1166        assert!(!is_shared_library("onnxruntime.pdb"));
1167        assert!(!is_shared_library("libonnxruntime.pc"));
1168    }
1169
1170    /// A versioned library is installed under the name the loader asks for.
1171    #[test]
1172    fn a_versioned_library_loses_its_version() {
1173        assert_eq!(
1174            unversioned("libonnxruntime_providers_cuda.so.1.22.0"),
1175            "libonnxruntime_providers_cuda.so"
1176        );
1177        assert_eq!(
1178            unversioned("libonnxruntime_providers_shared.so"),
1179            "libonnxruntime_providers_shared.so"
1180        );
1181        assert_eq!(
1182            unversioned("onnxruntime_providers_cuda.dll"),
1183            "onnxruntime_providers_cuda.dll"
1184        );
1185        assert_eq!(
1186            unversioned("libonnxruntime_providers_cuda.1.22.0.dylib"),
1187            "libonnxruntime_providers_cuda.dylib"
1188        );
1189    }
1190
1191    /// The three archive layouts each yield the main library under the name this
1192    /// platform's loader will look for, and the macOS debug bundle is skipped.
1193    ///
1194    /// The names are the ones the real archives hold, listed from
1195    /// `onnxruntime-{win-x64,linux-x64,osx-universal2}-1.22.0`.
1196    #[test]
1197    fn each_archive_layout_yields_the_library_the_loader_wants() {
1198        let members: Vec<Member> = [
1199            "onnxruntime-win-x64-1.22.0/lib/onnxruntime.dll",
1200            "onnxruntime-win-x64-1.22.0/lib/onnxruntime.lib",
1201            "onnxruntime-win-x64-1.22.0/lib/onnxruntime.pdb",
1202            "onnxruntime-linux-x64-1.22.0/lib/libonnxruntime.so.1.22.0",
1203            "onnxruntime-linux-x64-1.22.0/lib/libonnxruntime.pc",
1204            "onnxruntime-osx-universal2-1.22.0/lib/libonnxruntime.1.22.0.dylib",
1205            "onnxruntime-osx-universal2-1.22.0/lib/libonnxruntime.1.22.0.dylib.dSYM/Contents/Resources/DWARF/libonnxruntime.1.22.0.dylib",
1206        ]
1207        .into_iter()
1208        .map(|name| Member { name: name.to_string(), bytes: Vec::new(), executable: false })
1209        .collect();
1210
1211        let picked = shared_libraries(&members);
1212        let names: Vec<&str> = picked.iter().map(|(_, name)| name.as_str()).collect();
1213        assert!(
1214            names
1215                .iter()
1216                .all(|name| *name == install::runtime_library_name()),
1217            "every main library is installed under this platform's name: {names:?}"
1218        );
1219        assert!(
1220            !picked
1221                .iter()
1222                .any(|(member, _)| member.name.contains(".dSYM/")),
1223            "the macOS debug bundle is not a library"
1224        );
1225        // Three archives are listed, so three main libraries are found - one per
1226        // layout. A real install only ever reads one archive.
1227        assert_eq!(picked.len(), 3, "{names:?}");
1228    }
1229
1230    /// A provider library keeps its own name, because the main library loads it
1231    /// by that name.
1232    #[test]
1233    fn a_provider_library_keeps_its_own_name() {
1234        let members = vec![
1235            Member {
1236                name: "onnxruntime-win-x64-gpu-1.22.0/lib/onnxruntime.dll".to_string(),
1237                bytes: Vec::new(),
1238                executable: false,
1239            },
1240            Member {
1241                name: "onnxruntime-win-x64-gpu-1.22.0/lib/onnxruntime_providers_cuda.dll"
1242                    .to_string(),
1243                bytes: Vec::new(),
1244                executable: false,
1245            },
1246        ];
1247        let picked = shared_libraries(&members);
1248        assert_eq!(picked.len(), 2);
1249        assert!(picked
1250            .iter()
1251            .any(|(_, name)| name == "onnxruntime_providers_cuda.dll"));
1252    }
1253}