use anyhow::{Context, Result};
use std::fs;
use std::path::{Path, PathBuf};
const CODERANK_BASE_URL: &str =
"https://huggingface.co/MisterTK/CodeRankEmbed-onnx-int8/resolve/main";
pub(crate) const CODERANK_FILES: &[&str] = &[
"model_int8.onnx",
"model_int8.onnx.data",
"tokenizer.json",
"config.json",
];
pub(crate) const CODERANK_DIR: &str = "CodeRankEmbed";
pub(crate) const CODERANK_ONNX: &str = "model_int8.onnx";
pub fn ensure_coderank_model(models_dir: &Path) -> Result<PathBuf> {
let dir = models_dir.join(CODERANK_DIR);
if is_coderank_downloaded(models_dir) {
return Ok(dir);
}
fs::create_dir_all(&dir)
.with_context(|| format!("Failed to create model dir: {}", dir.display()))?;
tracing::info!("Downloading CodeRankEmbed single-vector ONNX model...");
for file_name in CODERANK_FILES {
let dest = dir.join(file_name);
if !dest.exists() {
let url = format!("{CODERANK_BASE_URL}/{file_name}");
crate::embedding::model_manager::download_file(&url, &dest)
.with_context(|| format!("Failed to download {file_name} for CodeRankEmbed"))?;
}
}
Ok(dir)
}
pub fn is_coderank_downloaded(models_dir: &Path) -> bool {
let dir = models_dir.join(CODERANK_DIR);
CODERANK_FILES.iter().all(|f| dir.join(f).exists())
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn coderank_not_downloaded_in_empty_dir() {
let tmp = tempfile::TempDir::new().unwrap();
assert!(!is_coderank_downloaded(tmp.path()));
}
#[test]
fn coderank_detected_when_files_present() {
let tmp = tempfile::TempDir::new().unwrap();
let dir = tmp.path().join(CODERANK_DIR);
fs::create_dir_all(&dir).unwrap();
for f in CODERANK_FILES {
fs::write(dir.join(f), b"stub").unwrap();
}
assert!(is_coderank_downloaded(tmp.path()));
}
#[test]
fn coderank_partial_download_is_not_complete() {
let tmp = tempfile::TempDir::new().unwrap();
let dir = tmp.path().join(CODERANK_DIR);
fs::create_dir_all(&dir).unwrap();
fs::write(dir.join("model_int8.onnx"), b"stub").unwrap();
fs::write(dir.join("tokenizer.json"), b"stub").unwrap();
fs::write(dir.join("config.json"), b"stub").unwrap();
assert!(!is_coderank_downloaded(tmp.path()));
}
}