fix: eliminate ~5s delay when opening Switch Models panel with local provider (#8626)

Signed-off-by: jh-block <jhugo@block.xyz>
This commit is contained in:
jh-block
2026-04-17 21:08:41 +02:00
committed by GitHub
parent 3f3b519b62
commit 1a18c2748c
9 changed files with 166 additions and 62 deletions
+1
View File
@@ -674,6 +674,7 @@ pub struct ApiDoc;
super::routes::dictation::cancel_download, super::routes::dictation::cancel_download,
super::routes::dictation::delete_model, super::routes::dictation::delete_model,
super::routes::local_inference::list_local_models, super::routes::local_inference::list_local_models,
super::routes::local_inference::sync_featured_models,
super::routes::local_inference::search_hf_models, super::routes::local_inference::search_hf_models,
super::routes::local_inference::get_repo_files, super::routes::local_inference::get_repo_files,
super::routes::local_inference::download_hf_model, super::routes::local_inference::download_hf_model,
@@ -8,6 +8,7 @@ use axum::{
routing::{delete, get, post}, routing::{delete, get, post},
Json, Router, Json, Router,
}; };
use futures::future::join_all;
use goose::config::paths::Paths; use goose::config::paths::Paths;
use goose::download_manager::{get_download_manager, DownloadProgress}; use goose::download_manager::{get_download_manager, DownloadProgress};
use goose::providers::local_inference::hf_models::{self, HfModelInfo, HfQuantVariant}; use goose::providers::local_inference::hf_models::{self, HfModelInfo, HfQuantVariant};
@@ -55,9 +56,16 @@ pub struct LocalModelResponse {
} }
async fn ensure_featured_models_in_registry() -> Result<(), ErrorResponse> { async fn ensure_featured_models_in_registry() -> Result<(), ErrorResponse> {
let mut entries_to_add = Vec::new();
let mut mmproj_downloads_needed: Vec<(String, String, PathBuf)> = Vec::new(); let mut mmproj_downloads_needed: Vec<(String, String, PathBuf)> = Vec::new();
struct PendingResolve {
spec: &'static str,
repo_id: String,
quantization: String,
model_id: String,
}
let mut to_resolve = Vec::new();
for featured in FEATURED_MODELS { for featured in FEATURED_MODELS {
let (repo_id, quantization) = match hf_models::parse_model_spec(featured.spec) { let (repo_id, quantization) = match hf_models::parse_model_spec(featured.spec) {
Ok(parts) => parts, Ok(parts) => parts,
@@ -90,50 +98,65 @@ async fn ensure_featured_models_in_registry() -> Result<(), ErrorResponse> {
if !needs_backfill { if !needs_backfill {
continue; continue;
} }
// Fall through to build the entry for sync_with_featured backfill // Fall through to resolve for backfill
} }
} }
let hf_file = match resolve_model_spec(featured.spec).await { to_resolve.push(PendingResolve {
Ok((_repo, file)) => file, spec: featured.spec,
Err(_) => {
let filename = format!(
"{}-{}.gguf",
repo_id.split('/').next_back().unwrap_or("model"),
quantization
);
HfGgufFile {
filename: filename.clone(),
size_bytes: 0,
quantization: quantization.to_string(),
download_url: format!(
"https://huggingface.co/{}/resolve/main/{}",
repo_id, filename
),
}
}
};
let local_path = Paths::in_data_dir("models").join(&hf_file.filename);
// enrich_with_featured_mmproj is called by sync_with_featured/add_model,
// so we don't need to populate mmproj fields here.
entries_to_add.push(LocalModelEntry {
id: model_id.clone(),
repo_id, repo_id,
filename: hf_file.filename,
quantization, quantization,
local_path, model_id,
source_url: hf_file.download_url,
settings: default_settings_for_model(&model_id),
size_bytes: hf_file.size_bytes,
mmproj_path: None,
mmproj_source_url: None,
mmproj_size_bytes: 0,
shard_files: vec![],
}); });
} }
let resolved: Vec<(PendingResolve, HfGgufFile)> =
join_all(to_resolve.into_iter().map(|pending| async move {
let hf_file = match resolve_model_spec(pending.spec).await {
Ok((_repo, file)) => file,
Err(_) => {
let filename = format!(
"{}-{}.gguf",
pending.repo_id.split('/').next_back().unwrap_or("model"),
pending.quantization
);
HfGgufFile {
filename: filename.clone(),
size_bytes: 0,
quantization: pending.quantization.to_string(),
download_url: format!(
"https://huggingface.co/{}/resolve/main/{}",
pending.repo_id, filename
),
}
}
};
(pending, hf_file)
}))
.await;
let entries_to_add: Vec<LocalModelEntry> = resolved
.into_iter()
.map(|(pending, hf_file)| {
let local_path = Paths::in_data_dir("models").join(&hf_file.filename);
let settings = default_settings_for_model(&pending.model_id);
LocalModelEntry {
id: pending.model_id,
repo_id: pending.repo_id,
filename: hf_file.filename,
quantization: pending.quantization,
local_path,
source_url: hf_file.download_url,
settings,
size_bytes: hf_file.size_bytes,
mmproj_path: None,
mmproj_source_url: None,
mmproj_size_bytes: 0,
shard_files: vec![],
}
})
.collect();
{ {
let mut registry = get_registry() let mut registry = get_registry()
.lock() .lock()
@@ -185,6 +208,18 @@ async fn ensure_featured_models_in_registry() -> Result<(), ErrorResponse> {
Ok(()) Ok(())
} }
#[utoipa::path(
post,
path = "/local-inference/sync-featured",
responses(
(status = 200, description = "Featured models synced to registry")
)
)]
pub async fn sync_featured_models() -> Result<StatusCode, ErrorResponse> {
ensure_featured_models_in_registry().await?;
Ok(StatusCode::OK)
}
#[utoipa::path( #[utoipa::path(
get, get,
path = "/local-inference/models", path = "/local-inference/models",
@@ -195,8 +230,6 @@ async fn ensure_featured_models_in_registry() -> Result<(), ErrorResponse> {
pub async fn list_local_models( pub async fn list_local_models(
axum::extract::State(state): axum::extract::State<Arc<AppState>>, axum::extract::State(state): axum::extract::State<Arc<AppState>>,
) -> Result<Json<Vec<LocalModelResponse>>, ErrorResponse> { ) -> Result<Json<Vec<LocalModelResponse>>, ErrorResponse> {
ensure_featured_models_in_registry().await?;
let recommended_id = recommend_local_model(&state.inference_runtime); let recommended_id = recommend_local_model(&state.inference_runtime);
let registry = get_registry() let registry = get_registry()
@@ -635,6 +668,7 @@ pub fn routes(state: Arc<AppState>) -> Router {
Router::new() Router::new()
.route("/local-inference/models", get(list_local_models)) .route("/local-inference/models", get(list_local_models))
.route("/local-inference/sync-featured", post(sync_featured_models))
.route("/local-inference/search", get(search_hf_models)) .route("/local-inference/search", get(search_hf_models))
.route( .route(
"/local-inference/repo/{author}/{repo}/files", "/local-inference/repo/{author}/{repo}/files",
+13
View File
@@ -2283,6 +2283,19 @@
} }
} }
}, },
"/local-inference/sync-featured": {
"post": {
"tags": [
"super::routes::local_inference"
],
"operationId": "sync_featured_models",
"responses": {
"200": {
"description": "Featured models synced to registry"
}
}
}
},
"/mcp-ui-proxy": { "/mcp-ui-proxy": {
"get": { "get": {
"tags": [ "tags": [
File diff suppressed because one or more lines are too long
File diff suppressed because one or more lines are too long
+14
View File
@@ -3412,6 +3412,20 @@ export type SearchHfModelsResponses = {
export type SearchHfModelsResponse = SearchHfModelsResponses[keyof SearchHfModelsResponses]; export type SearchHfModelsResponse = SearchHfModelsResponses[keyof SearchHfModelsResponses];
export type SyncFeaturedModelsData = {
body?: never;
path?: never;
query?: never;
url: '/local-inference/sync-featured';
};
export type SyncFeaturedModelsResponses = {
/**
* Featured models synced to registry
*/
200: unknown;
};
export type McpUiProxyData = { export type McpUiProxyData = {
body?: never; body?: never;
path?: never; path?: never;
@@ -1,6 +1,7 @@
import { useState, useEffect, useCallback, useRef } from 'react'; import { useState, useEffect, useCallback, useRef } from 'react';
import { import {
listLocalModels, listLocalModels,
syncFeaturedModels,
downloadHfModel, downloadHfModel,
getLocalModelDownloadProgress, getLocalModelDownloadProgress,
cancelLocalModelDownload, cancelLocalModelDownload,
@@ -128,6 +129,7 @@ export default function LocalModelPicker({ onConfigured, onBack }: LocalModelPic
useEffect(() => { useEffect(() => {
const load = async () => { const load = async () => {
try { try {
await syncFeaturedModels();
const response = await listLocalModels({ throwOnError: true }); const response = await listLocalModels({ throwOnError: true });
if (response.data) { if (response.data) {
setModels(response.data); setModels(response.data);
@@ -5,6 +5,7 @@ import { useModelAndProvider } from '../../ModelAndProviderContext';
import { defineMessages, useIntl } from '../../../i18n'; import { defineMessages, useIntl } from '../../../i18n';
import { import {
listLocalModels, listLocalModels,
syncFeaturedModels,
downloadHfModel, downloadHfModel,
getLocalModelDownloadProgress, getLocalModelDownloadProgress,
cancelLocalModelDownload, cancelLocalModelDownload,
@@ -159,6 +160,7 @@ export const LocalInferenceSettings = () => {
const loadModels = useCallback(async (): Promise<LocalModelResponse[] | undefined> => { const loadModels = useCallback(async (): Promise<LocalModelResponse[] | undefined> => {
try { try {
await syncFeaturedModels();
const response = await listLocalModels(); const response = await listLocalModels();
if (response.data) { if (response.data) {
setModels(response.data); setModels(response.data);
@@ -300,6 +300,10 @@ export const SwitchModelModal = ({
const [userClearedModel, setUserClearedModel] = useState(false); const [userClearedModel, setUserClearedModel] = useState(false);
const [providerErrors, setProviderErrors] = useState<Record<string, string>>({}); const [providerErrors, setProviderErrors] = useState<Record<string, string>>({});
const [providerWarnings, setProviderWarnings] = useState<Record<string, string>>({}); const [providerWarnings, setProviderWarnings] = useState<Record<string, string>>({});
const [activeProvidersList, setActiveProvidersList] = useState<
import('../../../../api').ProviderDetails[]
>([]);
const fetchedProviders = useRef<Set<string>>(new Set());
const [thinkingLevel, setThinkingLevel] = useState<string>('low'); const [thinkingLevel, setThinkingLevel] = useState<string>('low');
const [claudeThinkingType, setClaudeThinkingType] = useState<string>('disabled'); const [claudeThinkingType, setClaudeThinkingType] = useState<string>('disabled');
const [claudeThinkingEffort, setClaudeThinkingEffort] = useState<string>('high'); const [claudeThinkingEffort, setClaudeThinkingEffort] = useState<string>('high');
@@ -465,18 +469,16 @@ export const SwitchModelModal = ({
}, [currentModel, currentProvider, usePredefinedModels, provider, model, initialProvider]); }, [currentModel, currentProvider, usePredefinedModels, provider, model, initialProvider]);
useEffect(() => { useEffect(() => {
// Load predefined models if enabled
if (usePredefinedModels) { if (usePredefinedModels) {
const models = getPredefinedModelsFromEnv(); const models = getPredefinedModelsFromEnv();
setPredefinedModels(models); setPredefinedModels(models);
} }
// Load providers for manual model selection
(async () => { (async () => {
try { try {
const providersResponse = await getProviders(false); const providersResponse = await getProviders(false);
const activeProviders = providersResponse.filter((provider) => provider.is_configured); const activeProviders = providersResponse.filter((provider) => provider.is_configured);
// Create provider options and add "Use other provider" option setActiveProvidersList(activeProviders);
setProviderOptions([ setProviderOptions([
...activeProviders.map(({ metadata, name }) => ({ ...activeProviders.map(({ metadata, name }) => ({
value: name, value: name,
@@ -487,24 +489,43 @@ export const SwitchModelModal = ({
label: intl.formatMessage(i18n.useOtherProvider), label: intl.formatMessage(i18n.useOtherProvider),
}, },
]); ]);
} catch (error: unknown) {
console.error('Failed to query providers:', error);
}
})();
}, [getProviders, usePredefinedModels, read, intl]);
setLoadingModels(true); useEffect(() => {
if (!provider || usePredefinedModels) return;
if (fetchedProviders.current.has(provider)) {
setLoadingModels(false);
return;
}
const results = await fetchModelsForProviders(activeProviders); const activeProvider = activeProvidersList.find((p) => p.name === provider);
if (!activeProvider) return;
// Process results and build grouped options let cancelled = false;
const groupedOptions: {
(async () => {
setLoadingModels(true);
try {
const results = await fetchModelsForProviders([activeProvider]);
if (cancelled) return;
const newGroupedOptions: {
options: { value: string; label: string; provider: string; providerType: ProviderType }[]; options: { value: string; label: string; provider: string; providerType: ProviderType }[];
}[] = []; }[] = [];
const errorMap: Record<string, string> = {}; const newErrors: Record<string, string> = {};
const warningMap: Record<string, string> = {}; const newWarnings: Record<string, string> = {};
results.forEach(({ provider: p, models, error, warning }) => { results.forEach(({ provider: p, models, error, warning }) => {
if (warning) { if (warning) {
warningMap[p.name] = warning; newWarnings[p.name] = warning;
} }
if (error) { if (error) {
errorMap[p.name] = error; newErrors[p.name] = error;
return; return;
} }
@@ -532,23 +553,38 @@ export const SwitchModelModal = ({
} }
if (options.length > 0) { if (options.length > 0) {
groupedOptions.push({ options }); newGroupedOptions.push({ options });
} }
}); });
// Save provider errors and warnings to state setProviderErrors((prev) => {
setProviderErrors(errorMap); const next = { ...prev, ...newErrors };
setProviderWarnings(warningMap); if (!newErrors[activeProvider.name]) delete next[activeProvider.name];
return next;
});
setProviderWarnings((prev) => {
const next = { ...prev, ...newWarnings };
if (!newWarnings[activeProvider.name]) delete next[activeProvider.name];
return next;
});
setModelOptions(groupedOptions); setModelOptions((prev) => [...prev, ...newGroupedOptions]);
setOriginalModelOptions(groupedOptions); setOriginalModelOptions((prev) => [...prev, ...newGroupedOptions]);
fetchedProviders.current.add(provider);
} catch (error: unknown) { } catch (error: unknown) {
console.error('Failed to query providers:', error); console.error(`Failed to fetch models for ${provider}:`, error);
} finally { } finally {
setLoadingModels(false); if (!cancelled) {
setLoadingModels(false);
}
} }
})(); })();
}, [getProviders, usePredefinedModels, read, intl]);
return () => {
cancelled = true;
setLoadingModels(false);
};
}, [provider, activeProvidersList, usePredefinedModels, intl]);
const filteredModelOptions = provider const filteredModelOptions = provider
? modelOptions.filter((group) => group.options[0]?.provider === provider) ? modelOptions.filter((group) => group.options[0]?.provider === provider)