use crate::error::{Result, YantrikDbError};
use crate::types::Embedder;
use model2vec_rs::model::StaticModel;
use std::path::PathBuf;
struct DownloadableModel {
release_tag: &'static str,
asset: &'static str,
sha256: &'static str,
dim: usize,
}
fn registry(name: &str) -> Option<DownloadableModel> {
match name {
"potion-base-8M" => Some(DownloadableModel {
release_tag: "v0.1.0",
asset: "potion-base-8M.tar.gz",
sha256: "89dd960591c4fa0c7f7a45ed4cb94167ce4e09886f39bae008b8072b42439ac5",
dim: 256,
}),
"potion-base-32M" => Some(DownloadableModel {
release_tag: "v0.1.0",
asset: "potion-base-32M.tar.gz",
sha256: "428163e9aa596b38bf98f6d41ff7cb1b3d7e6d21f58e1edc8124cd9d180f93ad",
dim: 512,
}),
"potion-multilingual-128M" => Some(DownloadableModel {
release_tag: "v0.2.0",
asset: "potion-multilingual-128M.tar.gz",
sha256: "bbd9b15fa1303538206911f82c85b5d52e3fa0a334479f988e8a370b0e2e7a52",
dim: 256,
}),
_ => None,
}
}
fn asset_url(model: &DownloadableModel) -> String {
format!(
"https://github.com/yantrikos/yantrikdb-models/releases/download/{}/{}",
model.release_tag, model.asset
)
}
fn cache_dir_for(name: &str, release_tag: &str) -> Result<PathBuf> {
let base = dirs::cache_dir().ok_or_else(|| {
YantrikDbError::InvalidInput(
"could not resolve user cache dir; set XDG_CACHE_HOME or HOME".into(),
)
})?;
Ok(base
.join("yantrikdb")
.join("models")
.join(format!("{name}-{release_tag}")))
}
fn cache_is_populated(dir: &std::path::Path) -> bool {
for f in &[
"model.safetensors",
"tokenizer.json",
"config.json",
"modules.json",
] {
match std::fs::metadata(dir.join(f)) {
Ok(m) if m.len() > 0 => continue,
_ => return false,
}
}
true
}
fn extract_tarball_to(bytes: &[u8], dest_dir: &std::path::Path) -> Result<()> {
let gz = flate2::read::GzDecoder::new(bytes);
let mut archive = tar::Archive::new(gz);
for entry in archive
.entries()
.map_err(|e| YantrikDbError::InvalidInput(format!("tar open: {e}")))?
{
let mut entry =
entry.map_err(|e| YantrikDbError::InvalidInput(format!("tar entry: {e}")))?;
if entry.header().entry_type().is_dir() {
continue;
}
let path = entry
.path()
.map_err(|e| YantrikDbError::InvalidInput(format!("tar path: {e}")))?;
let n_components = path.components().count();
let stripped: PathBuf = if n_components >= 2 {
path.components().skip(1).collect()
} else {
path.into_owned()
};
if stripped.as_os_str().is_empty() {
continue;
}
let dest = dest_dir.join(&stripped);
if let Some(parent) = dest.parent() {
std::fs::create_dir_all(parent).ok();
}
entry.unpack(&dest).map_err(|e| {
YantrikDbError::InvalidInput(format!("tar unpack {}: {e}", dest.display()))
})?;
}
Ok(())
}
const DOWNLOAD_ATTEMPTS: u32 = 4;
fn download_with_retry(url: &str) -> Result<Vec<u8>> {
use std::io::Read;
let mut last_err = String::new();
for attempt in 1..=DOWNLOAD_ATTEMPTS {
let outcome = ureq::get(url)
.set(
"User-Agent",
concat!("yantrikdb/", env!("CARGO_PKG_VERSION")),
)
.call();
let retryable = match &outcome {
Ok(_) => false,
Err(ureq::Error::Status(code, _)) => *code == 429 || *code >= 500,
Err(ureq::Error::Transport(_)) => true,
};
match outcome {
Ok(resp) => {
let mut bytes = Vec::with_capacity(64 * 1024 * 1024);
match resp.into_reader().read_to_end(&mut bytes) {
Ok(_) => return Ok(bytes),
Err(e) => last_err = format!("read body: {e}"),
}
}
Err(e) => {
last_err = format!("download {url}: {e}");
if !retryable {
return Err(YantrikDbError::InvalidInput(last_err));
}
}
}
if attempt < DOWNLOAD_ATTEMPTS {
let backoff = std::time::Duration::from_secs(1 << (attempt - 1));
tracing::warn!(
target: "yantrikdb::embedder::download",
url = url,
attempt,
of = DOWNLOAD_ATTEMPTS,
backoff_secs = backoff.as_secs(),
error = %last_err,
"artifact download failed; retrying"
);
std::thread::sleep(backoff);
}
}
Err(YantrikDbError::InvalidInput(format!(
"{last_err} (after {DOWNLOAD_ATTEMPTS} attempts)"
)))
}
fn fetch_and_extract(model: &DownloadableModel, name: &str) -> Result<PathBuf> {
use sha2::{Digest, Sha256};
use std::io::Read;
let final_dir = cache_dir_for(name, model.release_tag)?;
if cache_is_populated(&final_dir) {
return Ok(final_dir);
}
let url = asset_url(model);
tracing::info!(
target: "yantrikdb::embedder::download",
name = name,
url = %url,
sha256 = model.sha256,
"downloading model artifact"
);
let bytes = download_with_retry(&url)?;
let mut hasher = Sha256::new();
hasher.update(&bytes);
let actual = hex::encode(hasher.finalize());
if actual != model.sha256 {
return Err(YantrikDbError::InvalidInput(format!(
"model {name}: SHA-256 mismatch — expected {}, got {actual}. Refusing to load \
(corrupted download or upstream tampering). Try again or fall back to \
set_embedder() with your own implementation.",
model.sha256,
)));
}
let parent = final_dir
.parent()
.ok_or_else(|| YantrikDbError::InvalidInput("cache dir has no parent".into()))?;
std::fs::create_dir_all(parent)
.map_err(|e| YantrikDbError::InvalidInput(format!("mkdir {}: {e}", parent.display())))?;
let tmp_dir = parent.join(format!(
"{}-{}.tmp.{}",
name,
model.release_tag,
std::process::id()
));
let _ = std::fs::remove_dir_all(&tmp_dir); std::fs::create_dir_all(&tmp_dir).map_err(|e| {
YantrikDbError::InvalidInput(format!("mkdir tmp {}: {e}", tmp_dir.display()))
})?;
extract_tarball_to(&bytes, &tmp_dir).map_err(|e| {
let _ = std::fs::remove_dir_all(&tmp_dir);
e
})?;
if !cache_is_populated(&final_dir) {
std::fs::rename(&tmp_dir, &final_dir).map_err(|e| {
YantrikDbError::InvalidInput(format!(
"rename {} -> {}: {e}",
tmp_dir.display(),
final_dir.display()
))
})?;
} else {
let _ = std::fs::remove_dir_all(&tmp_dir);
}
if !cache_is_populated(&final_dir) {
return Err(YantrikDbError::InvalidInput(format!(
"after extract, expected files missing in {}",
final_dir.display()
)));
}
tracing::info!(
target: "yantrikdb::embedder::download",
name = name,
cache_dir = %final_dir.display(),
"model ready"
);
Ok(final_dir)
}
pub struct DownloadedEmbedder {
model: std::sync::Arc<StaticModel>,
dim: usize,
name: String,
sha256: &'static str,
}
impl DownloadedEmbedder {
pub fn registry_dim(name: &str) -> Option<usize> {
registry(name).map(|m| m.dim)
}
pub fn fetch(name: &str) -> Result<Self> {
let model = registry(name).ok_or_else(|| {
YantrikDbError::InvalidInput(format!(
"unknown embedder name {name:?}; known: \
potion-base-8M, potion-base-32M, potion-multilingual-128M"
))
})?;
let dir = fetch_and_extract(&model, name)?;
let static_model = StaticModel::from_pretrained(&dir, None, None, None).map_err(|e| {
YantrikDbError::InvalidInput(format!("model2vec_rs load from {}: {e}", dir.display()))
})?;
Ok(Self {
model: std::sync::Arc::new(static_model),
dim: model.dim,
name: name.to_string(),
sha256: model.sha256,
})
}
}
impl Embedder for DownloadedEmbedder {
fn embed(
&self,
text: &str,
) -> std::result::Result<Vec<f32>, Box<dyn std::error::Error + Send + Sync>> {
Ok(self.model.encode_single(text))
}
fn embed_batch(
&self,
texts: &[&str],
) -> std::result::Result<Vec<Vec<f32>>, Box<dyn std::error::Error + Send + Sync>> {
let owned: Vec<String> = texts.iter().map(|s| (*s).to_string()).collect();
Ok(self.model.encode(&owned))
}
fn dim(&self) -> usize {
self.dim
}
fn fingerprint(&self) -> Option<String> {
Some(format!("sha256:{}", self.sha256))
}
fn name(&self) -> Option<String> {
Some(self.name.clone())
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn registry_unknown_name_returns_none() {
assert!(registry("definitely-not-real").is_none());
}
#[test]
fn registry_known_models_have_consistent_dim() {
let m8 = registry("potion-base-8M").expect("registered");
assert_eq!(m8.dim, 256);
let m32 = registry("potion-base-32M").expect("registered");
assert_eq!(m32.dim, 512);
assert_eq!(m8.sha256.len(), 64);
assert_eq!(m32.sha256.len(), 64);
}
#[test]
fn asset_url_format_matches_release_pattern() {
let m = registry("potion-base-8M").unwrap();
let url = asset_url(&m);
assert_eq!(
url,
"https://github.com/yantrikos/yantrikdb-models/releases/download/v0.1.0/potion-base-8M.tar.gz"
);
}
#[test]
fn registry_includes_potion_multilingual_128m() {
let m = registry("potion-multilingual-128M").expect("registered");
assert_eq!(m.dim, 256);
assert_eq!(m.release_tag, "v0.2.0");
assert_eq!(m.sha256.len(), 64);
}
fn build_tar_gz(entries: &[(&str, &[u8])]) -> Vec<u8> {
use std::io::Write;
let mut gz = flate2::write::GzEncoder::new(Vec::new(), flate2::Compression::default());
{
let mut builder = tar::Builder::new(&mut gz);
for (path, content) in entries {
let mut header = tar::Header::new_gnu();
header.set_size(content.len() as u64);
header.set_mode(0o644);
header.set_cksum();
builder
.append_data(&mut header, path, *content)
.expect("tar append");
}
builder.finish().expect("tar finish");
}
gz.finish().expect("gz finish")
}
#[test]
fn extract_tarball_files_at_root_layout() {
let bytes = build_tar_gz(&[
("model.safetensors", b"safetensors-bytes" as &[u8]),
("tokenizer.json", b"{\"tokenizer\": true}"),
("config.json", b"{\"hidden_dim\": 256}"),
("modules.json", b"[]"),
]);
let dir = tempfile::tempdir().expect("tmpdir");
extract_tarball_to(&bytes, dir.path()).expect("extract succeeds");
for filename in [
"model.safetensors",
"tokenizer.json",
"config.json",
"modules.json",
] {
let p = dir.path().join(filename);
assert!(p.exists(), "missing expected file: {}", p.display());
}
}
#[test]
fn extract_tarball_files_under_prefix_layout() {
let bytes = build_tar_gz(&[
(
"potion-base-8M/model.safetensors",
b"safetensors-bytes" as &[u8],
),
("potion-base-8M/tokenizer.json", b"{\"tokenizer\": true}"),
("potion-base-8M/config.json", b"{\"hidden_dim\": 256}"),
("potion-base-8M/modules.json", b"[]"),
]);
let dir = tempfile::tempdir().expect("tmpdir");
extract_tarball_to(&bytes, dir.path()).expect("extract succeeds");
for filename in [
"model.safetensors",
"tokenizer.json",
"config.json",
"modules.json",
] {
let p = dir.path().join(filename);
assert!(p.exists(), "missing expected file: {}", p.display());
}
assert!(
!dir.path().join("potion-base-8M").exists(),
"prefix directory should have been stripped"
);
}
#[test]
fn live_download_potion_8m_smoke() {
if std::env::var_os("YANTRIKDB_TEST_LIVE_DOWNLOAD").is_none() {
return; }
let e = DownloadedEmbedder::fetch("potion-base-8M").expect("fetch potion-base-8M live");
assert_eq!(e.dim(), 256);
let v = e.embed("Alice met Acme yesterday").unwrap();
assert_eq!(v.len(), 256);
let norm: f32 = v.iter().map(|x| x * x).sum::<f32>().sqrt();
assert!((norm - 1.0).abs() < 1e-3, "expected unit norm; got {norm}");
}
#[test]
fn only_transient_failures_are_retryable() {
fn retryable(e: &ureq::Error) -> bool {
match e {
ureq::Error::Status(code, _) => *code == 429 || *code >= 500,
ureq::Error::Transport(_) => true,
}
}
let resp = |code: u16| {
ureq::Error::Status(
code,
ureq::Response::new(code, "x", "").expect("synthetic response"),
)
};
assert!(retryable(&resp(500)), "5xx must retry");
assert!(retryable(&resp(503)), "503 must retry");
assert!(retryable(&resp(429)), "429 must retry");
assert!(
!retryable(&resp(404)),
"404 must NOT retry — asset is absent"
);
assert!(!retryable(&resp(403)), "403 must NOT retry");
assert!(!retryable(&resp(400)), "400 must NOT retry");
}
#[test]
fn backoff_is_bounded() {
let total: u64 = (1..DOWNLOAD_ATTEMPTS).map(|a| 1u64 << (a - 1)).sum();
assert_eq!(total, 7, "expected 1s+2s+4s of backoff, got {total}s");
assert!(total <= 10, "an offline caller must not wait minutes");
}
}