use std::path::{Path, PathBuf};
use basemyai_core::{CoreError, Device};
use serde::{Deserialize, Serialize};
use sha2::{Digest, Sha256};
use sysinfo::System;
pub const BASELINE_MODEL_ID: &str = "all-MiniLM-L6-v2";
pub const BASELINE_DIM: usize = 384;
const HF_BASE_URL: &str = "https://huggingface.co/sentence-transformers/all-MiniLM-L6-v2/resolve/main/";
const REQUIRED_MODEL_FILES: [&str; 3] = ["config.json", "tokenizer.json", "model.safetensors"];
const EXPECTED_SHA256: &[(&str, &str)] = &[
(
"config.json",
"953f9c0d463486b10a6871cc2fd59f223b2c70184f49815e7efbcab5d8908b41",
),
(
"tokenizer.json",
"be50c3628f2bf5bb5e3a7f17b1f74611b2561a3a27eeab05e5aa30f411572037",
),
(
"model.safetensors",
"53aa51172d142c89d9012cce15ae4d6cc0ca6895895114379cacb4fab128d9db",
),
];
#[derive(Debug, Clone)]
pub struct HardwareProfile {
pub total_ram_mb: u64,
pub gpu_vram_mb: Option<u64>,
pub cpu_cores: usize,
pub device: Device,
}
#[derive(Debug, Clone)]
pub struct ModelProvision {
pub model_id: String,
pub dim: usize,
pub model_path: PathBuf,
pub device: Device,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
enum PersistedDevice {
Cpu,
Cuda(usize),
Metal,
}
impl From<Device> for PersistedDevice {
fn from(d: Device) -> Self {
match d {
Device::Cpu => Self::Cpu,
Device::Cuda(i) => Self::Cuda(i),
Device::Metal => Self::Metal,
_ => Self::Cpu,
}
}
}
impl From<PersistedDevice> for Device {
fn from(d: PersistedDevice) -> Self {
match d {
PersistedDevice::Cpu => Self::Cpu,
PersistedDevice::Cuda(i) => Self::Cuda(i),
PersistedDevice::Metal => Self::Metal,
}
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
struct PersistedProvision {
model_id: String,
dim: usize,
model_path: PathBuf,
device: PersistedDevice,
}
impl From<&ModelProvision> for PersistedProvision {
fn from(p: &ModelProvision) -> Self {
Self {
model_id: p.model_id.clone(),
dim: p.dim,
model_path: p.model_path.clone(),
device: p.device.into(),
}
}
}
impl From<PersistedProvision> for ModelProvision {
fn from(p: PersistedProvision) -> Self {
Self {
model_id: p.model_id,
dim: p.dim,
model_path: p.model_path,
device: p.device.into(),
}
}
}
#[must_use]
pub fn detect_hardware() -> HardwareProfile {
let mut sys = System::new();
sys.refresh_memory();
let total_ram_mb = sys.total_memory() / (1024 * 1024);
let cpu_cores = std::thread::available_parallelism().map(usize::from).unwrap_or(1);
let gpu_vram_mb = detect_vram_mb();
let device = resolve_device();
HardwareProfile {
total_ram_mb,
gpu_vram_mb,
cpu_cores,
device,
}
}
pub async fn provision(consent_to_fetch: bool) -> crate::Result<ModelProvision> {
provision_inner(consent_to_fetch, |_, _| {}).await
}
pub async fn provision_with_progress(
consent_to_fetch: bool,
on_progress: impl Fn(u64, Option<u64>),
) -> crate::Result<ModelProvision> {
provision_inner(consent_to_fetch, on_progress).await
}
async fn provision_inner(
consent_to_fetch: bool,
on_progress: impl Fn(u64, Option<u64>),
) -> crate::Result<ModelProvision> {
if let Some(cached) = load_persisted_provision() {
return Ok(cached);
}
let hw = detect_hardware();
let model_path = baseline_cache_dir();
if model_present(&model_path) {
let result = ModelProvision {
model_id: BASELINE_MODEL_ID.to_string(),
dim: BASELINE_DIM,
model_path,
device: hw.device,
};
save_provision(&result);
return Ok(result);
}
if !consent_to_fetch {
return Err(CoreError::ModelNotProvisioned(format!(
"modèle '{BASELINE_MODEL_ID}' absent du cache ({}). Lancez le setup \
hardware-aware avec consentement explicite pour le récupérer.",
model_path.display()
))
.into());
}
fetch_model_files(&model_path, on_progress).await?;
let result = ModelProvision {
model_id: BASELINE_MODEL_ID.to_string(),
dim: BASELINE_DIM,
model_path,
device: hw.device,
};
save_provision(&result);
Ok(result)
}
async fn fetch_model_files(target_dir: &Path, on_progress: impl Fn(u64, Option<u64>)) -> crate::Result<()> {
std::fs::create_dir_all(target_dir).map_err(|e| {
CoreError::ModelNotProvisioned(format!(
"impossible de créer le dossier modèle {} : {e}",
target_dir.display()
))
})?;
let client = reqwest::Client::new();
for filename in REQUIRED_MODEL_FILES {
let url = format!("{HF_BASE_URL}{filename}");
let dest = target_dir.join(filename);
let expected = expected_sha256_for(filename);
download_and_verify(&client, &url, &dest, expected, &on_progress).await?;
}
Ok(())
}
async fn download_and_verify(
client: &reqwest::Client,
url: &str,
dest: &Path,
expected_sha256: Option<&str>,
on_progress: &impl Fn(u64, Option<u64>),
) -> crate::Result<()> {
let tmp = dest.with_extension("tmp");
let response = client
.get(url)
.send()
.await
.map_err(|e| CoreError::ModelNotProvisioned(format!("téléchargement échoué ({url}) : {e}")))?;
if !response.status().is_success() {
return Err(CoreError::ModelNotProvisioned(format!("HTTP {} pour {url}", response.status())).into());
}
let total = response.content_length();
let mut received = 0u64;
let mut hasher = Sha256::new();
let mut data: Vec<u8> = Vec::with_capacity(total.unwrap_or(0) as usize);
let mut response = response;
while let Some(chunk) = response
.chunk()
.await
.map_err(|e| CoreError::ModelNotProvisioned(format!("erreur stream ({url}) : {e}")))?
{
hasher.update(&chunk);
received += chunk.len() as u64;
data.extend_from_slice(&chunk);
on_progress(received, total);
}
let computed = format!("{:x}", hasher.finalize());
let sha_path = dest.with_extension("sha256");
match expected_sha256 {
Some(expected) => {
if computed != expected {
return Err(CoreError::ModelNotProvisioned(format!(
"SHA-256 mismatch pour {} : attendu {expected}, calculé {computed}",
dest.display()
))
.into());
}
}
None => {
if sha_path.exists() {
let stored = std::fs::read_to_string(&sha_path)
.map_err(|e| CoreError::ModelNotProvisioned(format!("lecture sha256 companion : {e}")))?;
if stored.trim() != computed {
return Err(CoreError::ModelNotProvisioned(format!(
"SHA-256 mismatch (companion) pour {} : stocké {}, calculé {computed}",
dest.display(),
stored.trim()
))
.into());
}
} else {
std::fs::write(&sha_path, &computed)
.map_err(|e| CoreError::ModelNotProvisioned(format!("écriture sha256 companion : {e}")))?;
}
}
}
std::fs::write(&tmp, &data)
.map_err(|e| CoreError::ModelNotProvisioned(format!("écriture tmp {} : {e}", tmp.display())))?;
std::fs::rename(&tmp, dest).map_err(|e| {
CoreError::ModelNotProvisioned(format!("renommage {} → {} : {e}", tmp.display(), dest.display()))
})?;
Ok(())
}
fn detect_vram_mb() -> Option<u64> {
detect_vram_nvidia_smi().or_else(detect_vram_platform)
}
fn detect_vram_nvidia_smi() -> Option<u64> {
let output = std::process::Command::new("nvidia-smi")
.args(["--query-gpu=memory.total", "--format=csv,noheader,nounits"])
.output()
.ok()?;
if !output.status.success() {
return None;
}
String::from_utf8_lossy(&output.stdout)
.lines()
.next()?
.trim()
.parse::<u64>()
.ok()
}
#[cfg(target_os = "macos")]
fn detect_vram_platform() -> Option<u64> {
detect_vram_macos()
}
#[cfg(not(target_os = "macos"))]
fn detect_vram_platform() -> Option<u64> {
None
}
#[cfg(target_os = "macos")]
fn detect_vram_macos() -> Option<u64> {
let output = std::process::Command::new("system_profiler")
.args(["SPDisplaysDataType", "-json"])
.output()
.ok()?;
let json: serde_json::Value = serde_json::from_slice(&output.stdout).ok()?;
let vram_str = json["SPDisplaysDataType"][0]["spdisplays_vram"].as_str()?;
parse_vram_mb(vram_str)
}
#[cfg(target_os = "macos")]
fn parse_vram_mb(s: &str) -> Option<u64> {
let s = s.trim().to_ascii_lowercase();
if let Some(n) = s.strip_suffix(" gb") {
n.trim().parse::<u64>().ok().map(|gb| gb * 1024)
} else if let Some(n) = s.strip_suffix(" mb") {
n.trim().parse::<u64>().ok()
} else {
None
}
}
#[must_use]
fn resolve_device() -> Device {
if cuda_available() {
Device::Cuda(0)
} else if metal_available() {
Device::Metal
} else {
Device::Cpu
}
}
fn cuda_available() -> bool {
std::env::var_os("CUDA_PATH").is_some() || std::env::var_os("CUDA_HOME").is_some()
}
fn metal_available() -> bool {
cfg!(target_os = "macos")
}
#[must_use]
fn baseline_cache_dir() -> PathBuf {
let base = dirs::cache_dir()
.or_else(dirs::home_dir)
.or_else(|| std::env::var_os("USERPROFILE").map(PathBuf::from))
.or_else(|| std::env::var_os("HOME").map(PathBuf::from))
.unwrap_or_else(|| PathBuf::from("."));
base.join("basemyai").join("models").join(BASELINE_MODEL_ID)
}
#[must_use]
fn model_present(dir: &Path) -> bool {
REQUIRED_MODEL_FILES.iter().all(|f| dir.join(f).is_file())
}
fn expected_sha256_for(filename: &str) -> Option<&'static str> {
EXPECTED_SHA256.iter().find(|(f, _)| *f == filename).map(|(_, h)| *h)
}
fn provision_config_path() -> PathBuf {
dirs::data_dir()
.or_else(dirs::home_dir)
.or_else(|| std::env::var_os("USERPROFILE").map(PathBuf::from))
.or_else(|| std::env::var_os("HOME").map(PathBuf::from))
.unwrap_or_else(|| PathBuf::from("."))
.join("basemyai")
.join("provision.json")
}
fn load_persisted_provision() -> Option<ModelProvision> {
let text = std::fs::read_to_string(provision_config_path()).ok()?;
let p: PersistedProvision = serde_json::from_str(&text).ok()?;
let result: ModelProvision = p.into();
if model_present(&result.model_path) {
Some(result)
} else {
None
}
}
fn save_provision(provision: &ModelProvision) {
let path = provision_config_path();
if let Some(parent) = path.parent() {
let _ = std::fs::create_dir_all(parent);
}
if let Ok(json) = serde_json::to_string_pretty(&PersistedProvision::from(provision)) {
let _ = std::fs::write(path, json);
}
}