use std::path::{Path, PathBuf};
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum Platform {
MacosArm64,
Standard,
}
impl Platform {
#[must_use]
pub fn host() -> Self {
if cfg!(all(target_os = "macos", target_arch = "aarch64")) {
Self::MacosArm64
} else {
Self::Standard
}
}
#[must_use]
pub fn as_str(self) -> &'static str {
match self {
Self::MacosArm64 => "macos-arm64",
Self::Standard => "standard",
}
}
}
#[derive(Debug, Clone, Copy)]
pub struct ModelFile {
pub name: &'static str,
pub url: &'static str,
pub sha256: &'static str,
}
#[derive(Debug, Clone, Copy)]
pub struct ModelVariant {
pub platform: Platform,
pub files: &'static [ModelFile],
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum ModelKind {
Embedding,
Generative,
Ocr,
Vision,
Audio,
}
impl ModelKind {
pub const ALL: [Self; 5] = [
Self::Embedding,
Self::Generative,
Self::Ocr,
Self::Vision,
Self::Audio,
];
#[must_use]
pub fn as_str(self) -> &'static str {
match self {
Self::Embedding => "embedding",
Self::Generative => "generative",
Self::Ocr => "ocr",
Self::Vision => "vision",
Self::Audio => "audio",
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum ResourceTier {
Low,
Mid,
High,
}
impl ResourceTier {
#[must_use]
pub fn as_str(self) -> &'static str {
match self {
Self::Low => "low",
Self::Mid => "mid",
Self::High => "high",
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum ModelRole {
None,
Instruct,
Coding,
Reasoning,
}
impl ModelRole {
#[must_use]
pub fn as_str(self) -> Option<&'static str> {
match self {
Self::None => None,
Self::Instruct => Some("instruct"),
Self::Coding => Some("coding"),
Self::Reasoning => Some("reasoning"),
}
}
}
#[derive(Debug, Clone, Copy)]
pub struct ModelSpec {
pub name: &'static str,
pub kind: ModelKind,
pub role: ModelRole,
pub tier: ResourceTier,
pub dim: usize,
pub licence: &'static str,
pub description: &'static str,
pub size_mib: u32,
pub variants: &'static [ModelVariant],
}
impl ModelSpec {
#[must_use]
pub fn variant_for(&self, platform: Platform) -> Option<&ModelVariant> {
self.variants
.iter()
.find(|v| v.platform == platform)
.or_else(|| {
self.variants
.iter()
.find(|v| v.platform == Platform::Standard)
})
}
}
pub const REGISTRY: &[ModelSpec] = &[
ModelSpec {
name: "bge-base-en-v1.5",
kind: ModelKind::Embedding,
role: ModelRole::None,
tier: ResourceTier::Mid,
dim: 768,
licence: "MIT",
description: "BAAI/bge-base-en-v1.5 (F16 GGUF) — stronger English embeddings (768-d)",
size_mib: 209,
variants: &[ModelVariant {
platform: Platform::Standard,
files: &[ModelFile {
name: "model.gguf",
url: "https://huggingface.co/CompendiumLabs/bge-base-en-v1.5-gguf/resolve/main/bge-base-en-v1.5-f16.gguf",
sha256: "88360fdf8521af0ac08d43818bd272da679ab97c685d9b273c48efd01a4187c2",
}],
}],
},
ModelSpec {
name: "bge-large-en-v1.5",
kind: ModelKind::Embedding,
role: ModelRole::None,
tier: ResourceTier::High,
dim: 1024,
licence: "MIT",
description: "BAAI/bge-large-en-v1.5 (F16 GGUF) — strongest English embeddings (1024-d)",
size_mib: 639,
variants: &[ModelVariant {
platform: Platform::Standard,
files: &[ModelFile {
name: "model.gguf",
url: "https://huggingface.co/CompendiumLabs/bge-large-en-v1.5-gguf/resolve/main/bge-large-en-v1.5-f16.gguf",
sha256: "3379a0e9cea28fc6d7136df8ea7a88ef99ccce5963b9a6f7af9609997be762e3",
}],
}],
},
ModelSpec {
name: "bge-small-en-v1.5-gguf",
kind: ModelKind::Embedding,
role: ModelRole::None,
tier: ResourceTier::Low,
dim: 384,
licence: "MIT",
description: "BAAI/bge-small-en-v1.5 (F16 GGUF) — small English embeddings (384-d), served via llama.cpp",
size_mib: 65,
variants: &[ModelVariant {
platform: Platform::Standard,
files: &[ModelFile {
name: "model.gguf",
url: "https://huggingface.co/CompendiumLabs/bge-small-en-v1.5-gguf/resolve/main/bge-small-en-v1.5-f16.gguf",
sha256: "f0b2fef971e8366438bfd2d9aefea1b0115919389448806d290237f638bae999",
}],
}],
},
ModelSpec {
name: "qwen3-0.6b",
kind: ModelKind::Generative,
role: ModelRole::Instruct,
tier: ResourceTier::Low,
dim: 0,
licence: "Apache-2.0",
description: "Qwen3-0.6B (Q4_K_M GGUF) — tiny offline instruct model, the `spec draft` default",
size_mib: 380,
variants: &[ModelVariant {
platform: Platform::Standard,
files: &[ModelFile {
name: "model.gguf",
url: "https://huggingface.co/unsloth/Qwen3-0.6B-GGUF/resolve/main/Qwen3-0.6B-Q4_K_M.gguf",
sha256: "ac2d97712095a558e31573f62f466a3f9d93990898b0ec79d7c974c1780d524a",
}],
}],
},
ModelSpec {
name: "qwen3-8b",
kind: ModelKind::Generative,
role: ModelRole::Instruct,
tier: ResourceTier::Mid,
dim: 0,
licence: "Apache-2.0",
description: "Qwen3-8B (Q4_K_M GGUF) — stronger offline drafting on a ~16 GB machine",
size_mib: 4795,
variants: &[ModelVariant {
platform: Platform::Standard,
files: &[ModelFile {
name: "model.gguf",
url: "https://huggingface.co/Qwen/Qwen3-8B-GGUF/resolve/main/Qwen3-8B-Q4_K_M.gguf",
sha256: "d98cdcbd03e17ce47681435b5150e34c1417f50b5c0019dd560e4882c5745785",
}],
}],
},
ModelSpec {
name: "qwen3-32b",
kind: ModelKind::Generative,
role: ModelRole::Instruct,
tier: ResourceTier::High,
dim: 0,
licence: "Apache-2.0",
description: "Qwen3-32B (Q4_K_M GGUF) — best offline drafting, for a workstation",
size_mib: 18845,
variants: &[ModelVariant {
platform: Platform::Standard,
files: &[ModelFile {
name: "model.gguf",
url: "https://huggingface.co/Qwen/Qwen3-32B-GGUF/resolve/main/Qwen3-32B-Q4_K_M.gguf",
sha256: "efd971561896866f0e910cce52761ca77b1b138090c7f15fe284676d57d1f689",
}],
}],
},
ModelSpec {
name: "qwen3.8-27b",
kind: ModelKind::Generative,
role: ModelRole::Instruct,
tier: ResourceTier::High,
dim: 0,
licence: "Apache-2.0",
description: "Qwen3.8-27B (Q4_K_M GGUF) — strongest offline instruct pick, tool-calling, for a workstation",
size_mib: 18095,
variants: &[ModelVariant {
platform: Platform::Standard,
files: &[ModelFile {
name: "model.gguf",
url: "https://huggingface.co/ggml-org/Qwen3.8-27B-GGUF/resolve/main/Qwen3.8-27B-Q4_K_M.gguf",
sha256: "31629f53165ab6a7dad8c9847dcfd1fdf55829dac1e6e748f4a68581b0033d34",
}],
}],
},
ModelSpec {
name: "qwen2.5-coder-3b",
kind: ModelKind::Generative,
role: ModelRole::Coding,
tier: ResourceTier::Mid,
dim: 0,
licence: "Apache-2.0",
description: "Qwen2.5-Coder-3B-Instruct (Q4_K_M GGUF) — code completion/Q&A, served via llama.cpp",
size_mib: 1841,
variants: &[ModelVariant {
platform: Platform::Standard,
files: &[ModelFile {
name: "model.gguf",
url: "https://huggingface.co/bartowski/Qwen2.5-Coder-3B-Instruct-GGUF/resolve/main/Qwen2.5-Coder-3B-Instruct-Q4_K_M.gguf",
sha256: "3da3afe6cf5c674ac195803ea0dd6fee7e1c228c2105c1ce8c66890d1d4ab460",
}],
}],
},
ModelSpec {
name: "qwen3-coder-30b-a3b",
kind: ModelKind::Generative,
role: ModelRole::Coding,
tier: ResourceTier::High,
dim: 0,
licence: "Apache-2.0",
description: "Qwen3-Coder-30B-A3B-Instruct (Q4_K_M GGUF) — 30B MoE coder (3B active), for a workstation",
size_mib: 17697,
variants: &[ModelVariant {
platform: Platform::Standard,
files: &[ModelFile {
name: "model.gguf",
url: "https://huggingface.co/unsloth/Qwen3-Coder-30B-A3B-Instruct-GGUF/resolve/main/Qwen3-Coder-30B-A3B-Instruct-Q4_K_M.gguf",
sha256: "fadc3e5f8d42bf7e894a785b05082e47daee4df26680389817e2093056f088ad",
}],
}],
},
ModelSpec {
name: "deepseek-r1-distill-qwen-1.5b",
kind: ModelKind::Generative,
role: ModelRole::Reasoning,
tier: ResourceTier::Low,
dim: 0,
licence: "MIT",
description: "DeepSeek-R1-Distill-Qwen-1.5B (Q4_K_M GGUF) — small reasoning model, served via llama.cpp",
size_mib: 1066,
variants: &[ModelVariant {
platform: Platform::Standard,
files: &[ModelFile {
name: "model.gguf",
url: "https://huggingface.co/bartowski/DeepSeek-R1-Distill-Qwen-1.5B-GGUF/resolve/main/DeepSeek-R1-Distill-Qwen-1.5B-Q4_K_M.gguf",
sha256: "1741e5b2d062b07acf048bf0d2c514dadf2a48f94e2b4aa0cfe069af3838ee2f",
}],
}],
},
ModelSpec {
name: "ocrs-text",
kind: ModelKind::Ocr,
role: ModelRole::None,
tier: ResourceTier::Low,
dim: 0,
licence: "CC-BY-SA-4.0",
description: "ocrs text detection + recognition (pure-Rust OCR for `image-ocr`)",
size_mib: 12,
variants: &[ModelVariant {
platform: Platform::Standard,
files: &[
ModelFile {
name: "text-detection.rten",
url: "https://ocrs-models.s3-accelerate.amazonaws.com/text-detection.rten",
sha256: "f15cfb56bd02c4bf478a20343986504a1f01e1665c2b3a0ad66340f054b1b5ca",
},
ModelFile {
name: "text-recognition.rten",
url: "https://ocrs-models.s3-accelerate.amazonaws.com/text-recognition.rten",
sha256: "e484866d4cce403175bd8d00b128feb08ab42e208de30e42cd9889d8f1735a6e",
},
],
}],
},
ModelSpec {
name: "smolvlm-500m-gguf",
kind: ModelKind::Vision,
role: ModelRole::None,
tier: ResourceTier::Low,
dim: 0,
licence: "Apache-2.0",
description: "SmolVLM-500M-Instruct (Q8_0 GGUF + mmproj) — small vision-language model served via llama.cpp",
size_mib: 520,
variants: &[ModelVariant {
platform: Platform::Standard,
files: &[
ModelFile {
name: "model.gguf",
url: "https://huggingface.co/ggml-org/SmolVLM-500M-Instruct-GGUF/resolve/main/SmolVLM-500M-Instruct-Q8_0.gguf",
sha256: "9d4612de6a42214499e301494a3ecc2be0abdd9de44e663bda63f1152fad1bf4",
},
ModelFile {
name: "mmproj.gguf",
url: "https://huggingface.co/ggml-org/SmolVLM-500M-Instruct-GGUF/resolve/main/mmproj-SmolVLM-500M-Instruct-Q8_0.gguf",
sha256: "d1eb8b6b23979205fdf63703ed10f788131a3f812c7b1f72e0119d5d81295150",
},
],
}],
},
ModelSpec {
name: "voxtral-mini-3b",
kind: ModelKind::Audio,
role: ModelRole::None,
tier: ResourceTier::Mid,
dim: 0,
licence: "Apache-2.0",
description: "Voxtral-Mini-3B (Mistral, Q4_K_M GGUF + Q8_0 audio mmproj) — speech transcription via llama.cpp mtmd",
size_mib: 3041,
variants: &[ModelVariant {
platform: Platform::Standard,
files: &[
ModelFile {
name: "model.gguf",
url: "https://huggingface.co/ggml-org/Voxtral-Mini-3B-2507-GGUF/resolve/main/Voxtral-Mini-3B-2507-Q4_K_M.gguf",
sha256: "4705be8ec22ca23d12632f4b4a3691faa95917d90a06d3cf3c3ec0e91958f1a8",
},
ModelFile {
name: "mmproj.gguf",
url: "https://huggingface.co/ggml-org/Voxtral-Mini-3B-2507-GGUF/resolve/main/mmproj-Voxtral-Mini-3B-2507-Q8_0.gguf",
sha256: "4f24c4ef3ce929d02ed9d1cfb050ae9a7365f057c0ddec0d489580982ebe0d02",
},
],
}],
},
];
#[must_use]
pub fn find(name: &str) -> Option<&'static ModelSpec> {
REGISTRY.iter().find(|m| m.name == name)
}
fn store_root_from(
model_store: Option<PathBuf>,
roteiro_home: Option<PathBuf>,
home: Option<PathBuf>,
) -> PathBuf {
if let Some(dir) = model_store {
return dir;
}
if let Some(dir) = roteiro_home {
return dir.join("models");
}
home.unwrap_or_else(|| PathBuf::from("."))
.join(".roteiro")
.join("models")
}
static MODEL_STORE_OVERRIDE: std::sync::OnceLock<PathBuf> = std::sync::OnceLock::new();
pub fn set_model_store(dir: PathBuf) {
let _ = MODEL_STORE_OVERRIDE.set(dir);
}
#[must_use]
pub fn store_root() -> PathBuf {
if let Some(dir) = MODEL_STORE_OVERRIDE.get() {
return dir.clone();
}
store_root_from(
std::env::var_os("ROTEIRO_MODEL_STORE").map(PathBuf::from),
std::env::var_os("ROTEIRO_HOME").map(PathBuf::from),
std::env::var_os("HOME")
.or_else(|| std::env::var_os("USERPROFILE"))
.map(PathBuf::from),
)
}
#[must_use]
pub fn model_dir(name: &str) -> PathBuf {
store_root().join(name)
}
#[must_use]
pub fn is_installed(name: &str, variant: &ModelVariant) -> bool {
let dir = model_dir(name);
variant.files.iter().all(|f| dir.join(f.name).exists())
}
#[must_use]
pub fn installed_size(name: &str) -> u64 {
dir_size(&model_dir(name))
}
fn dir_size(dir: &Path) -> u64 {
let Ok(entries) = std::fs::read_dir(dir) else {
return 0;
};
entries
.flatten()
.map(|e| match e.file_type() {
Ok(t) if t.is_dir() => dir_size(&e.path()),
Ok(_) => e.metadata().map_or(0, |m| m.len()),
Err(_) => 0,
})
.sum()
}
#[derive(Debug, Clone)]
pub struct Removal {
pub dir: PathBuf,
pub files: Vec<String>,
pub bytes: u64,
}
pub fn remove_model(name: &str) -> std::io::Result<Removal> {
let dir = model_dir(name);
let mut files = Vec::new();
let mut bytes = 0;
if let Ok(entries) = std::fs::read_dir(&dir) {
for entry in entries.flatten() {
bytes += match entry.file_type() {
Ok(t) if t.is_dir() => dir_size(&entry.path()),
_ => entry.metadata().map_or(0, |m| m.len()),
};
files.push(entry.file_name().to_string_lossy().into_owned());
}
}
files.sort();
if dir.exists() {
std::fs::remove_dir_all(&dir)?;
}
Ok(Removal { dir, files, bytes })
}
#[must_use]
pub fn sha256_hex(bytes: &[u8]) -> String {
use sha2::{Digest, Sha256};
let digest = Sha256::digest(bytes);
let mut out = String::with_capacity(64);
for byte in digest {
use std::fmt::Write as _;
let _ = write!(out, "{byte:02x}");
}
out
}
#[must_use]
pub fn verify_sha256(bytes: &[u8], expected: &str) -> bool {
expected.is_empty() || sha256_hex(bytes).eq_ignore_ascii_case(expected)
}
pub fn ensure_model_dir(name: &str) -> std::io::Result<PathBuf> {
let dir = model_dir(name);
std::fs::create_dir_all(&dir)?;
Ok(dir)
}
#[derive(Debug, thiserror::Error)]
pub enum DownloadError {
#[error("download io error: {0}")]
Io(#[from] std::io::Error),
#[error("checksum mismatch: expected {expected}, got {got} (partial discarded)")]
Checksum {
expected: String,
got: String,
},
#[error("transport error: {0}")]
Transport(Box<dyn std::error::Error + Send + Sync>),
#[error("range request failed: {0}")]
Range(String),
}
#[derive(Debug)]
pub enum RangeReply<R> {
Partial {
reader: R,
total: Option<u64>,
},
Full {
reader: R,
total: Option<u64>,
detail: String,
},
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum DownloadEvent {
DiscardedPartial {
bytes: u64,
reason: String,
},
Resuming {
offset: u64,
total: Option<u64>,
},
AlreadyComplete {
bytes: u64,
},
RangeUnsupported {
discarded: u64,
detail: String,
},
KeptPartial {
bytes: u64,
},
PoisonedPartial {
bytes: u64,
},
}
#[derive(serde::Serialize, serde::Deserialize)]
struct PartialMeta {
version: u32,
url: String,
sha256: String,
total: Option<u64>,
}
const PARTIAL_META_VERSION: u32 = 1;
#[must_use]
pub fn partial_path(dest: &Path) -> PathBuf {
dest.with_extension("partial")
}
#[must_use]
pub fn partial_meta_path(dest: &Path) -> PathBuf {
dest.with_extension("partial.json")
}
pub fn discard_partial(dest: &Path) -> std::io::Result<u64> {
let tmp = partial_path(dest);
let freed = std::fs::metadata(&tmp).map_or(0, |m| m.len());
remove_if_present(&tmp)?;
remove_if_present(&partial_meta_path(dest))?;
Ok(freed)
}
fn remove_if_present(path: &Path) -> std::io::Result<()> {
match std::fs::remove_file(path) {
Err(e) if e.kind() == std::io::ErrorKind::NotFound => Ok(()),
other => other,
}
}
pub fn interpret_range_response(
status: u16,
accept_ranges: Option<&str>,
content_range: Option<&str>,
content_length: Option<u64>,
requested_from: u64,
) -> Result<(RangeKind, Option<u64>), DownloadError> {
match status {
206 => {
let raw = content_range
.ok_or_else(|| DownloadError::Range("206 response without Content-Range".into()))?;
let (start, total) = parse_content_range(raw)?;
if start != requested_from {
return Err(DownloadError::Range(format!(
"server resumed at byte {start} but {requested_from} was requested \
(Content-Range: {raw})"
)));
}
Ok((RangeKind::Partial, total))
}
200 => {
let detail = match accept_ranges {
Some(v) => format!("200 OK, Accept-Ranges: {v}"),
None => "200 OK, no Accept-Ranges header".to_owned(),
};
Ok((RangeKind::Full { detail }, content_length))
}
other => Err(DownloadError::Range(format!(
"unexpected status {other} (expected 200 or 206)"
))),
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum RangeKind {
Partial,
Full {
detail: String,
},
}
fn parse_content_range(raw: &str) -> Result<(u64, Option<u64>), DownloadError> {
let bad = || DownloadError::Range(format!("unparseable Content-Range: {raw}"));
let rest = raw.trim().strip_prefix("bytes ").ok_or_else(bad)?;
let (range, total) = rest.split_once('/').ok_or_else(bad)?;
let (start, _end) = range.split_once('-').ok_or_else(bad)?;
let start: u64 = start.trim().parse().map_err(|_| bad())?;
let total = match total.trim() {
"*" => None,
n => Some(n.parse().map_err(|_| bad())?),
};
Ok((start, total))
}
pub fn download_verified(
mut reader: impl std::io::Read,
dest: &Path,
expected_sha256: &str,
) -> Result<(), DownloadError> {
use sha2::{Digest, Sha256};
struct PartialGuard<'a> {
path: &'a Path,
armed: bool,
}
impl Drop for PartialGuard<'_> {
fn drop(&mut self) {
if self.armed {
std::fs::remove_file(self.path).ok();
}
}
}
let tmp = dest.with_extension("partial");
let mut guard = PartialGuard {
path: &tmp,
armed: true,
};
let mut writer = std::io::BufWriter::new(std::fs::File::create(&tmp)?);
let mut hasher = Sha256::new();
let mut buf = vec![0u8; 1 << 16]; loop {
let n = reader.read(&mut buf)?;
if n == 0 {
break;
}
hasher.update(&buf[..n]);
std::io::Write::write_all(&mut writer, &buf[..n])?;
}
writer
.into_inner()
.map_err(std::io::IntoInnerError::into_error)?
.sync_all()?;
if !expected_sha256.is_empty() {
let mut got = String::with_capacity(64);
for byte in hasher.finalize() {
use std::fmt::Write as _;
let _ = write!(got, "{byte:02x}");
}
if !got.eq_ignore_ascii_case(expected_sha256) {
return Err(DownloadError::Checksum {
expected: expected_sha256.to_owned(),
got,
});
}
}
if dest.exists() {
std::fs::remove_file(dest)?;
}
std::fs::rename(&tmp, dest)?;
guard.armed = false; Ok(())
}
struct HashingWriter<'a, W> {
inner: W,
hasher: &'a mut sha2::Sha256,
}
impl<W: std::io::Write> std::io::Write for HashingWriter<'_, W> {
fn write(&mut self, buf: &[u8]) -> std::io::Result<usize> {
use sha2::Digest as _;
let n = self.inner.write(buf)?;
self.hasher.update(&buf[..n]);
Ok(n)
}
fn flush(&mut self) -> std::io::Result<()> {
self.inner.flush()
}
}
fn existing_len(path: &Path) -> u64 {
std::fs::metadata(path).map_or(0, |m| m.len())
}
fn hash_prefix(path: &Path, len: u64, hasher: &mut sha2::Sha256) -> std::io::Result<()> {
let mut src = std::io::Read::take(std::fs::File::open(path)?, len);
let mut sink = HashingWriter {
inner: std::io::sink(),
hasher,
};
let read = std::io::copy(&mut src, &mut sink)?;
if read != len {
return Err(std::io::Error::new(
std::io::ErrorKind::UnexpectedEof,
format!("partial download shrank while being re-hashed ({read} of {len} bytes)"),
));
}
Ok(())
}
fn load_partial_meta(path: &Path) -> Option<PartialMeta> {
let raw = std::fs::read(path).ok()?;
serde_json::from_slice(&raw).ok()
}
fn write_partial_meta(path: &Path, meta: &PartialMeta) -> std::io::Result<()> {
let json = serde_json::to_vec(meta).map_err(std::io::Error::other)?;
std::fs::write(path, json)
}
fn install_verified(
dest: &Path,
expected_sha256: &str,
hasher: sha2::Sha256,
on_event: &mut impl FnMut(DownloadEvent),
) -> Result<(), DownloadError> {
use sha2::Digest as _;
let tmp = partial_path(dest);
if !expected_sha256.is_empty() {
let mut got = String::with_capacity(64);
for byte in hasher.finalize() {
use std::fmt::Write as _;
let _ = write!(got, "{byte:02x}");
}
if !got.eq_ignore_ascii_case(expected_sha256) {
let bytes = existing_len(&tmp);
discard_partial(dest)?;
on_event(DownloadEvent::PoisonedPartial { bytes });
return Err(DownloadError::Checksum {
expected: expected_sha256.to_owned(),
got,
});
}
}
if dest.exists() {
std::fs::remove_file(dest)?;
}
std::fs::rename(&tmp, dest)?;
remove_if_present(&partial_meta_path(dest))?;
Ok(())
}
fn plan_resume(
dest: &Path,
url: &str,
expected_sha256: &str,
on_event: &mut impl FnMut(DownloadEvent),
) -> Result<(u64, Option<u64>), DownloadError> {
let on_disk = existing_len(&partial_path(dest));
if on_disk == 0 {
return Ok((0, None));
}
let mut known_total = None;
let reject = match load_partial_meta(&partial_meta_path(dest)) {
None => Some("no sidecar recording what it was started against".to_owned()),
Some(m) if m.version != PARTIAL_META_VERSION => Some(format!(
"its sidecar is format v{}, not v{PARTIAL_META_VERSION}",
m.version
)),
Some(m) if m.url != url => Some("it was started against a different URL".to_owned()),
Some(m) if m.sha256 != expected_sha256 => {
Some("the pinned checksum changed since it was started".to_owned())
}
Some(m) => match m.total {
Some(t) if on_disk > t => Some(format!(
"it is larger ({on_disk} bytes) than the recorded total ({t} bytes)"
)),
total => {
known_total = total;
None
}
},
};
if let Some(reason) = reject {
on_event(DownloadEvent::DiscardedPartial {
bytes: on_disk,
reason,
});
discard_partial(dest)?;
return Ok((0, None));
}
Ok((on_disk, known_total))
}
fn start_transfer<R, F, E>(
dest: &Path,
open: &mut F,
on_event: &mut E,
hasher: &mut sha2::Sha256,
resume_from: &mut u64,
known_total: &mut Option<u64>,
) -> Result<(R, bool), DownloadError>
where
R: std::io::Read,
F: FnMut(u64) -> Result<RangeReply<R>, DownloadError>,
E: FnMut(DownloadEvent),
{
use sha2::Digest as _;
let tmp = partial_path(dest);
loop {
if *resume_from > 0 {
hash_prefix(&tmp, *resume_from, hasher)?;
}
match open(*resume_from)? {
RangeReply::Partial { reader, total } => {
if let (Some(remote), Some(recorded)) = (total, *known_total)
&& remote != recorded
{
on_event(DownloadEvent::DiscardedPartial {
bytes: *resume_from,
reason: format!(
"the remote is now {remote} bytes but it was started against {recorded}"
),
});
discard_partial(dest)?;
*resume_from = 0;
*known_total = None;
*hasher = sha2::Sha256::new();
continue;
}
*known_total = total.or(*known_total);
on_event(DownloadEvent::Resuming {
offset: *resume_from,
total: *known_total,
});
return Ok((reader, true));
}
RangeReply::Full {
reader,
total,
detail,
} => {
if *resume_from > 0 {
on_event(DownloadEvent::RangeUnsupported {
discarded: *resume_from,
detail,
});
*resume_from = 0;
*hasher = sha2::Sha256::new();
}
*known_total = total;
return Ok((reader, false));
}
}
}
}
pub fn download_resumable<R, F, E>(
dest: &Path,
url: &str,
expected_sha256: &str,
mut open: F,
mut on_event: E,
) -> Result<(), DownloadError>
where
R: std::io::Read,
F: FnMut(u64) -> Result<RangeReply<R>, DownloadError>,
E: FnMut(DownloadEvent),
{
let result = download_attempt(dest, url, expected_sha256, &mut open, &mut on_event);
if result.is_err() {
let kept = existing_len(&partial_path(dest));
if kept > 0 {
on_event(DownloadEvent::KeptPartial { bytes: kept });
}
}
result
}
fn download_attempt<R, F, E>(
dest: &Path,
url: &str,
expected_sha256: &str,
open: &mut F,
on_event: &mut E,
) -> Result<(), DownloadError>
where
R: std::io::Read,
F: FnMut(u64) -> Result<RangeReply<R>, DownloadError>,
E: FnMut(DownloadEvent),
{
use sha2::Digest as _;
let tmp = partial_path(dest);
let meta_path = partial_meta_path(dest);
let (mut resume_from, mut known_total) = plan_resume(dest, url, expected_sha256, on_event)?;
if resume_from > 0 && known_total == Some(resume_from) {
on_event(DownloadEvent::AlreadyComplete { bytes: resume_from });
let mut hasher = sha2::Sha256::new();
hash_prefix(&tmp, resume_from, &mut hasher)?;
return install_verified(dest, expected_sha256, hasher, on_event);
}
let mut hasher = sha2::Sha256::new();
let (mut reader, append) = start_transfer(
dest,
open,
on_event,
&mut hasher,
&mut resume_from,
&mut known_total,
)?;
write_partial_meta(
&meta_path,
&PartialMeta {
version: PARTIAL_META_VERSION,
url: url.to_owned(),
sha256: expected_sha256.to_owned(),
total: known_total,
},
)?;
let mut file = std::fs::OpenOptions::new()
.create(true)
.write(true)
.truncate(false)
.open(&tmp)?;
if append {
std::io::Seek::seek(&mut file, std::io::SeekFrom::Start(resume_from))?;
} else {
file.set_len(0)?;
}
let mut sink = HashingWriter {
inner: std::io::BufWriter::with_capacity(1 << 20, file),
hasher: &mut hasher,
};
let streamed = std::io::copy(&mut reader, &mut sink).map(|_| ());
let durable = std::io::Write::flush(&mut sink).and_then(|()| sink.inner.get_ref().sync_all());
drop(sink);
if let Err(e) = streamed.and(durable) {
return Err(DownloadError::Io(e));
}
let on_disk = existing_len(&tmp);
if let Some(total) = known_total {
if on_disk < total {
return Err(DownloadError::Io(std::io::Error::new(
std::io::ErrorKind::UnexpectedEof,
format!("connection closed after {on_disk} of {total} bytes"),
)));
}
if on_disk > total {
discard_partial(dest)?;
on_event(DownloadEvent::PoisonedPartial { bytes: on_disk });
return Err(DownloadError::Range(format!(
"server sent {on_disk} bytes for a {total}-byte resource"
)));
}
}
install_verified(dest, expected_sha256, hasher, on_event)
}
#[cfg(test)]
mod testserver {
use std::io::{Read as _, Write as _};
use std::net::{SocketAddr, TcpListener, TcpStream};
use std::sync::atomic::{AtomicBool, Ordering};
use std::sync::{Arc, Mutex};
pub struct Behaviour {
pub body: Vec<u8>,
pub ranges: bool,
pub limit: Option<usize>,
}
pub type Hit = (u64, usize);
pub struct TestServer {
pub addr: SocketAddr,
pub behaviour: Arc<Mutex<Behaviour>>,
pub hits: Arc<Mutex<Vec<Hit>>>,
stop: Arc<AtomicBool>,
}
impl TestServer {
pub fn start(behaviour: Behaviour) -> Self {
let listener = TcpListener::bind("127.0.0.1:0").expect("bind");
let addr = listener.local_addr().expect("addr");
let behaviour = Arc::new(Mutex::new(behaviour));
let hits = Arc::new(Mutex::new(Vec::new()));
let stop = Arc::new(AtomicBool::new(false));
{
let (behaviour, hits, stop) = (behaviour.clone(), hits.clone(), stop.clone());
std::thread::spawn(move || {
for sock in listener.incoming() {
if stop.load(Ordering::SeqCst) {
break;
}
if let Ok(sock) = sock {
handle(sock, &behaviour, &hits);
}
}
});
}
Self {
addr,
behaviour,
hits,
stop,
}
}
pub fn hits(&self) -> Vec<Hit> {
self.hits.lock().expect("hits").clone()
}
pub fn set_limit(&self, limit: Option<usize>) {
self.behaviour.lock().expect("behaviour").limit = limit;
}
}
impl Drop for TestServer {
fn drop(&mut self) {
self.stop.store(true, Ordering::SeqCst);
let _ = TcpStream::connect(self.addr);
}
}
fn handle(mut sock: TcpStream, behaviour: &Arc<Mutex<Behaviour>>, hits: &Arc<Mutex<Vec<Hit>>>) {
let Some(head) = read_head(&mut sock) else {
return;
};
let requested = head
.lines()
.find_map(|l| {
l.to_ascii_lowercase()
.strip_prefix("range:")
.map(str::trim)
.map(str::to_owned)
})
.and_then(|v| v.strip_prefix("bytes=").map(str::to_owned))
.and_then(|v| v.split('-').next().and_then(|n| n.parse::<u64>().ok()));
let b = behaviour.lock().expect("behaviour");
let total = b.body.len();
let start = match requested {
Some(from) if b.ranges => usize::try_from(from).expect("offset fits"),
_ => 0,
};
let partial = b.ranges && requested.is_some();
let remainder = &b.body[start.min(total)..];
let serve = b.limit.unwrap_or(remainder.len()).min(remainder.len());
let status_line = if partial {
format!(
"HTTP/1.1 206 Partial Content\r\nContent-Range: bytes {start}-{}/{total}\r\n",
total.saturating_sub(1)
)
} else {
"HTTP/1.1 200 OK\r\n".to_owned()
};
let accept = if b.ranges {
"Accept-Ranges: bytes\r\n"
} else {
"Accept-Ranges: none\r\n"
};
let resp = format!(
"{status_line}{accept}Content-Length: {}\r\nConnection: close\r\n\r\n",
remainder.len()
);
let _ = sock.write_all(resp.as_bytes());
let _ = sock.write_all(&remainder[..serve]);
let _ = sock.flush();
hits.lock()
.expect("hits")
.push((u64::try_from(start).unwrap_or(0), serve));
}
fn read_head(sock: &mut TcpStream) -> Option<String> {
let mut buf = Vec::new();
let mut byte = [0u8; 1];
while !buf.ends_with(b"\r\n\r\n") {
match sock.read(&mut byte) {
Ok(0) | Err(_) => return None,
Ok(_) => buf.push(byte[0]),
}
}
Some(String::from_utf8_lossy(&buf).into_owned())
}
}
#[cfg(test)]
fn test_get_range(
addr: std::net::SocketAddr,
from: u64,
) -> Result<RangeReply<std::net::TcpStream>, DownloadError> {
use std::io::{Read as _, Write as _};
let transport = |e: std::io::Error| DownloadError::Transport(Box::new(e));
let mut sock = std::net::TcpStream::connect(addr).map_err(transport)?;
let range = if from > 0 {
format!("Range: bytes={from}-\r\n")
} else {
String::new()
};
let req =
format!("GET /model.bin HTTP/1.1\r\nHost: localhost\r\nConnection: close\r\n{range}\r\n");
sock.write_all(req.as_bytes()).map_err(transport)?;
let mut head = Vec::new();
let mut byte = [0u8; 1];
while !head.ends_with(b"\r\n\r\n") {
match sock.read(&mut byte) {
Ok(0) => return Err(DownloadError::Range("connection closed in headers".into())),
Ok(_) => head.push(byte[0]),
Err(e) => return Err(transport(e)),
}
}
let head = String::from_utf8_lossy(&head).into_owned();
let mut lines = head.lines();
let status: u16 = lines
.next()
.and_then(|l| l.split_whitespace().nth(1).and_then(|s| s.parse().ok()))
.ok_or_else(|| DownloadError::Range("no status line".into()))?;
let (mut accept_ranges, mut content_range, mut content_length) = (None, None, None);
for line in lines {
let Some((k, v)) = line.split_once(':') else {
continue;
};
let v = v.trim().to_owned();
match k.trim().to_ascii_lowercase().as_str() {
"accept-ranges" => accept_ranges = Some(v),
"content-range" => content_range = Some(v),
"content-length" => content_length = v.parse().ok(),
_ => {}
}
}
let (kind, total) = interpret_range_response(
status,
accept_ranges.as_deref(),
content_range.as_deref(),
content_length,
from,
)?;
Ok(match kind {
RangeKind::Partial => RangeReply::Partial {
reader: sock,
total,
},
RangeKind::Full { detail } => RangeReply::Full {
reader: sock,
total,
detail,
},
})
}
#[cfg(test)]
mod tests {
use super::testserver::{Behaviour, TestServer};
use super::{
DownloadError, DownloadEvent, ModelKind, Platform, REGISTRY, RangeKind, RangeReply,
ResourceTier, download_resumable, download_verified, find, installed_size,
interpret_range_response, partial_meta_path, partial_path, sha256_hex, store_root,
test_get_range, verify_sha256,
};
use std::path::{Path, PathBuf};
fn scratch(tag: &str) -> PathBuf {
let dir = std::env::temp_dir().join(format!("roteiro-dl-{tag}-{}", std::process::id()));
std::fs::remove_dir_all(&dir).ok();
std::fs::create_dir_all(&dir).expect("mkdir");
dir
}
fn u(n: usize) -> u64 {
u64::try_from(n).expect("length fits u64")
}
fn payload(len: usize) -> Vec<u8> {
(0..len)
.map(|i| u8::try_from(i % 251).unwrap_or(0))
.collect()
}
fn attempt(
server: &TestServer,
dest: &Path,
sha: &str,
) -> (Result<(), DownloadError>, Vec<DownloadEvent>) {
let addr = server.addr;
let mut events = Vec::new();
let url = format!("http://{addr}/model.bin");
let result = download_resumable(
dest,
&url,
sha,
|from| test_get_range(addr, from),
|e| events.push(e),
);
(result, events)
}
#[test]
fn resumable_clean_download() {
let dir = scratch("clean");
let dest = dir.join("model.bin");
let body = payload(50_000);
let sha = sha256_hex(&body);
let server = TestServer::start(Behaviour {
body: body.clone(),
ranges: true,
limit: None,
});
let (result, events) = attempt(&server, &dest, &sha);
result.expect("clean download");
assert_eq!(std::fs::read(&dest).expect("installed"), body);
assert!(!partial_path(&dest).exists());
assert!(!partial_meta_path(&dest).exists());
assert_eq!(server.hits(), vec![(0, 50_000)]);
assert!(events.is_empty(), "unexpected events: {events:?}");
std::fs::remove_dir_all(&dir).ok();
}
#[test]
fn resumable_interrupted_then_resumed_transfers_only_the_remainder() {
const CUT: usize = 20_000;
let dir = scratch("resume");
let dest = dir.join("model.bin");
let body = payload(50_000);
let sha = sha256_hex(&body);
let server = TestServer::start(Behaviour {
body: body.clone(),
ranges: true,
limit: Some(CUT),
});
let (result, events) = attempt(&server, &dest, &sha);
let err = result.expect_err("interrupted");
assert!(
matches!(err, DownloadError::Io(_)),
"a dropped connection must be an I/O failure, not a checksum one: {err:?}"
);
assert_eq!(events, vec![DownloadEvent::KeptPartial { bytes: u(CUT) }]);
assert_eq!(
std::fs::metadata(partial_path(&dest))
.expect("partial kept")
.len(),
u(CUT)
);
assert!(partial_meta_path(&dest).exists());
assert!(!dest.exists());
server.set_limit(None);
let (result, events) = attempt(&server, &dest, &sha);
result.expect("resumed");
assert_eq!(std::fs::read(&dest).expect("installed"), body);
assert_eq!(
events,
vec![DownloadEvent::Resuming {
offset: u(CUT),
total: Some(50_000),
}]
);
assert_eq!(
server.hits(),
vec![(0, CUT), (u(CUT), 50_000 - CUT)],
"the resumed attempt must transfer only the remainder"
);
assert!(!partial_path(&dest).exists());
assert!(!partial_meta_path(&dest).exists());
std::fs::remove_dir_all(&dir).ok();
}
#[test]
fn resumable_server_without_range_support_restarts_and_says_so() {
const CUT: usize = 15_000;
let dir = scratch("norange");
let dest = dir.join("model.bin");
let body = payload(40_000);
let sha = sha256_hex(&body);
let server = TestServer::start(Behaviour {
body: body.clone(),
ranges: false,
limit: Some(CUT),
});
let (result, _) = attempt(&server, &dest, &sha);
assert!(matches!(result, Err(DownloadError::Io(_))));
assert_eq!(
std::fs::metadata(partial_path(&dest)).expect("kept").len(),
u(CUT)
);
server.set_limit(None);
let (result, events) = attempt(&server, &dest, &sha);
result.expect("restarted");
assert_eq!(
std::fs::read(&dest).expect("installed"),
body,
"restarting must not append a whole-file body onto the stale prefix"
);
match events.as_slice() {
[DownloadEvent::RangeUnsupported { discarded, detail }] => {
assert_eq!(*discarded, u(CUT));
assert!(
detail.contains("200") && detail.to_ascii_lowercase().contains("accept-ranges"),
"the message must explain why: {detail}"
);
}
other => panic!("expected a single RangeUnsupported event, got {other:?}"),
}
assert_eq!(server.hits(), vec![(0, CUT), (0, 40_000)]);
std::fs::remove_dir_all(&dir).ok();
}
#[test]
fn resumable_stale_partial_is_discarded() {
let dir = scratch("stale");
let body = payload(30_000);
let sha = sha256_hex(&body);
let server = TestServer::start(Behaviour {
body: body.clone(),
ranges: true,
limit: None,
});
let dest = dir.join("nometa.bin");
std::fs::write(partial_path(&dest), b"anonymous bytes").expect("write");
let (result, events) = attempt(&server, &dest, &sha);
result.expect("re-downloaded");
assert_eq!(std::fs::read(&dest).expect("installed"), body);
assert!(
matches!(
events.first(),
Some(DownloadEvent::DiscardedPartial { bytes: 15, reason }) if reason.contains("sidecar")
),
"{events:?}"
);
let dest = dir.join("othersha.bin");
std::fs::write(partial_path(&dest), &body[..1000]).expect("write");
std::fs::write(
partial_meta_path(&dest),
format!(
r#"{{"version":1,"url":"http://{}/model.bin","sha256":"{}","total":30000}}"#,
server.addr,
"0".repeat(64)
),
)
.expect("write meta");
let (result, events) = attempt(&server, &dest, &sha);
result.expect("re-downloaded");
assert_eq!(std::fs::read(&dest).expect("installed"), body);
assert!(
matches!(
events.first(),
Some(DownloadEvent::DiscardedPartial { reason, .. }) if reason.contains("checksum")
),
"{events:?}"
);
let dest = dir.join("otherurl.bin");
std::fs::write(partial_path(&dest), &body[..2000]).expect("write");
std::fs::write(
partial_meta_path(&dest),
format!(
r#"{{"version":1,"url":"http://elsewhere.invalid/model.bin","sha256":"{sha}","total":30000}}"#
),
)
.expect("write meta");
let (result, events) = attempt(&server, &dest, &sha);
result.expect("re-downloaded");
assert!(
matches!(
events.first(),
Some(DownloadEvent::DiscardedPartial { reason, .. }) if reason.contains("URL")
),
"{events:?}"
);
let dest = dir.join("othersize.bin");
std::fs::write(partial_path(&dest), &body[..3000]).expect("write");
std::fs::write(
partial_meta_path(&dest),
format!(
r#"{{"version":1,"url":"http://{}/model.bin","sha256":"{sha}","total":999999}}"#,
server.addr
),
)
.expect("write meta");
let (result, events) = attempt(&server, &dest, &sha);
result.expect("re-downloaded");
assert_eq!(std::fs::read(&dest).expect("installed"), body);
assert!(
matches!(
events.first(),
Some(DownloadEvent::DiscardedPartial { reason, .. }) if reason.contains("30000")
),
"{events:?}"
);
std::fs::remove_dir_all(&dir).ok();
}
#[test]
fn a_transport_that_never_opens_still_reports_the_partial_it_kept() {
let dir = scratch("openfail");
let dest = dir.join("model.bin");
let body = payload(30_000);
let sha = sha256_hex(&body);
let dead = "http://example.invalid/model.bin";
std::fs::write(partial_path(&dest), &body[..9_000]).expect("write");
std::fs::write(
partial_meta_path(&dest),
format!(r#"{{"version":1,"url":"{dead}","sha256":"{sha}","total":30000}}"#),
)
.expect("write meta");
let mut events = Vec::new();
let result = download_resumable(
&dest,
dead,
&sha,
|_from| -> Result<RangeReply<std::io::Empty>, DownloadError> {
Err(DownloadError::Transport("connection refused".into()))
},
|e| events.push(e),
);
let err = result.expect_err("transport failure");
assert!(matches!(err, DownloadError::Transport(_)), "{err:?}");
assert_eq!(events, vec![DownloadEvent::KeptPartial { bytes: 9_000 }]);
assert_eq!(
std::fs::metadata(partial_path(&dest)).expect("kept").len(),
9_000,
"the prefix itself must survive, not merely be announced"
);
assert!(partial_meta_path(&dest).exists(), "sidecar survives too");
let server = TestServer::start(Behaviour {
body: body.clone(),
ranges: true,
limit: None,
});
let live = format!("http://{}/model.bin", server.addr);
std::fs::write(
partial_meta_path(&dest),
format!(r#"{{"version":1,"url":"{live}","sha256":"{sha}","total":30000}}"#),
)
.expect("rewrite meta");
let (result, _) = attempt(&server, &dest, &sha);
result.expect("resumed");
assert_eq!(std::fs::read(&dest).expect("installed"), body);
assert_eq!(server.hits(), vec![(9_000, 21_000)], "only the remainder");
std::fs::remove_dir_all(&dir).ok();
}
#[test]
fn resumable_checksum_failure_discards_the_partial() {
let dir = scratch("poison");
let dest = dir.join("model.bin");
let body = payload(25_000);
let server = TestServer::start(Behaviour {
body,
ranges: true,
limit: None,
});
let (result, events) = attempt(&server, &dest, &"a".repeat(64));
let err = result.expect_err("checksum mismatch");
assert!(matches!(err, DownloadError::Checksum { .. }), "{err:?}");
assert_eq!(
events,
vec![DownloadEvent::PoisonedPartial { bytes: 25_000 }]
);
assert!(!partial_path(&dest).exists());
assert!(!partial_meta_path(&dest).exists());
assert!(!dest.exists());
std::fs::remove_dir_all(&dir).ok();
}
#[test]
fn resumable_complete_partial_only_needs_verifying() {
let dir = scratch("complete");
let dest = dir.join("model.bin");
let body = payload(12_345);
let sha = sha256_hex(&body);
let server = TestServer::start(Behaviour {
body: body.clone(),
ranges: true,
limit: None,
});
std::fs::write(partial_path(&dest), &body).expect("write");
std::fs::write(
partial_meta_path(&dest),
format!(
r#"{{"version":1,"url":"http://{}/model.bin","sha256":"{sha}","total":12345}}"#,
server.addr
),
)
.expect("write meta");
let (result, events) = attempt(&server, &dest, &sha);
result.expect("installed from the complete partial");
assert_eq!(std::fs::read(&dest).expect("installed"), body);
assert_eq!(
events,
vec![DownloadEvent::AlreadyComplete { bytes: 12_345 }]
);
assert!(server.hits().is_empty(), "no request should have been made");
std::fs::remove_dir_all(&dir).ok();
}
#[test]
fn range_response_interpretation() {
let (kind, total) =
interpret_range_response(206, Some("bytes"), Some("bytes 100-499/500"), None, 100)
.expect("206");
assert_eq!(kind, RangeKind::Partial);
assert_eq!(total, Some(500));
let (_, total) =
interpret_range_response(206, None, Some("bytes 10-19/*"), None, 10).expect("206 *");
assert_eq!(total, None);
let err = interpret_range_response(206, None, Some("bytes 0-499/500"), None, 100)
.expect_err("wrong offset");
assert!(matches!(err, DownloadError::Range(_)), "{err:?}");
assert!(matches!(
interpret_range_response(206, None, None, None, 5),
Err(DownloadError::Range(_))
));
assert!(matches!(
interpret_range_response(206, None, Some("chunks 1-2/3"), None, 1),
Err(DownloadError::Range(_))
));
let (kind, total) =
interpret_range_response(200, Some("none"), None, Some(500), 100).expect("200");
match kind {
RangeKind::Full { detail } => assert!(detail.contains("none"), "{detail}"),
RangeKind::Partial => panic!("expected Full, got Partial"),
}
assert_eq!(total, Some(500));
let (kind, _) = interpret_range_response(200, None, None, None, 0).expect("200");
match kind {
RangeKind::Full { detail } => assert!(detail.contains("no Accept-Ranges"), "{detail}"),
RangeKind::Partial => panic!("expected Full, got Partial"),
}
assert!(matches!(
interpret_range_response(416, None, None, None, 10),
Err(DownloadError::Range(_))
));
}
#[test]
fn installed_size_and_removal() {
use super::{model_dir, remove_model, set_model_store};
let dir = scratch("store");
set_model_store(dir.join("models"));
let root = super::store_root();
let name = "size-probe";
assert_eq!(installed_size(name), 0, "absent model occupies nothing");
let mdir = root.join(name);
std::fs::create_dir_all(&mdir).expect("mkdir");
std::fs::write(mdir.join("model.gguf"), vec![7u8; 4096]).expect("write");
std::fs::write(mdir.join("model.partial"), vec![7u8; 1024]).expect("write");
std::fs::write(mdir.join("model.partial.json"), b"{}").expect("write");
assert_eq!(model_dir(name), mdir);
assert_eq!(installed_size(name), 4096 + 1024 + 2);
let removed = remove_model(name).expect("removed");
assert_eq!(removed.bytes, 4096 + 1024 + 2);
assert_eq!(
removed.files,
vec!["model.gguf", "model.partial", "model.partial.json"],
"an orphaned partial is cleaned up with the model"
);
assert!(!mdir.exists());
let again = remove_model(name).expect("idempotent");
assert_eq!(again.bytes, 0);
assert!(again.files.is_empty());
std::fs::remove_dir_all(&dir).ok();
}
#[test]
fn download_verified_streams_and_checks() {
struct FailReader(usize);
impl std::io::Read for FailReader {
fn read(&mut self, buf: &mut [u8]) -> std::io::Result<usize> {
if self.0 == 0 {
return Err(std::io::Error::new(std::io::ErrorKind::BrokenPipe, "boom"));
}
let n = buf.len().min(self.0);
buf[..n].fill(b'x');
self.0 -= n;
Ok(n)
}
}
let dir = std::env::temp_dir().join(format!("roteiro-dl-{}", std::process::id()));
std::fs::create_dir_all(&dir).expect("mkdir");
let payload = b"the streamed model bytes";
let sha = sha256_hex(payload);
let good = dir.join("good.bin");
download_verified(&payload[..], &good, &sha).expect("verified");
assert_eq!(std::fs::read(&good).expect("read"), payload);
let bad = dir.join("bad.bin");
let err = download_verified(&payload[..], &bad, &"0".repeat(64)).unwrap_err();
assert!(matches!(err, DownloadError::Checksum { .. }));
assert!(!bad.exists());
assert!(!bad.with_extension("partial").exists());
let dropped = dir.join("dropped.bin");
let err = download_verified(FailReader(100), &dropped, "").unwrap_err();
assert!(matches!(err, DownloadError::Io(_)));
assert!(!dropped.exists());
assert!(!dropped.with_extension("partial").exists());
let unpinned = dir.join("unpinned.bin");
download_verified(&payload[..], &unpinned, "").expect("unpinned");
assert!(unpinned.exists());
std::fs::remove_dir_all(&dir).ok();
}
#[test]
fn registry_entries_are_well_formed() {
assert!(!REGISTRY.is_empty());
for spec in REGISTRY {
assert!(!spec.name.is_empty());
assert_eq!(
spec.dim > 0,
spec.kind == ModelKind::Embedding,
"{}",
spec.name
);
assert!(!spec.variants.is_empty());
assert!(
spec.variants
.iter()
.any(|v| v.platform == Platform::Standard),
"{} needs a Standard variant",
spec.name,
);
let v = spec.variant_for(Platform::host()).expect("host variant");
assert!(!v.files.is_empty());
assert!(!spec.tier.as_str().is_empty());
}
}
#[test]
fn every_section_has_a_low_tier_floor() {
for kind in [
ModelKind::Embedding,
ModelKind::Generative,
ModelKind::Ocr,
ModelKind::Vision,
] {
assert!(
REGISTRY
.iter()
.any(|s| s.kind == kind && s.tier == ResourceTier::Low),
"section {} needs a Low-tier entry",
kind.as_str(),
);
}
}
#[test]
fn variant_selection_falls_back_to_standard() {
let spec = find("bge-base-en-v1.5").expect("registered");
let mac = spec.variant_for(Platform::MacosArm64).expect("mac");
let std = spec.variant_for(Platform::Standard).expect("std");
assert_eq!(mac.platform, Platform::Standard);
assert_eq!(std.platform, Platform::Standard);
}
#[test]
fn platform_host_is_stable() {
let p = Platform::host();
assert!(matches!(p, Platform::MacosArm64 | Platform::Standard));
assert!(!p.as_str().is_empty());
}
#[test]
fn store_root_resolution() {
use super::store_root_from;
use std::path::PathBuf;
assert_eq!(
store_root_from(
Some(PathBuf::from("/data/models")),
Some(PathBuf::from("/opt/rt")),
Some(PathBuf::from("/home/u"))
),
Path::new("/data/models"),
);
assert_eq!(
store_root_from(
None,
Some(PathBuf::from("/opt/rt")),
Some(PathBuf::from("/home/u"))
),
Path::new("/opt/rt/models"),
);
assert_eq!(
store_root_from(None, None, Some(PathBuf::from("/home/u"))),
Path::new("/home/u/.roteiro/models"),
);
assert!(store_root().ends_with("models"));
}
#[test]
fn sha256_and_verify() {
let want = "ba7816bf8f01cfea414140de5dae2223b00361a396177a9cb410ff61f20015ad";
assert_eq!(sha256_hex(b"abc"), want);
assert!(verify_sha256(b"abc", want));
assert!(verify_sha256(b"abc", &want.to_uppercase()));
assert!(!verify_sha256(b"abc", "00"));
assert!(verify_sha256(b"anything", ""));
}
}