Improve local inference settings and model downloader (#8467)

Signed-off-by: jh-block <jhugo@block.xyz>
This commit is contained in:
jh-block
2026-04-14 10:50:50 +02:00
committed by GitHub
parent 2e529b7497
commit 65e5d97cac
10 changed files with 1009 additions and 297 deletions
+1
View File
@@ -1607,6 +1607,7 @@ async fn handle_local_models_command(command: LocalModelsCommand) -> Result<()>
mmproj_path: None,
mmproj_source_url: None,
mmproj_size_bytes: 0,
shard_files: vec![],
};
{
@@ -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"),
&registered_paths,
);
Router::new()
.route("/local-inference/models", get(list_local_models))
+370 -77
View File
@@ -16,14 +16,37 @@ fn partial_path_for(destination: &Path) -> PathBuf {
)
}
/// Remove any leftover `.part` files in the given directory.
pub fn cleanup_partial_downloads(dir: &Path) {
/// Remove orphaned `.part` files in the given directory (and one level of subdirectories).
/// Preserves `.part` files whose final destination is in `registered_paths` so that
/// in-progress shard downloads can resume after a restart.
pub fn cleanup_partial_downloads(
dir: &Path,
registered_paths: &std::collections::HashSet<PathBuf>,
) {
let should_keep = |part_path: &Path| -> bool {
// Derive the final path by stripping the trailing ".part" extension
let final_path = part_path.with_extension("");
registered_paths.contains(&final_path)
};
if let Ok(entries) = std::fs::read_dir(dir) {
for entry in entries.flatten() {
let path = entry.path();
if path.extension().is_some_and(|e| e == "part") {
if path.extension().is_some_and(|e| e == "part") && !should_keep(&path) {
let _ = std::fs::remove_file(&path);
}
if path.is_dir() {
if let Ok(sub_entries) = std::fs::read_dir(&path) {
for sub in sub_entries.flatten() {
let sub_path = sub.path();
if sub_path.extension().is_some_and(|e| e == "part")
&& !should_keep(&sub_path)
{
let _ = std::fs::remove_file(&sub_path);
}
}
}
}
}
}
}
@@ -104,7 +127,18 @@ impl DownloadManager {
destination: PathBuf,
on_complete: Option<Box<dyn FnOnce() + Send + 'static>>,
) -> Result<()> {
info!(model_id = %model_id, url = %url, destination = ?destination, "Starting model download");
self.download_model_sharded(model_id, vec![(url, destination)], 0, on_complete)
.await
}
pub async fn download_model_sharded(
&self,
model_id: String,
files: Vec<(String, PathBuf)>,
total_size_hint: u64,
on_complete: Option<Box<dyn FnOnce() + Send + 'static>>,
) -> Result<()> {
info!(model_id = %model_id, file_count = files.len(), "Starting model download");
{
let mut downloads = self
.downloads
@@ -128,7 +162,7 @@ impl DownloadManager {
model_id: model_id.clone(),
status: DownloadStatus::Downloading,
bytes_downloaded: 0,
total_bytes: 0,
total_bytes: total_size_hint,
progress_percent: 0.0,
speed_bps: None,
eta_seconds: None,
@@ -138,21 +172,24 @@ impl DownloadManager {
);
}
// Create parent directory if it doesn't exist
if let Some(parent) = destination.parent() {
tokio::fs::create_dir_all(parent)
.await
.map_err(|e| anyhow::anyhow!("Failed to create directory: {}", e))?;
// Create parent directories for all files
for (_, dest) in &files {
if let Some(parent) = dest.parent() {
tokio::fs::create_dir_all(parent)
.await
.map_err(|e| anyhow::anyhow!("Failed to create directory: {}", e))?;
}
}
let downloads = self.downloads.clone();
let model_id_clone = model_id.clone();
let files_for_cleanup: Vec<PathBuf> = files.iter().map(|(_, d)| d.clone()).collect();
let destination_for_cleanup = destination.clone();
// Download in background task
tokio::spawn(async move {
match Self::download_file(&url, &destination, &downloads, &model_id_clone).await {
let result =
Self::download_files_sequentially(&files, &downloads, &model_id_clone).await;
match result {
Ok(_) => {
info!(model_id = %model_id_clone, "Download completed successfully");
if let Ok(mut downloads) = downloads.lock() {
@@ -168,9 +205,10 @@ impl DownloadManager {
}
}
Err(e) => {
// Clean up partial file on failure
let partial = partial_path_for(&destination_for_cleanup);
let _ = tokio::fs::remove_file(&partial).await;
for dest in &files_for_cleanup {
let partial = partial_path_for(dest);
let _ = tokio::fs::remove_file(&partial).await;
}
if let Ok(mut downloads) = downloads.lock() {
if let Some(progress) = downloads.get_mut(&model_id_clone) {
@@ -188,93 +226,348 @@ impl DownloadManager {
Ok(())
}
async fn download_file(
url: &str,
destination: &PathBuf,
const MAX_RETRIES: u32 = 10;
const RETRY_BASE_DELAY: std::time::Duration = std::time::Duration::from_secs(2);
const RETRY_MAX_DELAY: std::time::Duration = std::time::Duration::from_secs(60);
async fn cancellable_sleep(
delay: std::time::Duration,
downloads: &DownloadMap,
model_id: &str,
) -> Result<(), anyhow::Error> {
let check_interval = std::time::Duration::from_millis(500);
let start = std::time::Instant::now();
while start.elapsed() < delay {
if Self::is_cancelled(downloads, model_id) {
anyhow::bail!("Download cancelled");
}
let remaining = delay.saturating_sub(start.elapsed());
tokio::time::sleep(std::cmp::min(check_interval, remaining)).await;
}
Ok(())
}
fn is_cancelled(downloads: &DownloadMap, model_id: &str) -> bool {
if let Ok(downloads) = downloads.lock() {
if let Some(progress) = downloads.get(model_id) {
return progress.status == DownloadStatus::Cancelled;
}
}
false
}
#[allow(clippy::too_many_arguments)]
/// Download multiple files sequentially, tracking cumulative progress under one model_id.
async fn download_files_sequentially(
files: &[(String, PathBuf)],
downloads: &DownloadMap,
model_id: &str,
) -> Result<(), anyhow::Error> {
let client = reqwest::Client::builder()
.connect_timeout(std::time::Duration::from_secs(30))
.read_timeout(std::time::Duration::from_secs(60))
.read_timeout(std::time::Duration::from_secs(120))
.build()?;
let mut response = client.get(url).send().await?;
if !response.status().is_success() {
anyhow::bail!("Failed to download: HTTP {}", response.status());
// HEAD each file to get accurate total size. Only replace the hint if
// every file returned a size; partial results would underestimate.
let mut total: u64 = 0;
let mut all_resolved = true;
for (url, _) in files {
let size = client
.head(url)
.send()
.await
.ok()
.and_then(|r| r.content_length())
.unwrap_or(0);
if size == 0 {
all_resolved = false;
}
total += size;
}
let total_bytes = response.content_length().unwrap_or(0);
{
if let Ok(mut downloads) = downloads.lock() {
if let Some(progress) = downloads.get_mut(model_id) {
progress.total_bytes = total_bytes;
if all_resolved && total > 0 {
if let Ok(mut dl) = downloads.lock() {
if let Some(progress) = dl.get_mut(model_id) {
progress.total_bytes = total;
}
}
}
let partial_path = partial_path_for(destination);
let mut file = tokio::fs::File::create(&partial_path).await?;
let mut bytes_downloaded = 0u64;
let start_time = std::time::Instant::now();
while let Some(chunk) = response.chunk().await? {
// Check if cancelled
let should_cancel = {
if let Ok(downloads) = downloads.lock() {
if let Some(progress) = downloads.get(model_id) {
progress.status == DownloadStatus::Cancelled
} else {
false
}
} else {
false
let mut cumulative_bytes: u64 = 0;
// Account for already-downloaded shards
for (_, dest) in files {
let partial = partial_path_for(dest);
if dest.exists() {
if let Ok(meta) = tokio::fs::metadata(dest).await {
cumulative_bytes += meta.len();
}
};
} else if partial.exists() {
if let Ok(meta) = tokio::fs::metadata(&partial).await {
cumulative_bytes += meta.len();
}
}
}
let bytes_at_start = cumulative_bytes;
if should_cancel {
for (url, destination) in files {
if Self::is_cancelled(downloads, model_id) {
anyhow::bail!("Download cancelled");
}
// Skip already-completed shards
if destination.exists() {
continue;
}
Self::download_one_file(
&client,
url,
destination,
downloads,
model_id,
&mut cumulative_bytes,
start_time,
bytes_at_start,
)
.await?;
}
Ok(())
}
#[allow(clippy::too_many_arguments)]
async fn download_one_file(
client: &reqwest::Client,
url: &str,
destination: &Path,
downloads: &DownloadMap,
model_id: &str,
cumulative_bytes: &mut u64,
start_time: std::time::Instant,
bytes_at_start: u64,
) -> Result<(), anyhow::Error> {
let partial_path = partial_path_for(destination);
let mut retries = 0u32;
let mut file_bytes: u64 = if partial_path.exists() {
tokio::fs::metadata(&partial_path).await?.len()
} else {
0
};
// Get this file's total size
let mut file_total: u64 = client
.head(url)
.send()
.await
.ok()
.and_then(|r| r.content_length())
.unwrap_or(0);
// If partial matches expected size exactly, promote it
if file_total > 0 && file_bytes == file_total {
tokio::fs::rename(&partial_path, destination).await?;
// cumulative_bytes already accounts for this file from the pre-scan
return Ok(());
}
// If partial is oversized or remote changed, discard and re-download
if file_total > 0 && file_bytes > file_total {
info!(model_id = %model_id, file_bytes, file_total, "Partial file oversized, re-downloading");
*cumulative_bytes = cumulative_bytes.saturating_sub(file_bytes);
file_bytes = 0;
let _ = tokio::fs::remove_file(&partial_path).await;
}
loop {
if Self::is_cancelled(downloads, model_id) {
let _ = tokio::fs::remove_file(&partial_path).await;
anyhow::bail!("Download cancelled");
}
file.write_all(&chunk).await?;
bytes_downloaded += chunk.len() as u64;
let mut request = client.get(url);
if file_bytes > 0 {
request = request.header("Range", format!("bytes={}-", file_bytes));
}
// Update progress
let elapsed = start_time.elapsed().as_secs_f64();
let speed_bps = if elapsed > 0.0 {
Some((bytes_downloaded as f64 / elapsed) as u64)
} else {
None
};
let eta_seconds = if let Some(speed) = speed_bps {
if speed > 0 && total_bytes > 0 {
Some(total_bytes.saturating_sub(bytes_downloaded) / speed)
} else {
None
let response = match request.send().await {
Ok(r) => r,
Err(e) => {
if retries >= Self::MAX_RETRIES {
anyhow::bail!("Download failed after {} retries: {}", retries, e);
}
retries += 1;
let delay = std::cmp::min(
Self::RETRY_BASE_DELAY * 2u32.saturating_pow(retries - 1),
Self::RETRY_MAX_DELAY,
);
info!(model_id = %model_id, retry = retries, delay_secs = ?delay.as_secs(), error = %e, "Retrying download after connection error");
Self::cancellable_sleep(delay, downloads, model_id).await?;
continue;
}
} else {
None
};
if let Ok(mut downloads) = downloads.lock() {
if let Some(progress) = downloads.get_mut(model_id) {
progress.bytes_downloaded = bytes_downloaded;
progress.progress_percent = if total_bytes > 0 {
(bytes_downloaded as f64 / total_bytes as f64 * 100.0) as f32
} else {
0.0
};
progress.speed_bps = speed_bps;
progress.eta_seconds = eta_seconds;
let status = response.status();
if status == reqwest::StatusCode::RANGE_NOT_SATISFIABLE {
if file_total > 0 && file_bytes == file_total {
break;
}
*cumulative_bytes = cumulative_bytes.saturating_sub(file_bytes);
file_bytes = 0;
let _ = tokio::fs::remove_file(&partial_path).await;
continue;
}
if !status.is_success() && status != reqwest::StatusCode::PARTIAL_CONTENT {
let is_transient = status.is_server_error()
|| status == reqwest::StatusCode::REQUEST_TIMEOUT
|| status == reqwest::StatusCode::TOO_MANY_REQUESTS;
if !is_transient || retries >= Self::MAX_RETRIES {
anyhow::bail!("Failed to download: HTTP {}", status);
}
retries += 1;
let delay = std::cmp::min(
Self::RETRY_BASE_DELAY * 2u32.saturating_pow(retries - 1),
Self::RETRY_MAX_DELAY,
);
info!(model_id = %model_id, retry = retries, http_status = %status, "Retrying download after transient HTTP error");
Self::cancellable_sleep(delay, downloads, model_id).await?;
continue;
}
if file_bytes > 0 && status == reqwest::StatusCode::OK {
info!(model_id = %model_id, "Server ignored Range header, restarting file from scratch");
// Subtract already-counted partial bytes from cumulative
*cumulative_bytes = cumulative_bytes.saturating_sub(file_bytes);
file_bytes = 0;
let _ = tokio::fs::remove_file(&partial_path).await;
}
// If HEAD didn't return this file's size, learn it from the GET response.
// This block only fires once per file (file_total stays non-zero after),
// so retries don't double-count. Since download_files_sequentially's HEAD
// pass contributed 0 for this file, we add the discovered size to the
// shared total so progress/ETA are accurate.
if file_total == 0 {
let new_file_total = if file_bytes > 0 {
response
.headers()
.get("content-range")
.and_then(|v| v.to_str().ok())
.and_then(|s| s.rsplit('/').next())
.and_then(|s| s.parse::<u64>().ok())
} else {
response.content_length()
};
if let Some(t) = new_file_total {
file_total = t;
if let Ok(mut dl) = downloads.lock() {
if let Some(progress) = dl.get_mut(model_id) {
progress.total_bytes = progress.total_bytes.saturating_add(t);
}
}
}
}
let mut file = tokio::fs::OpenOptions::new()
.create(true)
.append(true)
.open(&partial_path)
.await?;
let file_len = tokio::fs::metadata(&partial_path).await?.len();
if file_len != file_bytes {
file.set_len(file_bytes).await?;
}
let mut stream_error = false;
let mut resp = response;
loop {
let chunk_result = resp.chunk().await;
match chunk_result {
Ok(Some(chunk)) => {
if Self::is_cancelled(downloads, model_id) {
let _ = tokio::fs::remove_file(&partial_path).await;
anyhow::bail!("Download cancelled");
}
file.write_all(&chunk).await?;
let chunk_len = chunk.len() as u64;
file_bytes += chunk_len;
*cumulative_bytes += chunk_len;
let elapsed = start_time.elapsed().as_secs_f64();
let bytes_this_session = cumulative_bytes.saturating_sub(bytes_at_start);
let speed_bps = if elapsed > 0.0 {
Some((bytes_this_session as f64 / elapsed) as u64)
} else {
None
};
let current_total = if let Ok(dl) = downloads.lock() {
dl.get(model_id).map(|p| p.total_bytes).unwrap_or(0)
} else {
0
};
let eta_seconds = if let Some(speed) = speed_bps {
if speed > 0 && current_total > 0 {
Some(current_total.saturating_sub(*cumulative_bytes) / speed)
} else {
None
}
} else {
None
};
if let Ok(mut dl) = downloads.lock() {
if let Some(progress) = dl.get_mut(model_id) {
progress.bytes_downloaded = *cumulative_bytes;
progress.progress_percent = if current_total > 0 {
(*cumulative_bytes as f64 / current_total as f64 * 100.0) as f32
} else {
0.0
};
progress.speed_bps = speed_bps;
progress.eta_seconds = eta_seconds;
}
}
}
Ok(None) => break,
Err(e) => {
info!(model_id = %model_id, bytes = *cumulative_bytes, error = %e, "Download stream interrupted, will retry");
stream_error = true;
break;
}
}
}
file.flush().await?;
drop(file);
if stream_error {
if retries >= Self::MAX_RETRIES {
anyhow::bail!(
"Download failed after {} retries due to stream interruption",
retries
);
}
retries += 1;
let delay = std::cmp::min(
Self::RETRY_BASE_DELAY * 2u32.saturating_pow(retries - 1),
Self::RETRY_MAX_DELAY,
);
info!(model_id = %model_id, retry = retries, delay_secs = ?delay.as_secs(), "Retrying download with resume");
Self::cancellable_sleep(delay, downloads, model_id).await?;
continue;
}
break;
}
file.flush().await?;
drop(file);
tokio::fs::rename(&partial_path, destination).await?;
Ok(())
}
@@ -33,6 +33,15 @@ pub struct HfQuantVariant {
pub download_url: String,
pub description: &'static str,
pub quality_rank: u8,
#[serde(default)]
pub sharded: bool,
}
/// Result of resolving a model spec — may contain multiple shard files.
#[derive(Debug, Clone)]
pub struct ResolvedModel {
pub files: Vec<HfGgufFile>,
pub total_size: u64,
}
#[derive(Debug, Deserialize)]
@@ -55,52 +64,71 @@ struct QuantInfo {
quality_rank: u8,
}
// quality_rank groups quants by bit-level so that all N-bit variants sort
// together. Within a group, higher rank = higher quality.
//
// 1-bit: 1019 4-bit: 4049 8-bit: 8089
// 2-bit: 2029 5-bit: 5059 16-bit: 9094
// 3-bit: 3039 6-bit: 6069 32-bit: 9599
//
const QUANT_TABLE: &[(&str, &str, u8)] = &[
("IQ1_S", "Extremely small, very low quality", 1),
("IQ1_M", "Extremely small, very low quality", 2),
("IQ2_XXS", "Very small, low quality", 3),
("IQ2_XS", "Very small, low quality", 4),
("IQ2_S", "Very small, low quality", 5),
("IQ2_M", "Very small, low quality", 6),
("Q2_K", "Small, low quality", 7),
("Q2_K_S", "Small, low quality", 7),
("IQ3_XXS", "Very small, moderate quality loss", 8),
("IQ3_XS", "Small, moderate quality loss", 9),
("IQ3_S", "Small, moderate quality loss", 9),
("Q3_K_S", "Small, moderate quality loss", 10),
("IQ3_M", "Small, moderate quality loss", 11),
("Q3_K_M", "Small, balanced quality/size", 12),
("Q3_K_L", "Medium-small, decent quality", 13),
("IQ4_XS", "Medium, good quality", 14),
("IQ4_NL", "Medium, good quality", 15),
("Q4_0", "Medium, good quality", 16),
("Q4_1", "Medium, good quality", 17),
("Q4_K_S", "Medium, good quality/size balance", 18),
// 1-bit
("TQ1_0", "Tiny, ternary quantization", 10),
("IQ1_S", "Extremely small, very low quality", 11),
("IQ1_M", "Extremely small, very low quality", 12),
// 2-bit
("IQ2_XXS", "Very small, low quality", 20),
("IQ2_XS", "Very small, low quality", 21),
("IQ2_S", "Very small, low quality", 22),
("IQ2_M", "Very small, low quality", 23),
("Q2_K", "Small, low quality", 24),
("Q2_K_S", "Small, low quality", 24),
("Q2_K_L", "Small, low quality", 25),
("Q2_K_XL", "Small, low quality", 26),
// 3-bit
("IQ3_XXS", "Very small, moderate quality loss", 30),
("IQ3_XS", "Small, moderate quality loss", 31),
("IQ3_S", "Small, moderate quality loss", 32),
("IQ3_M", "Small, moderate quality loss", 33),
("Q3_K_S", "Small, moderate quality loss", 34),
("Q3_K_M", "Small, balanced quality/size", 35),
("Q3_K_L", "Medium-small, decent quality", 36),
("Q3_K_XL", "Medium-small, decent quality", 37),
// 4-bit
("IQ4_XS", "Medium, good quality", 40),
("IQ4_NL", "Medium, good quality", 41),
("Q4_0", "Medium, good quality", 42),
("Q4_1", "Medium, good quality", 43),
("Q4_K_S", "Medium, good quality/size balance", 44),
(
"Q4_K_M",
"Medium, recommended balance of quality and size",
19,
45,
),
("Q5_0", "Medium-large, high quality", 20),
("Q5_1", "Medium-large, high quality", 21),
("Q5_K_S", "Medium-large, high quality", 22),
("Q5_K_M", "Medium-large, very high quality", 23),
("Q6_K", "Large, near-lossless quality", 24),
("Q8_0", "Large, near-lossless quality", 25),
("F16", "Full size, original quality (16-bit)", 26),
("BF16", "Full size, original quality (bfloat16)", 27),
("F32", "Full size, original quality (32-bit)", 28),
("Q4_K_L", "Medium, good quality", 46),
("Q4_K_XL", "Medium, good quality", 47),
(
"MXFP4_MOE",
"Medium, mixed-precision 4-bit for MoE models",
18,
48,
),
("TQ1_0", "Tiny, ternary quantization", 1),
("Q2_K_XL", "Extended-layer variant", 15),
("Q3_K_XL", "Extended-layer variant", 15),
("Q4_K_XL", "Extended-layer variant", 15),
("Q2_K_L", "Small, low quality (large variant)", 8),
("Q4_K_L", "Medium, good quality (large variant)", 20),
// 5-bit
("Q5_0", "Medium-large, high quality", 50),
("Q5_1", "Medium-large, high quality", 51),
("Q5_K_S", "Medium-large, high quality", 52),
("Q5_K_M", "Medium-large, very high quality", 53),
("Q5_K_XL", "Medium-large, very high quality", 54),
// 6-bit
("Q6_K", "Large, near-lossless quality", 60),
("Q6_K_XL", "Large, near-lossless quality", 61),
// 8-bit
("Q8_0", "Large, near-lossless quality", 80),
("Q8_K_XL", "Large, near-lossless quality", 81),
// 16-bit
("F16", "Full size, original quality (16-bit)", 90),
("BF16", "Full size, original quality (bfloat16)", 91),
// 32-bit
("F32", "Full size, original quality (32-bit)", 95),
];
fn quant_info(quant: &str) -> QuantInfo {
@@ -113,7 +141,7 @@ fn quant_info(quant: &str) -> QuantInfo {
})
.unwrap_or(QuantInfo {
description: "",
quality_rank: 15,
quality_rank: 45,
})
}
@@ -168,47 +196,132 @@ fn looks_like_quant(s: &str) -> bool {
fn is_shard_file(filename: &str) -> bool {
// Matches patterns like "-00001-of-00003.gguf"
parse_shard_index(filename).is_some()
}
/// Parse the shard index (1-based) from a filename like "model-BF16-00001-of-00002.gguf".
fn parse_shard_index(filename: &str) -> Option<u32> {
let basename = filename.rsplit('/').next().unwrap_or(filename);
let stem = basename.trim_end_matches(".gguf");
if let Some(pos) = stem.rfind("-of-") {
stem.get(..pos)
.and_then(|before| before.rsplit('-').next())
.map(|s| !s.is_empty() && s.chars().all(|c| c.is_ascii_digit()))
.unwrap_or(false)
let pos = stem.rfind("-of-")?;
let before = stem.get(..pos)?;
let idx_str = before.rsplit('-').next()?;
if !idx_str.is_empty() && idx_str.chars().all(|c| c.is_ascii_digit()) {
idx_str.parse().ok()
} else {
false
None
}
}
/// Parse the total shard count from a filename like "model-BF16-00001-of-00002.gguf".
fn parse_shard_total(filename: &str) -> Option<u32> {
let basename = filename.rsplit('/').next().unwrap_or(filename);
let stem = basename.trim_end_matches(".gguf");
let pos = stem.rfind("-of-")?;
let total_str = stem.get(pos + 4..)?;
total_str.parse().ok()
}
fn build_download_url(repo_id: &str, filename: &str) -> String {
format!("{}/{}/resolve/main/{}", HF_DOWNLOAD_BASE, repo_id, filename)
}
/// Collect single-file GGUFs into quantization variants (sharded files are excluded).
/// Derive the expected model filename stem from a repo_id.
/// e.g. "unsloth/gemma-4-26B-A4B-it-GGUF" → "gemma-4-26b-a4b-it" (lowercased)
fn model_stem_from_repo(repo_id: &str) -> String {
let repo_name = repo_id.rsplit('/').next().unwrap_or(repo_id);
let stem = repo_name
.strip_suffix("-GGUF")
.or_else(|| repo_name.strip_suffix("-gguf"))
.unwrap_or(repo_name);
stem.to_lowercase()
}
/// Check whether a GGUF file belongs to the main model (vs auxiliary files like mmproj).
/// Matches files whose basename starts with the model stem derived from the repo name.
fn is_model_file(filename: &str, model_stem_lower: &str) -> bool {
let basename = filename.rsplit('/').next().unwrap_or(filename);
basename.to_lowercase().starts_with(model_stem_lower)
}
/// Collect GGUF files into quantization variants.
/// Single-file quants use the file directly.
/// Sharded quants (multiple files for one quantization) aggregate sizes and use the
/// first shard filename as the representative — the download path must handle all shards.
fn group_into_variants(repo_id: &str, files: Vec<HfApiSibling>) -> Vec<HfQuantVariant> {
let mut variants: Vec<HfQuantVariant> = files
use std::collections::HashMap;
let stem = model_stem_from_repo(repo_id);
let gguf_files: Vec<_> = files
.into_iter()
.filter(|s| {
s.rfilename.ends_with(".gguf")
&& !is_shard_file(&s.rfilename)
&& is_model_file(&s.rfilename, &stem)
&& parse_quantization(&s.rfilename) != "unknown"
})
.map(|s| {
let quant = parse_quantization(&s.rfilename);
let info = quant_info(&quant);
let download_url = build_download_url(repo_id, &s.rfilename);
HfQuantVariant {
quantization: quant,
size_bytes: s.size.unwrap_or(0),
filename: s.rfilename,
download_url,
description: info.description,
quality_rank: info.quality_rank,
}
})
.collect();
variants.sort_by_key(|v| v.quality_rank);
// Separate single files from shards
let mut single_files: Vec<&HfApiSibling> = Vec::new();
let mut shard_groups: HashMap<String, Vec<&HfApiSibling>> = HashMap::new();
for file in &gguf_files {
if is_shard_file(&file.rfilename) {
let quant = parse_quantization(&file.rfilename);
shard_groups.entry(quant).or_default().push(file);
} else {
single_files.push(file);
}
}
let mut variants: Vec<HfQuantVariant> = Vec::new();
let mut seen_quants: std::collections::HashSet<String> = std::collections::HashSet::new();
// Add single-file variants
for s in single_files {
let quant = parse_quantization(&s.rfilename);
seen_quants.insert(quant.clone());
let info = quant_info(&quant);
let download_url = build_download_url(repo_id, &s.rfilename);
variants.push(HfQuantVariant {
quantization: quant,
size_bytes: s.size.unwrap_or(0),
filename: s.rfilename.clone(),
download_url,
description: info.description,
quality_rank: info.quality_rank,
sharded: false,
});
}
// Add shard-only variants (quants that only exist as sharded files)
for (quant, mut shards) in shard_groups {
if seen_quants.contains(&quant) {
continue;
}
shards.sort_by(|a, b| a.rfilename.cmp(&b.rfilename));
let total_size: u64 = shards.iter().map(|s| s.size.unwrap_or(0)).sum();
let info = quant_info(&quant);
let first_filename = &shards[0].rfilename;
let download_url = build_download_url(repo_id, first_filename);
variants.push(HfQuantVariant {
quantization: quant,
size_bytes: total_size,
filename: first_filename.clone(),
download_url,
description: info.description,
quality_rank: info.quality_rank,
sharded: true,
});
}
// Sort descending by quality_rank, then by size descending as tiebreaker
variants.sort_by(|a, b| {
b.quality_rank
.cmp(&a.quality_rank)
.then_with(|| b.size_bytes.cmp(&a.size_bytes))
});
variants
}
@@ -323,10 +436,13 @@ pub async fn get_repo_gguf_files(repo_id: &str) -> Result<Vec<HfGgufFile>> {
let model: HfApiModel = response.json().await?;
let siblings = model.siblings.unwrap_or_default();
let stem = model_stem_from_repo(repo_id);
let files = siblings
.into_iter()
.filter(|s| s.rfilename.ends_with(".gguf"))
.filter(|s| !is_shard_file(&s.rfilename))
.filter(|s| is_model_file(&s.rfilename, &stem))
.map(|s| {
let quantization = parse_quantization(&s.rfilename);
let download_url = build_download_url(repo_id, &s.rfilename);
@@ -358,23 +474,146 @@ pub fn parse_model_spec(spec: &str) -> Result<(String, String)> {
Ok((repo_id.to_string(), quant.to_string()))
}
/// Resolve a model spec to all GGUF files for that quantization (handles shards).
pub async fn resolve_model_spec_full(spec: &str) -> Result<(String, ResolvedModel)> {
let (repo_id, quant) = parse_model_spec(spec)?;
let client = reqwest::Client::new();
let url = format!("{}/{}?blobs=true", HF_API_BASE, repo_id);
let response = client
.get(&url)
.header("User-Agent", "goose-ai-agent")
.send()
.await?;
if !response.status().is_success() {
bail!(
"HuggingFace API returned status {} for repo {}",
response.status(),
repo_id
);
}
let model: HfApiModel = response.json().await?;
let siblings = model.siblings.unwrap_or_default();
let stem = model_stem_from_repo(&repo_id);
// Collect all GGUF files matching the quantization
let matching: Vec<_> = siblings
.into_iter()
.filter(|s| {
s.rfilename.ends_with(".gguf")
&& is_model_file(&s.rfilename, &stem)
&& parse_quantization(&s.rfilename).eq_ignore_ascii_case(&quant)
})
.collect();
if matching.is_empty() {
bail!(
"No GGUF file with quantization '{}' found in {}",
quant,
repo_id
);
}
// Separate single files from shards
let mut single_files: Vec<&HfApiSibling> = Vec::new();
let mut shard_files: Vec<&HfApiSibling> = Vec::new();
for f in &matching {
if is_shard_file(&f.rfilename) {
shard_files.push(f);
} else {
single_files.push(f);
}
}
// Prefer single file if available
if let Some(single) = single_files.first() {
let file = HfGgufFile {
filename: single.rfilename.clone(),
size_bytes: single.size.unwrap_or(0),
quantization: quant,
download_url: build_download_url(&repo_id, &single.rfilename),
};
let total_size = file.size_bytes;
return Ok((
repo_id,
ResolvedModel {
files: vec![file],
total_size,
},
));
}
// Use shards, sorted by filename so shard 1 is first
shard_files.sort_by(|a, b| a.rfilename.cmp(&b.rfilename));
// Validate shard set completeness: every file must parse to the same
// -of-N total, and indices must be contiguous 1..=N.
let expected_total = parse_shard_total(&shard_files[0].rfilename).ok_or_else(|| {
anyhow::anyhow!(
"Cannot parse shard total from '{}'",
shard_files[0].rfilename
)
})?;
if shard_files.len() != expected_total as usize {
bail!(
"Incomplete shard set for '{}' in {}: found {} of {} shards",
quant,
repo_id,
shard_files.len(),
expected_total
);
}
for (i, shard) in shard_files.iter().enumerate() {
let shard_total = parse_shard_total(&shard.rfilename);
if shard_total != Some(expected_total) {
bail!(
"Inconsistent shard totals for '{}' in {}: shard '{}' has total {:?}, expected {}",
quant,
repo_id,
shard.rfilename,
shard_total,
expected_total
);
}
let idx = parse_shard_index(&shard.rfilename);
if idx != Some((i + 1) as u32) {
bail!(
"Non-contiguous shard set for '{}' in {}: expected shard {} but found {:?}",
quant,
repo_id,
i + 1,
idx
);
}
}
let files: Vec<HfGgufFile> = shard_files
.iter()
.map(|s| HfGgufFile {
filename: s.rfilename.clone(),
size_bytes: s.size.unwrap_or(0),
quantization: quant.clone(),
download_url: build_download_url(&repo_id, &s.rfilename),
})
.collect();
let total_size: u64 = files.iter().map(|f| f.size_bytes).sum();
Ok((repo_id, ResolvedModel { files, total_size }))
}
/// Resolve a model spec to a specific GGUF file from the repo.
pub async fn resolve_model_spec(spec: &str) -> Result<(String, HfGgufFile)> {
let (repo_id, quant) = parse_model_spec(spec)?;
let files = get_repo_gguf_files(&repo_id).await?;
let file = files
.into_iter()
.find(|f| f.quantization.eq_ignore_ascii_case(&quant))
.ok_or_else(|| {
anyhow::anyhow!(
"No GGUF file with quantization '{}' found in {}",
quant,
repo_id
)
})?;
Ok((repo_id, file))
let (repo_id, resolved) = resolve_model_spec_full(spec).await?;
if resolved.files.len() > 1 {
bail!(
"Model '{}' is sharded ({} files) — use resolve_model_spec_full instead",
spec,
resolved.files.len()
);
}
Ok((repo_id, resolved.files.into_iter().next().unwrap()))
}
/// Recommend which quantization variant to use based on available memory.
@@ -459,7 +698,8 @@ mod tests {
filename: "m-Q2_K.gguf".into(),
download_url: String::new(),
description: "Small",
quality_rank: 7,
quality_rank: 24,
sharded: false,
},
HfQuantVariant {
quantization: "Q4_K_M".into(),
@@ -467,7 +707,8 @@ mod tests {
filename: "m-Q4_K_M.gguf".into(),
download_url: String::new(),
description: "Medium",
quality_rank: 19,
quality_rank: 45,
sharded: false,
},
HfQuantVariant {
quantization: "Q8_0".into(),
@@ -475,7 +716,8 @@ mod tests {
filename: "m-Q8_0.gguf".into(),
download_url: String::new(),
description: "Large",
quality_rank: 25,
quality_rank: 80,
sharded: false,
},
];
@@ -483,4 +725,95 @@ mod tests {
assert_eq!(recommend_variant(&variants, 10_000_000_000), Some(2));
assert_eq!(recommend_variant(&variants, 1_000_000_000), None);
}
#[test]
fn test_model_stem_from_repo() {
assert_eq!(
model_stem_from_repo("unsloth/gemma-4-26B-A4B-it-GGUF"),
"gemma-4-26b-a4b-it"
);
assert_eq!(
model_stem_from_repo("bartowski/Llama-3.2-3B-Instruct-GGUF"),
"llama-3.2-3b-instruct"
);
assert_eq!(model_stem_from_repo("someone/SomeModel"), "somemodel");
}
#[test]
fn test_is_model_file() {
let stem = "gemma-3-27b-it";
assert!(is_model_file("gemma-3-27b-it-Q4_K_M.gguf", stem));
assert!(is_model_file(
"BF16/gemma-3-27b-it-BF16-00001-of-00002.gguf",
stem
));
assert!(!is_model_file("mmproj-BF16.gguf", stem));
assert!(!is_model_file("vision-encoder-Q4_K_M.gguf", stem));
}
#[test]
fn test_group_into_variants_filters_auxiliary_files() {
let files = vec![
HfApiSibling {
rfilename: "gemma-3-27b-it-Q4_K_M.gguf".into(),
size: Some(4_000_000_000),
},
HfApiSibling {
rfilename: "mmproj-BF16.gguf".into(),
size: Some(800_000_000),
},
];
let variants = group_into_variants("unsloth/gemma-3-27b-it-GGUF", files);
assert_eq!(variants.len(), 1);
assert_eq!(variants[0].quantization, "Q4_K_M");
}
#[test]
fn test_group_into_variants_includes_shard_only_quants() {
let files = vec![
HfApiSibling {
rfilename: "BF16/gemma-3-27b-it-BF16-00001-of-00002.gguf".into(),
size: Some(40_000_000_000),
},
HfApiSibling {
rfilename: "BF16/gemma-3-27b-it-BF16-00002-of-00002.gguf".into(),
size: Some(10_000_000_000),
},
HfApiSibling {
rfilename: "gemma-3-27b-it-Q4_K_M.gguf".into(),
size: Some(4_000_000_000),
},
];
let variants = group_into_variants("unsloth/gemma-3-27b-it-GGUF", files);
assert_eq!(variants.len(), 2);
// Sorted descending by quality_rank: BF16 (91) > Q4_K_M (45)
assert_eq!(variants[0].quantization, "BF16");
assert!(variants[0].sharded);
assert_eq!(variants[0].size_bytes, 50_000_000_000);
assert_eq!(variants[1].quantization, "Q4_K_M");
assert!(!variants[1].sharded);
}
#[test]
fn test_group_into_variants_sorted_descending() {
let files = vec![
HfApiSibling {
rfilename: "Model-IQ1_S.gguf".into(),
size: Some(500_000_000),
},
HfApiSibling {
rfilename: "Model-Q4_K_M.gguf".into(),
size: Some(4_000_000_000),
},
HfApiSibling {
rfilename: "Model-Q8_0.gguf".into(),
size: Some(8_000_000_000),
},
];
let variants = group_into_variants("someone/Model-GGUF", files);
assert_eq!(variants.len(), 3);
assert_eq!(variants[0].quantization, "Q8_0");
assert_eq!(variants[1].quantization, "Q4_K_M");
assert_eq!(variants[2].quantization, "IQ1_S");
}
}
@@ -233,6 +233,14 @@ pub fn get_registry() -> &'static Mutex<LocalModelRegistry> {
})
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ShardFile {
pub filename: String,
pub local_path: PathBuf,
pub source_url: String,
pub size_bytes: u64,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct LocalModelEntry {
pub id: String,
@@ -254,6 +262,8 @@ pub struct LocalModelEntry {
/// Size of the mmproj file in bytes.
#[serde(default)]
pub mmproj_size_bytes: u64,
#[serde(default)]
pub shard_files: Vec<ShardFile>,
}
impl LocalModelEntry {
@@ -282,7 +292,14 @@ impl LocalModelEntry {
}
pub fn is_downloaded(&self) -> bool {
self.local_path.exists()
self.local_path.exists() && self.shard_files.iter().all(|s| s.local_path.exists())
}
/// Returns all GGUF model file paths (primary + shards).
/// Does NOT include mmproj — that has separate shared-ownership deletion logic.
pub fn all_local_paths(&self) -> impl Iterator<Item = &std::path::Path> {
std::iter::once(self.local_path.as_path())
.chain(self.shard_files.iter().map(|s| s.local_path.as_path()))
}
pub fn is_downloading(&self) -> bool {
@@ -292,7 +309,7 @@ impl LocalModelEntry {
}
pub fn download_status(&self) -> ModelDownloadStatus {
if self.local_path.exists() {
if self.is_downloaded() {
return ModelDownloadStatus::Downloaded;
}
+17 -1
View File
@@ -5522,6 +5522,9 @@
"quantization": {
"type": "string"
},
"sharded": {
"type": "boolean"
},
"size_bytes": {
"type": "integer",
"format": "int64",
@@ -7261,9 +7264,22 @@
"RepoVariantsResponse": {
"type": "object",
"required": [
"variants"
"variants",
"available_memory_bytes",
"downloaded_quants"
],
"properties": {
"available_memory_bytes": {
"type": "integer",
"format": "int64",
"minimum": 0
},
"downloaded_quants": {
"type": "array",
"items": {
"type": "string"
}
},
"recommended_index": {
"type": "integer",
"nullable": true,
+3
View File
@@ -541,6 +541,7 @@ export type HfQuantVariant = {
filename: string;
quality_rank: number;
quantization: string;
sharded?: boolean;
size_bytes: number;
};
@@ -1095,6 +1096,8 @@ export type RemoveExtensionRequest = {
};
export type RepoVariantsResponse = {
available_memory_bytes: number;
downloaded_quants: Array<string>;
recommended_index?: number | null;
variants: Array<HfQuantVariant>;
};
@@ -1,5 +1,5 @@
import { useState, useCallback, useRef } from 'react';
import { Search, Download, ChevronDown, ChevronUp, Loader2, Star } from 'lucide-react';
import { Search, Download, ChevronDown, ChevronUp, Loader2, Star, Check, AlertTriangle } from 'lucide-react';
import { Button } from '../../ui/button';
import {
searchHfModels,
@@ -8,8 +8,6 @@ import {
type HfModelInfo,
type HfQuantVariant,
} from '../../../api';
import { toastError } from '../../../toasts';
import { errorMessage } from '../../../utils/conversionUtils';
import { defineMessages, useIntl } from '../../../i18n';
const i18n = defineMessages({
@@ -33,21 +31,17 @@ const i18n = defineMessages({
id: 'huggingFaceModelSearch.download',
defaultMessage: 'Download',
},
directDownload: {
id: 'huggingFaceModelSearch.directDownload',
defaultMessage: 'Direct Download',
downloaded: {
id: 'huggingFaceModelSearch.downloaded',
defaultMessage: 'Downloaded',
},
directDownloadDescription: {
id: 'huggingFaceModelSearch.directDownloadDescription',
defaultMessage: 'Specify a model directly: {format}',
downloading: {
id: 'huggingFaceModelSearch.downloading',
defaultMessage: 'Downloading…',
},
directDownloadFailed: {
id: 'huggingFaceModelSearch.directDownloadFailed',
defaultMessage: 'Direct download failed',
},
directDownloadErrorMsg: {
id: 'huggingFaceModelSearch.directDownloadErrorMsg',
defaultMessage: 'Failed to start the download. Check the spec: {error}',
tooLarge: {
id: 'huggingFaceModelSearch.tooLarge',
defaultMessage: 'May not fit in memory ({size} model, {available} available)',
},
noGgufModels: {
id: 'huggingFaceModelSearch.noGgufModels',
@@ -84,13 +78,19 @@ const formatDownloads = (n: number): string => {
interface RepoData {
variants: HfQuantVariant[];
recommendedIndex: number | null;
availableMemoryBytes: number;
downloadedQuants: Set<string>;
}
interface Props {
onDownloadStarted: (modelId: string) => void;
/** Model IDs (repo:quant) with an active download in progress */
activeDownloadIds?: Set<string>;
/** Model IDs (repo:quant) confirmed downloaded on disk */
downloadedModelIds?: Set<string>;
}
export const HuggingFaceModelSearch = ({ onDownloadStarted }: Props) => {
export const HuggingFaceModelSearch = ({ onDownloadStarted, activeDownloadIds, downloadedModelIds }: Props) => {
const intl = useIntl();
const [query, setQuery] = useState('');
const [results, setResults] = useState<HfModelInfo[]>([]);
@@ -99,7 +99,6 @@ export const HuggingFaceModelSearch = ({ onDownloadStarted }: Props) => {
const [searching, setSearching] = useState(false);
const [downloading, setDownloading] = useState<Set<string>>(new Set());
const [loadingFiles, setLoadingFiles] = useState<Set<string>>(new Set());
const [directSpec, setDirectSpec] = useState('');
const [error, setError] = useState<string | null>(null);
const debounceRef = useRef<ReturnType<typeof setTimeout> | null>(null);
@@ -134,7 +133,7 @@ export const HuggingFaceModelSearch = ({ onDownloadStarted }: Props) => {
const validResults = modelsWithVariants.filter(Boolean) as {
model: HfModelInfo;
data: { variants: HfQuantVariant[]; recommended_index?: number | null };
data: { variants: HfQuantVariant[]; recommended_index?: number | null; available_memory_bytes: number; downloaded_quants: string[] };
}[];
setResults(validResults.map((r) => r.model));
@@ -144,6 +143,8 @@ export const HuggingFaceModelSearch = ({ onDownloadStarted }: Props) => {
next[r.model.repo_id] = {
variants: r.data.variants,
recommendedIndex: r.data.recommended_index ?? null,
availableMemoryBytes: r.data.available_memory_bytes,
downloadedQuants: new Set(r.data.downloaded_quants),
};
}
return next;
@@ -188,12 +189,13 @@ export const HuggingFaceModelSearch = ({ onDownloadStarted }: Props) => {
path: { author, repo },
});
if (response.data) {
const variants = response.data.variants;
setRepoData((prev) => ({
...prev,
[repoId]: {
variants,
variants: response.data!.variants,
recommendedIndex: response.data!.recommended_index ?? null,
availableMemoryBytes: response.data!.available_memory_bytes,
downloadedQuants: new Set(response.data!.downloaded_quants),
},
}));
}
@@ -230,34 +232,6 @@ export const HuggingFaceModelSearch = ({ onDownloadStarted }: Props) => {
}
};
const startDirectDownload = async () => {
const spec = directSpec.trim();
if (!spec) return;
const key = `direct:${spec}`;
setDownloading((prev) => new Set(prev).add(key));
try {
const response = await downloadHfModel({
body: { spec },
throwOnError: true,
});
if (response.data) {
onDownloadStarted(response.data);
setDirectSpec('');
}
} catch (e) {
toastError({
title: intl.formatMessage(i18n.directDownloadFailed),
msg: intl.formatMessage(i18n.directDownloadErrorMsg, { error: errorMessage(e) }),
});
} finally {
setDownloading((prev) => {
const next = new Set(prev);
next.delete(key);
return next;
});
}
};
return (
<div className="space-y-4">
<div>
@@ -280,12 +254,14 @@ export const HuggingFaceModelSearch = ({ onDownloadStarted }: Props) => {
{error && !searching && <p className="text-xs text-text-muted">{error}</p>}
{results.length > 0 && (
<div className="space-y-1 max-h-96 overflow-y-auto">
<div className="space-y-1">
{results.map((model) => {
const isExpanded = expandedRepo === model.repo_id;
const data = repoData[model.repo_id];
const variants = data?.variants || [];
const recommendedIndex = data?.recommendedIndex ?? null;
const availableMemory = data?.availableMemoryBytes ?? 0;
const downloadedQuants = data?.downloadedQuants ?? new Set<string>();
return (
<div key={model.repo_id} className="border border-border-subtle rounded-lg">
@@ -324,14 +300,22 @@ export const HuggingFaceModelSearch = ({ onDownloadStarted }: Props) => {
const dlKey = `${model.repo_id}:${variant.quantization}`;
const isStarting = downloading.has(dlKey);
const isRecommended = idx === recommendedIndex;
const modelId = `${model.repo_id}:${variant.quantization}`;
const isActiveDownload = activeDownloadIds?.has(modelId) ?? false;
const isDownloaded = downloadedModelIds
? downloadedModelIds.has(modelId)
: downloadedQuants.has(variant.quantization);
const tooLarge = availableMemory > 0 && variant.size_bytes > availableMemory * 0.85;
return (
<div
key={variant.quantization}
className={`flex items-center justify-between py-2 px-2 rounded ${
isRecommended
? 'bg-blue-500/5 border border-blue-500/20'
: 'hover:bg-background-subtle'
isDownloaded
? 'bg-green-500/5 border border-green-500/20'
: isRecommended
? 'bg-blue-500/5 border border-blue-500/20'
: 'hover:bg-background-subtle'
}`}
>
<div className="flex flex-col gap-0.5 min-w-0 flex-1 mr-3">
@@ -342,7 +326,7 @@ export const HuggingFaceModelSearch = ({ onDownloadStarted }: Props) => {
<span className="text-xs text-text-muted">
{formatBytes(variant.size_bytes)}
</span>
{isRecommended && (
{isRecommended && !isDownloaded && (
<span className="inline-flex items-center gap-1 text-xs bg-blue-500 text-white px-1.5 py-0.5 rounded">
<Star className="w-3 h-3" />
{intl.formatMessage(i18n.recommended)}
@@ -352,22 +336,53 @@ export const HuggingFaceModelSearch = ({ onDownloadStarted }: Props) => {
{variant.description && (
<span className="text-xs text-text-muted">{variant.description}</span>
)}
</div>
<Button
variant="outline"
size="sm"
disabled={isStarting}
onClick={() => startDownload(model.repo_id, variant.quantization)}
>
{isStarting ? (
<Loader2 className="w-3 h-3 animate-spin" />
) : (
<>
<Download className="w-3 h-3 mr-1" />
{intl.formatMessage(i18n.download)}
</>
{tooLarge && !isDownloaded && (
<span className="inline-flex items-center gap-1 text-xs text-amber-500">
<AlertTriangle className="w-3 h-3" />
{intl.formatMessage(i18n.tooLarge, {
size: formatBytes(variant.size_bytes),
available: formatBytes(availableMemory),
})}
</span>
)}
</Button>
</div>
{isDownloaded ? (
<Button
variant="outline"
size="sm"
disabled
className="opacity-60"
>
<Check className="w-3 h-3 mr-1" />
{intl.formatMessage(i18n.downloaded)}
</Button>
) : isActiveDownload ? (
<Button
variant="outline"
size="sm"
disabled
className="opacity-60"
>
<Loader2 className="w-3 h-3 mr-1 animate-spin" />
{intl.formatMessage(i18n.downloading)}
</Button>
) : (
<Button
variant="outline"
size="sm"
disabled={isStarting}
onClick={() => startDownload(model.repo_id, variant.quantization)}
>
{isStarting ? (
<Loader2 className="w-3 h-3 animate-spin" />
) : (
<>
<Download className="w-3 h-3 mr-1" />
{intl.formatMessage(i18n.download)}
</>
)}
</Button>
)}
</div>
);
})}
@@ -379,41 +394,6 @@ export const HuggingFaceModelSearch = ({ onDownloadStarted }: Props) => {
</div>
)}
<div>
<h4 className="text-sm font-medium text-text-default mb-2">{intl.formatMessage(i18n.directDownload)}</h4>
<p className="text-xs text-text-muted mb-2">
{intl.formatMessage(i18n.directDownloadDescription, {
format: 'user/repo:quantization',
})}
</p>
<div className="flex gap-2">
<input
type="text"
value={directSpec}
onChange={(e) => setDirectSpec(e.target.value)}
placeholder="bartowski/Llama-3.2-1B-Instruct-GGUF:Q4_K_M"
className="flex-1 px-3 py-2 text-sm border border-border-subtle rounded-lg bg-background-default text-text-default placeholder:text-text-muted focus:outline-none focus:border-accent-primary"
onKeyDown={(e) => {
if (e.key === 'Enter') startDirectDownload();
}}
/>
<Button
variant="outline"
size="sm"
disabled={!directSpec.trim() || downloading.has(`direct:${directSpec}`)}
onClick={startDirectDownload}
>
{downloading.has(`direct:${directSpec}`) ? (
<Loader2 className="w-4 h-4 animate-spin" />
) : (
<>
<Download className="w-4 h-4 mr-1" />
{intl.formatMessage(i18n.download)}
</>
)}
</Button>
</div>
</div>
</div>
);
};
@@ -520,7 +520,11 @@ export const LocalInferenceSettings = () => {
{/* HuggingFace Search */}
<div className="border-t border-border-subtle pt-4">
<HuggingFaceModelSearch onDownloadStarted={handleHfDownloadStarted} />
<HuggingFaceModelSearch
onDownloadStarted={handleHfDownloadStarted}
activeDownloadIds={new Set(downloads.keys())}
downloadedModelIds={new Set(models.filter(m => m.status.state === 'Downloaded').map(m => m.id))}
/>
</div>
{models.length === 0 && (
+9 -12
View File
@@ -1436,21 +1436,15 @@
"headersSection.value": {
"defaultMessage": "Value"
},
"huggingFaceModelSearch.directDownload": {
"defaultMessage": "Direct Download"
},
"huggingFaceModelSearch.directDownloadDescription": {
"defaultMessage": "Specify a model directly: {format}"
},
"huggingFaceModelSearch.directDownloadErrorMsg": {
"defaultMessage": "Failed to start the download. Check the spec: {error}"
},
"huggingFaceModelSearch.directDownloadFailed": {
"defaultMessage": "Direct download failed"
},
"huggingFaceModelSearch.download": {
"defaultMessage": "Download"
},
"huggingFaceModelSearch.downloaded": {
"defaultMessage": "Downloaded"
},
"huggingFaceModelSearch.downloading": {
"defaultMessage": "Downloading\u2026"
},
"huggingFaceModelSearch.loadingVariants": {
"defaultMessage": "Loading variants..."
},
@@ -1475,6 +1469,9 @@
"huggingFaceModelSearch.searchPlaceholder": {
"defaultMessage": "Search for GGUF models..."
},
"huggingFaceModelSearch.tooLarge": {
"defaultMessage": "May not fit in memory ({size} model, {available} available)"
},
"imagePreview.altText": {
"defaultMessage": "goose image"
},