Improve local inference settings and model downloader (#8467)
Signed-off-by: jh-block <jhugo@block.xyz>
This commit is contained in:
@@ -13,11 +13,11 @@ use goose::download_manager::{get_download_manager, DownloadProgress};
|
||||
use goose::providers::local_inference::hf_models::{self, HfModelInfo, HfQuantVariant};
|
||||
use goose::providers::local_inference::{
|
||||
available_inference_memory_bytes,
|
||||
hf_models::{resolve_model_spec, HfGgufFile},
|
||||
hf_models::{resolve_model_spec, resolve_model_spec_full, HfGgufFile},
|
||||
local_model_registry::{
|
||||
default_settings_for_model, featured_mmproj_spec, get_registry, is_featured_model,
|
||||
model_id_from_repo, LocalModelEntry, ModelDownloadStatus as RegistryDownloadStatus,
|
||||
ModelSettings, FEATURED_MODELS,
|
||||
ModelSettings, ShardFile, FEATURED_MODELS,
|
||||
},
|
||||
recommend_local_model,
|
||||
};
|
||||
@@ -130,6 +130,7 @@ async fn ensure_featured_models_in_registry() -> Result<(), ErrorResponse> {
|
||||
mmproj_path: None,
|
||||
mmproj_source_url: None,
|
||||
mmproj_size_bytes: 0,
|
||||
shard_files: vec![],
|
||||
});
|
||||
}
|
||||
|
||||
@@ -284,6 +285,8 @@ pub struct SearchQuery {
|
||||
pub struct RepoVariantsResponse {
|
||||
pub variants: Vec<HfQuantVariant>,
|
||||
pub recommended_index: Option<usize>,
|
||||
pub available_memory_bytes: u64,
|
||||
pub downloaded_quants: Vec<String>,
|
||||
}
|
||||
|
||||
#[utoipa::path(
|
||||
@@ -327,9 +330,23 @@ pub async fn get_repo_files(
|
||||
let available_memory = available_inference_memory_bytes(&state.inference_runtime);
|
||||
let recommended_index = hf_models::recommend_variant(&variants, available_memory);
|
||||
|
||||
let downloaded_quants = {
|
||||
let registry = get_registry()
|
||||
.lock()
|
||||
.map_err(|_| ErrorResponse::internal("Failed to acquire registry lock"))?;
|
||||
registry
|
||||
.list_models()
|
||||
.iter()
|
||||
.filter(|m| m.repo_id == repo_id && m.is_downloaded())
|
||||
.map(|m| m.quantization.clone())
|
||||
.collect()
|
||||
};
|
||||
|
||||
Ok(Json(RepoVariantsResponse {
|
||||
variants,
|
||||
recommended_index,
|
||||
available_memory_bytes: available_memory,
|
||||
downloaded_quants,
|
||||
}))
|
||||
}
|
||||
|
||||
@@ -354,26 +371,44 @@ pub async fn download_hf_model(
|
||||
let (repo_id, quantization) = hf_models::parse_model_spec(&req.spec)
|
||||
.map_err(|e| ErrorResponse::bad_request(format!("Invalid spec format: {e}")))?;
|
||||
|
||||
let (_repo, hf_file) = resolve_model_spec(&req.spec)
|
||||
let (_repo, resolved) = resolve_model_spec_full(&req.spec)
|
||||
.await
|
||||
.map_err(|e| ErrorResponse::bad_request(format!("Invalid spec: {}", e)))?;
|
||||
|
||||
let model_id = model_id_from_repo(&repo_id, &quantization);
|
||||
let local_path = Paths::in_data_dir("models").join(&hf_file.filename);
|
||||
let download_url = hf_file.download_url.clone();
|
||||
let models_dir = Paths::in_data_dir("models");
|
||||
let first_file = &resolved.files[0];
|
||||
let first_local_path = models_dir.join(&first_file.filename);
|
||||
|
||||
let shard_files: Vec<ShardFile> = if resolved.files.len() > 1 {
|
||||
resolved
|
||||
.files
|
||||
.iter()
|
||||
.skip(1)
|
||||
.map(|f| ShardFile {
|
||||
filename: f.filename.clone(),
|
||||
local_path: models_dir.join(&f.filename),
|
||||
source_url: f.download_url.clone(),
|
||||
size_bytes: f.size_bytes,
|
||||
})
|
||||
.collect()
|
||||
} else {
|
||||
vec![]
|
||||
};
|
||||
|
||||
let entry = LocalModelEntry {
|
||||
id: model_id.clone(),
|
||||
repo_id,
|
||||
filename: hf_file.filename,
|
||||
filename: first_file.filename.clone(),
|
||||
quantization,
|
||||
local_path: local_path.clone(),
|
||||
source_url: download_url.clone(),
|
||||
local_path: first_local_path.clone(),
|
||||
source_url: first_file.download_url.clone(),
|
||||
settings: default_settings_for_model(&model_id),
|
||||
size_bytes: hf_file.size_bytes,
|
||||
size_bytes: resolved.total_size,
|
||||
mmproj_path: None,
|
||||
mmproj_source_url: None,
|
||||
mmproj_size_bytes: 0,
|
||||
shard_files: shard_files.clone(),
|
||||
};
|
||||
|
||||
// add_model enriches the entry with mmproj metadata from the featured table
|
||||
@@ -393,10 +428,16 @@ pub async fn download_hf_model(
|
||||
};
|
||||
|
||||
let dm = get_download_manager();
|
||||
dm.download_model(
|
||||
let all_files: Vec<(String, std::path::PathBuf)> = resolved
|
||||
.files
|
||||
.iter()
|
||||
.map(|f| (f.download_url.clone(), models_dir.join(&f.filename)))
|
||||
.collect();
|
||||
|
||||
dm.download_model_sharded(
|
||||
format!("{}-model", model_id),
|
||||
download_url,
|
||||
local_path,
|
||||
all_files,
|
||||
resolved.total_size,
|
||||
None,
|
||||
)
|
||||
.await
|
||||
@@ -470,28 +511,39 @@ pub async fn cancel_local_model_download(
|
||||
)
|
||||
)]
|
||||
pub async fn delete_local_model(Path(model_id): Path<String>) -> Result<StatusCode, ErrorResponse> {
|
||||
let (local_path, mmproj_path, other_uses_mmproj) = {
|
||||
let (all_paths, primary_path, mmproj_path, other_uses_mmproj) = {
|
||||
let registry = get_registry()
|
||||
.lock()
|
||||
.map_err(|_| ErrorResponse::internal("Failed to acquire registry lock"))?;
|
||||
let entry = registry
|
||||
.get_model(&model_id)
|
||||
.ok_or_else(|| ErrorResponse::not_found("Model not found"))?;
|
||||
let lp = entry.local_path.clone();
|
||||
let paths: Vec<std::path::PathBuf> =
|
||||
entry.all_local_paths().map(|p| p.to_path_buf()).collect();
|
||||
let primary = entry.local_path.clone();
|
||||
let mp = entry.mmproj_path.clone();
|
||||
// Check if another downloaded model shares this mmproj file
|
||||
let shared = mp.as_ref().is_some_and(|target| {
|
||||
registry.list_models().iter().any(|m| {
|
||||
m.id != model_id && m.is_downloaded() && m.mmproj_path.as_ref() == Some(target)
|
||||
})
|
||||
});
|
||||
(lp, mp, shared)
|
||||
(paths, primary, mp, shared)
|
||||
};
|
||||
|
||||
if local_path.exists() {
|
||||
tokio::fs::remove_file(&local_path)
|
||||
.await
|
||||
.map_err(|e| ErrorResponse::internal(format!("Failed to delete: {}", e)))?;
|
||||
for path in &all_paths {
|
||||
if path.exists() {
|
||||
tokio::fs::remove_file(path)
|
||||
.await
|
||||
.map_err(|e| ErrorResponse::internal(format!("Failed to delete: {}", e)))?;
|
||||
}
|
||||
}
|
||||
|
||||
// Clean up empty parent directories (e.g. BF16/ subdirectory)
|
||||
if let Some(parent) = primary_path.parent() {
|
||||
let models_dir = Paths::in_data_dir("models");
|
||||
if parent != models_dir {
|
||||
let _ = tokio::fs::remove_dir(parent).await;
|
||||
}
|
||||
}
|
||||
|
||||
if !other_uses_mmproj {
|
||||
@@ -563,7 +615,23 @@ pub async fn update_model_settings(
|
||||
}
|
||||
|
||||
pub fn routes(state: Arc<AppState>) -> Router {
|
||||
goose::download_manager::cleanup_partial_downloads(&Paths::in_data_dir("models"));
|
||||
let registered_paths: std::collections::HashSet<std::path::PathBuf> = get_registry()
|
||||
.lock()
|
||||
.map(|reg| {
|
||||
reg.list_models()
|
||||
.iter()
|
||||
.flat_map(|m| {
|
||||
m.all_local_paths()
|
||||
.map(|p| p.to_path_buf())
|
||||
.chain(m.mmproj_path.as_deref().map(|p| p.to_path_buf()))
|
||||
})
|
||||
.collect()
|
||||
})
|
||||
.unwrap_or_default();
|
||||
goose::download_manager::cleanup_partial_downloads(
|
||||
&Paths::in_data_dir("models"),
|
||||
®istered_paths,
|
||||
);
|
||||
|
||||
Router::new()
|
||||
.route("/local-inference/models", get(list_local_models))
|
||||
|
||||
Reference in New Issue
Block a user