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:
@@ -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",
|
||||||
|
|||||||
@@ -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
@@ -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)
|
||||||
|
|||||||
Reference in New Issue
Block a user