1use 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
46pub const DEFAULT_RUNTIME: &str = "1.22.0";
48
49const RUNTIME_BASE: &str = "https://github.com/microsoft/onnxruntime/releases/download";
51
52const MODEL_BASE: &str = "https://huggingface.co/nomic-ai/nomic-embed-text-v1.5/resolve/main";
54
55struct RuntimeArchive {
57 os: &'static str,
59 arch: &'static str,
61 gpu: bool,
63 asset: &'static str,
65 sha256: &'static str,
71}
72
73const 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
136struct ModelFile {
138 remote: &'static str,
140 local: &'static str,
142 sha256: &'static str,
144 bytes: u64,
146}
147
148const 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
189const RERANKER_BASE: &str = "https://huggingface.co/Alibaba-NLP/gte-reranker-modernbert-base/resolve/f7481e6055501a30fb19d090657df9ec1f79ab2c";
197
198const 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
236struct ModelSpec {
239 id: &'static str,
241 base: &'static str,
243 source: &'static str,
245 files: &'static [ModelFile],
247 manifest: fn() -> ModelManifest,
249}
250
251const 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
260const 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#[derive(Debug, Clone, Copy, PartialEq, Eq)]
271enum Component {
272 All,
274 Runtime,
276 Model,
278 Reranker,
280}
281
282impl Component {
283 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
299pub 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 let named = arguments.text("component").map(str::trim).unwrap_or("");
319 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 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 let from = folder_to_copy_from(context, arguments, component)?;
362
363 let mut state = install::read_state(&root).unwrap_or_default();
364 let mut lines: Vec<String> = Vec::new();
365 let mut fields: Vec<(String, Json)> = Vec::new();
366
367 if matches!(component, Component::All | Component::Runtime) {
368 let installed = install_runtime(
369 &root,
370 &version,
371 if gpu {
372 Accelerator::Gpu
373 } else {
374 Accelerator::Cpu
375 },
376 if force {
377 Reinstall::Always
378 } else {
379 Reinstall::WhenMissing
380 },
381 )?;
382 lines.push(format!(
383 "ONNX Runtime {} -> {}{}",
384 installed.version,
385 installed.library,
386 if installed.verified {
387 ""
388 } else {
389 " (digest not pinned in this build)"
390 }
391 ));
392 fields.push(("runtime".to_string(), runtime_json(&installed)));
393 state.runtime = Some(installed);
394 }
395
396 if matches!(component, Component::All | Component::Model) {
397 let installed = install_model(&root, &EMBEDDER, force, from.as_deref())?;
398 lines.push(format!("{} -> {}", installed.id, installed.dir));
399 fields.push(("model".to_string(), model_json(&installed)));
400 state.put_model(installed);
401 }
402
403 if component == Component::Reranker {
404 let installed = install_model(&root, &RERANKER, force, from.as_deref())?;
405 lines.push(format!("{} -> {}", installed.id, installed.dir));
406 fields.push(("reranker".to_string(), model_json(&installed)));
407 state.put_model(installed);
408 }
409
410 if let Some(residency) = residency {
411 state.residency = Some(residency.label());
412 lines.push(format!("residency profile: {}", residency.label()));
413 }
414 machine.record(&mut state);
415 lines.extend(machine.lines());
416 let effective = state
417 .residency
418 .as_deref()
419 .and_then(|text| Residency::parse(text).ok())
420 .unwrap_or_default();
421 fields.push(("residency".to_string(), json::text(effective.label())));
422
423 install::write_state(&root, &state).map_err(|error| {
424 Failed::misuse(format!("the install state could not be written: {error}"))
425 })?;
426 let _ = std::fs::remove_dir(install::downloads_dir(&root));
431
432 lines.push(String::new());
433 lines.push(format!("Installed under {}.", root.display()));
434 lines.push(
435 "Nothing to export: the engine finds both of these on its own. Check it with".to_string(),
436 );
437 lines.push(" inillucent --db test.rdb query \"SELECT length(embed('hello'))\"".to_string());
438
439 let mut outcome = Outcome::said("setup-embeddings", lines.join("\n"));
440 outcome = outcome.with("root", json::text(root.display().to_string()));
441 for (name, value) in fields {
442 outcome = outcome.with(&name, value);
443 }
444 Ok(outcome)
445}
446
447fn folder_to_copy_from(
453 context: &mut Context,
454 arguments: &Arguments,
455 component: Component,
456) -> Result<Option<std::path::PathBuf>, Failed> {
457 let Some(named) = arguments.text("from") else {
458 return Ok(None);
459 };
460 if component == Component::Runtime {
461 return Err(Failed::misuse(
462 "--from installs the model or the reranker from a folder; the runtime is downloaded. \
463 Use 'setup-embeddings model --from <folder>', or leave out --from.",
464 ));
465 }
466 Ok(Some(context.confine(named)?))
467}
468
469struct MachineSettings {
471 threads: Option<usize>,
472 device: Option<String>,
473}
474
475impl MachineSettings {
476 fn any(&self) -> bool {
478 self.threads.is_some() || self.device.is_some()
479 }
480
481 fn record(&self, state: &mut install::State) {
485 if let Some(threads) = self.threads {
486 state.threads = Some(threads);
487 }
488 if let Some(device) = self.device.as_ref() {
489 state.device = Some(device.clone());
490 }
491 }
492
493 fn lines(&self) -> Vec<String> {
495 let mut lines = Vec::new();
496 if let Some(threads) = self.threads {
497 lines.push(format!("threads: {threads}"));
498 }
499 if let Some(device) = self.device.as_ref() {
500 lines.push(format!("device: {device}"));
501 }
502 lines
503 }
504}
505
506fn machine_settings(arguments: &Arguments) -> Result<MachineSettings, Failed> {
510 let threads = match arguments.integer("threads") {
511 Some(count) => Some(install::parse_threads(&count.to_string()).map_err(Failed::misuse)?),
512 None => None,
513 };
514 let device = match arguments.text("device") {
515 Some(text) => Some(install::parse_device(text).map_err(Failed::misuse)?),
516 None => None,
517 };
518 Ok(MachineSettings { threads, device })
519}
520
521fn setting_line(name: &str, value: &str, source: install::SettingSource, variable: &str) -> String {
528 format!(
529 "{name}: {value} ({}). {variable} overrides it for one process",
530 source.label()
531 )
532}
533
534fn status(root: &Path) -> Outcome {
542 let state = install::read_state(root).unwrap_or_default();
543 let mut lines = vec![format!("Install root: {}", root.display())];
544
545 match install::runtime_library() {
546 Some(library) => {
547 let recorded = state.runtime.as_ref();
548 let version = recorded
549 .map(|r| r.version.as_str())
550 .unwrap_or("unknown version");
551 let unverified = recorded.is_some_and(|r| !r.verified);
552 lines.push(format!(
553 "ONNX Runtime: {} ({version}){}",
554 library.display(),
555 if unverified {
556 " (digest not pinned in this build)"
557 } else {
558 ""
559 }
560 ));
561 }
562 None => lines.push(
563 "ONNX Runtime: not installed. Run: inillucent setup-embeddings runtime".to_string(),
564 ),
565 }
566
567 match install::model_dir(install::DEFAULT_MODEL) {
568 Some(dir) => lines.push(format!("{}: {}", install::DEFAULT_MODEL, dir.display())),
569 None => lines.push(format!(
570 "{}: not installed. Run: inillucent setup-embeddings model",
571 install::DEFAULT_MODEL
572 )),
573 }
574
575 let effective = Residency::configured();
576 lines.push(format!("Residency profile: {}", effective.label()));
577 lines.push(match effective {
578 Residency::Resident => {
579 " loaded on first use and kept, which is about 1.9 GB held and 12 to 36 ms a query"
580 .to_string()
581 }
582 Residency::OnDemand => {
583 " loaded per call and dropped, which is nothing held and about 0.8 s a query"
584 .to_string()
585 }
586 Residency::Idle(after) => format!(
587 " loaded on use and dropped after {}s idle: the first query in a burst pays about \
588 0.8 s and the rest pay 12 to 36 ms",
589 after.as_secs()
590 ),
591 });
592
593 let threads = install::configured_threads();
594 let device = install::configured_device();
595 let threads_text = threads
596 .value
597 .map_or("ONNX Runtime's own choice".to_string(), |count| {
598 count.to_string()
599 });
600 lines.push(setting_line(
601 "Threads",
602 &threads_text,
603 threads.source,
604 install::THREADS_VAR,
605 ));
606 lines.push(setting_line(
607 "Device",
608 &device.value,
609 device.source,
610 install::DEVICE_VAR,
611 ));
612
613 match install::model_dir(install::RERANKER_MODEL) {
614 Some(dir) => lines.push(format!("{}: {}", install::RERANKER_MODEL, dir.display())),
615 None => lines.push(format!(
616 "{}: not installed. Run: inillucent setup-embeddings reranker (about 600 MB, needed \
617 only for rerank() and a search that names question)",
618 install::RERANKER_MODEL
619 )),
620 }
621
622 let ready = install::runtime_library().is_some()
623 && install::model_dir(install::DEFAULT_MODEL).is_some();
624 Outcome::said("setup-embeddings", lines.join("\n"))
625 .with("root", json::text(root.display().to_string()))
626 .with("ready", Json::Bool(ready))
627 .with("residency", json::text(effective.label()))
628 .with(
629 "threads",
630 setting_json(threads.value.map(|count| count.to_string()), threads.source),
631 )
632 .with(
633 "device",
634 setting_json(Some(device.value.clone()), device.source),
635 )
636 .with(
637 "runtime",
638 match state.runtime.as_ref() {
639 Some(runtime) => runtime_json(runtime),
640 None => Json::Null,
641 },
642 )
643 .with(
644 "model",
645 match state.model(install::DEFAULT_MODEL) {
646 Some(model) => model_json(model),
647 None => Json::Null,
648 },
649 )
650 .with(
651 "reranker_installed",
652 Json::Bool(install::model_dir(install::RERANKER_MODEL).is_some()),
653 )
654 .with(
655 "reranker",
656 match state.model(install::RERANKER_MODEL) {
657 Some(model) => model_json(model),
658 None => Json::Null,
659 },
660 )
661}
662
663fn setting_json(value: Option<String>, source: install::SettingSource) -> Json {
668 json::object(vec![
669 ("value", value.map_or(Json::Null, json::text)),
670 ("source", json::text(source.label())),
671 ])
672}
673
674#[derive(Clone, Copy, Debug, Eq, PartialEq)]
688pub enum Accelerator {
689 Cpu,
691 Gpu,
693}
694
695#[derive(Clone, Copy, Debug, Eq, PartialEq)]
697pub enum Reinstall {
698 Always,
700 WhenMissing,
702}
703
704fn install_runtime(
705 root: &Path,
706 version: &str,
707 accelerator: Accelerator,
708 force: Reinstall,
709) -> Result<InstalledRuntime, Failed> {
710 let gpu = accelerator == Accelerator::Gpu;
711 let force = force == Reinstall::Always;
712 let archive_spec = pick_runtime(gpu)?;
713 let asset = archive_spec.asset.replace("{version}", version);
714 let directory = install::runtime_dir(root, version);
715 let library = directory.join("lib").join(install::runtime_library_name());
716
717 if library.exists() && !force {
718 return Ok(InstalledRuntime {
719 version: version.to_string(),
720 archive: asset,
721 library: library.display().to_string(),
722 verified: true,
723 gpu,
724 });
725 }
726
727 let expected = (version == DEFAULT_RUNTIME).then_some(archive_spec.sha256);
730 let url = format!("{RUNTIME_BASE}/v{version}/{asset}");
731 let downloaded = install::downloads_dir(root).join(&asset);
732 let mut progress = bar();
733 let fetched = http::download(&url, &downloaded, expected, &mut progress)
734 .map_err(|error| Failed::misuse(format!("{error}")))?;
735
736 let bytes = std::fs::read(&fetched.path).map_err(|error| {
737 Failed::misuse(format!(
738 "{} could not be read: {error}",
739 fetched.path.display()
740 ))
741 })?;
742 let members = archive::read(&bytes).map_err(|error| Failed::misuse(format!("{error}")))?;
743 let libraries = shared_libraries(&members);
744 if libraries.is_empty() {
745 return Err(Failed::misuse(format!(
746 "{asset} holds no {} - it is not an ONNX Runtime release, or its layout has changed",
747 install::runtime_library_name()
748 )));
749 }
750
751 let lib_dir = directory.join("lib");
752 std::fs::create_dir_all(&lib_dir).map_err(|error| {
753 Failed::misuse(format!(
754 "{} could not be created: {error}",
755 lib_dir.display()
756 ))
757 })?;
758 for (member, name) in &libraries {
759 archive::write_member(member, &lib_dir, Some(name))
760 .map_err(|error| Failed::misuse(format!("{error}")))?;
761 }
762 let _ = std::fs::remove_file(&fetched.path);
765
766 if !library.exists() {
767 return Err(Failed::misuse(format!(
768 "{asset} was extracted and {} is still not there",
769 library.display()
770 )));
771 }
772
773 Ok(InstalledRuntime {
774 version: version.to_string(),
775 archive: asset,
776 library: library.display().to_string(),
777 verified: expected.is_some(),
778 gpu,
779 })
780}
781
782fn pick_runtime(gpu: bool) -> Result<&'static RuntimeArchive, Failed> {
786 let os = std::env::consts::OS;
787 let arch = std::env::consts::ARCH;
788 RUNTIMES
789 .iter()
790 .find(|entry| entry.os == os && entry.arch == arch && entry.gpu == gpu)
791 .ok_or_else(|| {
792 let plain = RUNTIMES
793 .iter()
794 .any(|e| e.os == os && e.arch == arch && !e.gpu);
795 if gpu && plain {
796 Failed::misuse(format!(
797 "there is no CUDA build of ONNX Runtime for {os} on {arch}. Run the command \
798 without --gpu."
799 ))
800 } else {
801 Failed::misuse(format!(
802 "there is no ONNX Runtime release for {os} on {arch} that this command knows \
803 how to install. Build it, and point ORT_DYLIB_PATH at the result."
804 ))
805 }
806 })
807}
808
809fn shared_libraries(members: &[Member]) -> Vec<(&Member, String)> {
822 let wanted = install::runtime_library_name();
823 let mut out = Vec::new();
824 for member in members {
825 let name = member.name.rsplit('/').next().unwrap_or(&member.name);
826 if member.name.contains(".dSYM/") {
827 continue;
828 }
829 if !is_shared_library(name) {
830 continue;
831 }
832 if is_main_library(name) {
833 out.push((member, wanted.to_string()));
834 } else if name.contains("onnxruntime_providers") {
835 out.push((member, unversioned(name)));
839 }
840 }
841 out
842}
843
844fn is_shared_library(name: &str) -> bool {
848 name.ends_with(".dll") || name.ends_with(".dylib") || name.contains(".so")
849}
850
851fn is_main_library(name: &str) -> bool {
855 if name.contains("providers") {
856 return false;
857 }
858 name == "onnxruntime.dll"
859 || name.starts_with("libonnxruntime.so")
860 || (name.starts_with("libonnxruntime.") && name.ends_with(".dylib"))
861 || name == "libonnxruntime.dylib"
862}
863
864fn unversioned(name: &str) -> String {
872 if let Some(at) = name.find(".so") {
873 return format!("{}.so", name.get(..at).unwrap_or(name));
874 }
875 if name.ends_with(".dylib") {
876 let stem = name.trim_end_matches(".dylib");
877 let base = stem.split('.').next().unwrap_or(stem);
878 return format!("{base}.dylib");
879 }
880 name.to_string()
881}
882
883fn install_model(
890 root: &Path,
891 spec: &ModelSpec,
892 force: bool,
893 from: Option<&Path>,
894) -> Result<InstalledModel, Failed> {
895 let directory = install::models_root(root).join(spec.id);
896 std::fs::create_dir_all(&directory).map_err(|error| {
897 Failed::misuse(format!(
898 "{} could not be created: {error}",
899 directory.display()
900 ))
901 })?;
902
903 match from {
904 Some(folder) => copy_model_files(folder, &directory, spec, force)?,
905 None => {
906 for file in spec.files {
907 let destination = directory.join(file.local);
908 if destination.exists() && !force && already_correct(&destination, file) {
909 continue;
910 }
911 let url = format!("{}/{}", spec.base, file.remote);
912 let mut progress = bar();
913 http::download(&url, &destination, Some(file.sha256), &mut progress)
914 .map_err(|error| Failed::misuse(format!("{error}")))?;
915 }
916 }
917 }
918
919 let mut manifest = (spec.manifest)();
924 manifest.weights_sha256 = spec
925 .files
926 .iter()
927 .find(|f| f.local == "model.onnx")
928 .map(|f| f.sha256.to_string())
929 .unwrap_or_default();
930 manifest.tokenizer_sha256 = spec
931 .files
932 .iter()
933 .find(|f| f.local == "tokenizer.json")
934 .map(|f| f.sha256.to_string())
935 .unwrap_or_default();
936 manifest.source = Some(spec.source.to_string());
937 manifest
938 .write(&directory)
939 .map_err(|reason| Failed::misuse(format!("the model manifest: {reason:#}")))?;
940
941 Ok(InstalledModel {
942 id: spec.id.to_string(),
943 dir: directory.display().to_string(),
944 source: spec.source.to_string(),
945 verified: true,
946 })
947}
948
949fn copy_model_files(
963 folder: &Path,
964 directory: &Path,
965 spec: &ModelSpec,
966 force: bool,
967) -> Result<(), Failed> {
968 let mut found = Vec::with_capacity(spec.files.len());
969 for file in spec.files {
970 let candidates = [folder.join(file.local), folder.join(file.remote)];
971 let Some(source) = candidates.iter().find(|path| path.is_file()) else {
972 return Err(Failed::misuse(format!(
973 "{} is not in {} (looked for {} and {}). Nothing was installed",
974 file.local,
975 folder.display(),
976 file.local,
977 file.remote
978 )));
979 };
980 let actual = sha256_of(source).ok_or_else(|| {
981 Failed::misuse(format!(
982 "{} could not be read. Nothing was installed",
983 source.display()
984 ))
985 })?;
986 if !actual.eq_ignore_ascii_case(file.sha256) {
987 return Err(Failed::misuse(format!(
988 "{} is not the file this build installs: its SHA-256 is {actual}, and {} \
989 needs {}. Nothing was installed",
990 source.display(),
991 spec.id,
992 file.sha256
993 )));
994 }
995 found.push((file, source.clone()));
996 }
997 for (file, source) in found {
998 let destination = directory.join(file.local);
999 if destination.exists() && !force && already_correct(&destination, file) {
1000 continue;
1001 }
1002 let partial = directory.join(format!("{}.partial", file.local));
1005 std::fs::copy(&source, &partial)
1006 .and_then(|_| std::fs::rename(&partial, &destination))
1007 .map_err(|error| {
1008 let _ = std::fs::remove_file(&partial);
1009 Failed::misuse(format!(
1010 "{} could not be copied to {}: {error}",
1011 source.display(),
1012 destination.display()
1013 ))
1014 })?;
1015 }
1016 Ok(())
1017}
1018
1019fn sha256_of(path: &Path) -> Option<String> {
1023 let mut handle = std::fs::File::open(path).ok()?;
1024 let mut digest = inillucent_base::hash::Sha256::new();
1025 let mut buffer = vec![0u8; 1 << 20];
1026 loop {
1027 use std::io::Read;
1028 match handle.read(&mut buffer) {
1029 Ok(0) => break,
1030 Ok(read) => digest.update(buffer.get(..read).unwrap_or(&[])),
1031 Err(_) => return None,
1032 }
1033 }
1034 Some(inillucent_base::hash::to_hex(&digest.finish()))
1035}
1036
1037fn already_correct(path: &Path, file: &ModelFile) -> bool {
1047 let Ok(metadata) = std::fs::metadata(path) else {
1048 return false;
1049 };
1050 if metadata.len() != file.bytes {
1051 return false;
1052 }
1053 let Ok(mut handle) = std::fs::File::open(path) else {
1054 return false;
1055 };
1056 let mut digest = inillucent_base::hash::Sha256::new();
1057 let mut buffer = vec![0u8; 1 << 20];
1058 loop {
1059 use std::io::Read;
1060 match handle.read(&mut buffer) {
1061 Ok(0) => break,
1062 Ok(read) => digest.update(buffer.get(..read).unwrap_or(&[])),
1063 Err(_) => return false,
1064 }
1065 }
1066 inillucent_base::hash::to_hex(&digest.finish()).eq_ignore_ascii_case(file.sha256)
1067}
1068
1069fn runtime_json(runtime: &InstalledRuntime) -> Json {
1073 json::object(vec![
1074 ("version", json::text(&runtime.version)),
1075 ("archive", json::text(&runtime.archive)),
1076 ("library", json::text(&runtime.library)),
1077 ("verified", Json::Bool(runtime.verified)),
1078 ("gpu", Json::Bool(runtime.gpu)),
1079 ])
1080}
1081
1082fn model_json(model: &InstalledModel) -> Json {
1086 json::object(vec![
1087 ("id", json::text(&model.id)),
1088 ("dir", json::text(&model.dir)),
1089 ("source", json::text(&model.source)),
1090 ("verified", Json::Bool(model.verified)),
1091 ])
1092}
1093
1094fn bar() -> Bar {
1096 Bar {
1097 terminal: std::io::stderr().is_terminal(),
1098 name: String::new(),
1099 width: 0,
1100 last_percent: -1,
1101 started: std::time::Instant::now(),
1102 }
1103}
1104
1105struct Bar {
1114 terminal: bool,
1115 name: String,
1116 width: u64,
1118 last_percent: i64,
1120 started: std::time::Instant,
1121}
1122
1123impl Bar {
1124 fn draw(&self, done: u64) {
1128 let elapsed = self.started.elapsed().as_secs_f64().max(0.001);
1129 let rate = done as f64 / elapsed / (1024.0 * 1024.0);
1130 let mut line = if self.width > 0 {
1131 let share = (done as f64 / self.width as f64).clamp(0.0, 1.0);
1132 let filled = (share * 20.0).round() as usize;
1133 format!(
1134 " {:<38} [{}{}] {:3.0}% {:.1}/{:.1} MB {rate:.1} MB/s",
1135 self.name,
1136 "#".repeat(filled.min(20)),
1137 ".".repeat(20usize.saturating_sub(filled)),
1138 share * 100.0,
1139 done as f64 / (1024.0 * 1024.0),
1140 self.width as f64 / (1024.0 * 1024.0),
1141 )
1142 } else {
1143 format!(
1144 " {:<38} {:.1} MB {rate:.1} MB/s",
1145 self.name,
1146 done as f64 / (1024.0 * 1024.0)
1147 )
1148 };
1149 line.push('\r');
1150 let mut stderr = std::io::stderr();
1151 let _ = stderr.write_all(line.as_bytes());
1152 let _ = stderr.flush();
1153 }
1154}
1155
1156impl Progress for Bar {
1157 fn started(&mut self, name: &str, total: Option<u64>, resumed: u64) {
1159 self.name = name.to_string();
1160 self.width = total.unwrap_or(0);
1161 self.last_percent = -1;
1162 self.started = std::time::Instant::now();
1163 if resumed > 0 {
1164 eprintln!(
1165 " {name}: resuming at {:.1} MB",
1166 resumed as f64 / (1024.0 * 1024.0)
1167 );
1168 }
1169 }
1170
1171 fn advanced(&mut self, done: u64, _total: Option<u64>) {
1173 if self.terminal {
1174 self.draw(done);
1175 return;
1176 }
1177 if self.width == 0 {
1178 return;
1179 }
1180 let percent = (done as i64).saturating_mul(100) / self.width.max(1) as i64;
1181 if percent / 10 > self.last_percent / 10 {
1182 self.last_percent = percent;
1183 eprintln!(" {}: {percent}%", self.name);
1184 }
1185 }
1186
1187 fn finished(&mut self, done: u64) {
1189 if self.terminal {
1190 self.draw(done);
1191 }
1192 eprintln!(
1193 " {:<38} {:.1} MB",
1194 self.name,
1195 done as f64 / (1024.0 * 1024.0)
1196 );
1197 }
1198}
1199
1200#[cfg(test)]
1201mod tests {
1202 use super::*;
1203
1204 const FAKE_FILES: &[ModelFile] = &[
1207 ModelFile {
1208 remote: "a.txt",
1209 local: "a.txt",
1210 sha256: "2cf24dba5fb0a30e26e83b2ac5b9e29e1b161e5c1fa7425e73043362938b9824",
1211 bytes: 5,
1212 },
1213 ModelFile {
1214 remote: "onnx/b.bin",
1215 local: "b.bin",
1216 sha256: "486ea46224d1bb4fb680f34f7c9ad96a8f24ec88be73ea8e5a6c65260e9cb8a7",
1217 bytes: 5,
1218 },
1219 ];
1220
1221 const FAKE_MODEL: ModelSpec = ModelSpec {
1223 id: "fake-model",
1224 base: "",
1225 source: "",
1226 files: FAKE_FILES,
1227 manifest: EMBEDDER.manifest,
1228 };
1229
1230 fn scratch(name: &str) -> std::path::PathBuf {
1234 let path = std::env::temp_dir()
1235 .join(format!("inillucent-setup-from-{}", std::process::id()))
1236 .join(name);
1237 let _ = std::fs::remove_dir_all(&path);
1238 std::fs::create_dir_all(&path).unwrap();
1239 path
1240 }
1241
1242 #[test]
1247 fn a_model_is_installed_from_a_folder_only_when_every_file_matches() {
1248 let folder = scratch("good");
1249 std::fs::write(folder.join("a.txt"), b"hello").unwrap();
1250 std::fs::create_dir_all(folder.join("onnx")).unwrap();
1251 std::fs::write(folder.join("onnx").join("b.bin"), b"world").unwrap();
1252 let installed = scratch("installed");
1253 copy_model_files(&folder, &installed, &FAKE_MODEL, false).unwrap();
1254 assert_eq!(std::fs::read(installed.join("a.txt")).unwrap(), b"hello");
1255 assert_eq!(std::fs::read(installed.join("b.bin")).unwrap(), b"world");
1256
1257 let wrong = scratch("wrong");
1258 std::fs::write(wrong.join("a.txt"), b"hullo").unwrap();
1259 std::fs::write(wrong.join("b.bin"), b"world").unwrap();
1260 let untouched = scratch("untouched");
1261 let refused = copy_model_files(&wrong, &untouched, &FAKE_MODEL, false);
1262 let said = format!("{:?}", refused.err());
1263 assert!(said.contains("Nothing was installed"), "{said}");
1264 assert!(
1265 said.contains("2cf24dba"),
1266 "the refusal names the pinned digest: {said}"
1267 );
1268 assert!(
1269 !untouched.join("b.bin").exists(),
1270 "a good file was copied beside a bad one"
1271 );
1272
1273 let missing = scratch("missing");
1274 std::fs::write(missing.join("a.txt"), b"hello").unwrap();
1275 let refused = copy_model_files(&missing, &untouched, &FAKE_MODEL, false);
1276 let said = format!("{:?}", refused.err());
1277 assert!(said.contains("b.bin is not in"), "{said}");
1278 assert!(
1279 !untouched.join("a.txt").exists(),
1280 "a file was copied while another was missing"
1281 );
1282 }
1283
1284 #[test]
1291 fn every_platform_has_an_archive_with_a_pinned_digest() {
1292 for (os, arch) in [
1293 ("windows", "x86_64"),
1294 ("windows", "aarch64"),
1295 ("macos", "x86_64"),
1296 ("macos", "aarch64"),
1297 ("linux", "x86_64"),
1298 ("linux", "aarch64"),
1299 ] {
1300 let found = RUNTIMES
1301 .iter()
1302 .find(|entry| entry.os == os && entry.arch == arch && !entry.gpu);
1303 assert!(found.is_some(), "{os} on {arch} has no archive");
1304 }
1305 for entry in RUNTIMES {
1306 assert_eq!(entry.sha256.len(), 64, "{} has a short digest", entry.asset);
1307 assert!(
1308 entry.sha256.chars().all(|c| c.is_ascii_hexdigit()),
1309 "{} has a digest that is not hex",
1310 entry.asset
1311 );
1312 assert!(
1313 entry.asset.contains("{version}"),
1314 "{} has no version slot",
1315 entry.asset
1316 );
1317 }
1318 }
1319
1320 #[test]
1323 fn every_model_file_has_a_digest_and_a_size() {
1324 for file in NOMIC_FILES {
1325 assert_eq!(file.sha256.len(), 64, "{} has a short digest", file.local);
1326 assert!(file.bytes > 0, "{} has no size", file.local);
1327 }
1328 assert!(NOMIC_FILES.iter().any(|f| f.local == "model.onnx"));
1329 assert!(NOMIC_FILES.iter().any(|f| f.local == "tokenizer.json"));
1330 }
1331
1332 #[test]
1339 fn the_installed_weights_are_the_full_precision_export() {
1340 let weights = NOMIC_FILES
1341 .iter()
1342 .find(|f| f.local == "model.onnx")
1343 .expect("the weights are in the list");
1344 assert_eq!(weights.remote, "onnx/model.onnx");
1345 assert_eq!(weights.bytes, 547_310_275, "the fp32 export's size");
1346 }
1347
1348 #[test]
1350 fn a_component_parses_or_is_refused() {
1351 assert_eq!(Component::parse("all").unwrap(), Component::All);
1352 assert_eq!(Component::parse("").unwrap(), Component::All);
1353 assert_eq!(Component::parse(" RUNTIME ").unwrap(), Component::Runtime);
1354 assert_eq!(Component::parse("model").unwrap(), Component::Model);
1355 assert!(Component::parse("everything").is_err());
1356 }
1357
1358 #[test]
1361 fn the_main_library_is_told_apart_from_its_providers() {
1362 assert!(is_main_library("onnxruntime.dll"));
1363 assert!(is_main_library("libonnxruntime.so.1.22.0"));
1364 assert!(is_main_library("libonnxruntime.1.22.0.dylib"));
1365 assert!(!is_main_library("onnxruntime_providers_cuda.dll"));
1366 assert!(!is_main_library("libonnxruntime_providers_shared.so"));
1367 assert!(!is_shared_library("onnxruntime.lib"));
1368 assert!(!is_shared_library("onnxruntime.pdb"));
1369 assert!(!is_shared_library("libonnxruntime.pc"));
1370 }
1371
1372 #[test]
1374 fn a_versioned_library_loses_its_version() {
1375 assert_eq!(
1376 unversioned("libonnxruntime_providers_cuda.so.1.22.0"),
1377 "libonnxruntime_providers_cuda.so"
1378 );
1379 assert_eq!(
1380 unversioned("libonnxruntime_providers_shared.so"),
1381 "libonnxruntime_providers_shared.so"
1382 );
1383 assert_eq!(
1384 unversioned("onnxruntime_providers_cuda.dll"),
1385 "onnxruntime_providers_cuda.dll"
1386 );
1387 assert_eq!(
1388 unversioned("libonnxruntime_providers_cuda.1.22.0.dylib"),
1389 "libonnxruntime_providers_cuda.dylib"
1390 );
1391 }
1392
1393 #[test]
1399 fn each_archive_layout_yields_the_library_the_loader_wants() {
1400 let members: Vec<Member> = [
1401 "onnxruntime-win-x64-1.22.0/lib/onnxruntime.dll",
1402 "onnxruntime-win-x64-1.22.0/lib/onnxruntime.lib",
1403 "onnxruntime-win-x64-1.22.0/lib/onnxruntime.pdb",
1404 "onnxruntime-linux-x64-1.22.0/lib/libonnxruntime.so.1.22.0",
1405 "onnxruntime-linux-x64-1.22.0/lib/libonnxruntime.pc",
1406 "onnxruntime-osx-universal2-1.22.0/lib/libonnxruntime.1.22.0.dylib",
1407 "onnxruntime-osx-universal2-1.22.0/lib/libonnxruntime.1.22.0.dylib.dSYM/Contents/Resources/DWARF/libonnxruntime.1.22.0.dylib",
1408 ]
1409 .into_iter()
1410 .map(|name| Member { name: name.to_string(), bytes: Vec::new(), executable: false })
1411 .collect();
1412
1413 let picked = shared_libraries(&members);
1414 let names: Vec<&str> = picked.iter().map(|(_, name)| name.as_str()).collect();
1415 assert!(
1416 names
1417 .iter()
1418 .all(|name| *name == install::runtime_library_name()),
1419 "every main library is installed under this platform's name: {names:?}"
1420 );
1421 assert!(
1422 !picked
1423 .iter()
1424 .any(|(member, _)| member.name.contains(".dSYM/")),
1425 "the macOS debug bundle is not a library"
1426 );
1427 assert_eq!(picked.len(), 3, "{names:?}");
1430 }
1431
1432 #[test]
1435 fn a_provider_library_keeps_its_own_name() {
1436 let members = vec![
1437 Member {
1438 name: "onnxruntime-win-x64-gpu-1.22.0/lib/onnxruntime.dll".to_string(),
1439 bytes: Vec::new(),
1440 executable: false,
1441 },
1442 Member {
1443 name: "onnxruntime-win-x64-gpu-1.22.0/lib/onnxruntime_providers_cuda.dll"
1444 .to_string(),
1445 bytes: Vec::new(),
1446 executable: false,
1447 },
1448 ];
1449 let picked = shared_libraries(&members);
1450 assert_eq!(picked.len(), 2);
1451 assert!(picked
1452 .iter()
1453 .any(|(_, name)| name == "onnxruntime_providers_cuda.dll"));
1454 }
1455}