diff --git a/open-sse/providers/registry/huggingface.js b/open-sse/providers/registry/huggingface.js index 768b0ded95a..3260bd47133 100644 --- a/open-sse/providers/registry/huggingface.js +++ b/open-sse/providers/registry/huggingface.js @@ -25,9 +25,7 @@ export default { transport: null, models: [ { id: "black-forest-labs/FLUX.1-schnell", name: "FLUX.1 Schnell", params: [], kind: "image" }, - { id: "stabilityai/stable-diffusion-xl-base-1.0", name: "SDXL Base 1.0", params: [], kind: "image" }, { id: "openai/whisper-large-v3", name: "Whisper Large v3 (HF)", params: ["language"], kind: "stt" }, - { id: "openai/whisper-small", name: "Whisper Small (HF)", params: ["language"], kind: "stt" }, ], serviceKinds: ["image", "stt"], imageConfig: { baseUrl: "https://api-inference.huggingface.co/models" }, diff --git a/src/app/(dashboard)/dashboard/providers/components/ModelsCard.js b/src/app/(dashboard)/dashboard/providers/components/ModelsCard.js index 9bdf4b411bb..82bfb70124b 100644 --- a/src/app/(dashboard)/dashboard/providers/components/ModelsCard.js +++ b/src/app/(dashboard)/dashboard/providers/components/ModelsCard.js @@ -4,8 +4,9 @@ import { useState, useCallback, useEffect } from "react"; import PropTypes from "prop-types"; import { Card, Button, Modal } from "@/shared/components"; import { getModelsByProviderId, getModelKind } from "@/shared/constants/models"; -import { getProviderAlias } from "@/shared/constants/providers"; +import { getProviderAlias, getProviderModelsFetcher } from "@/shared/constants/providers"; import { useCopyToClipboard } from "@/shared/hooks/useCopyToClipboard"; +import { fetchSuggestedModels } from "@/shared/utils/providerModelsFetcher"; // ── ModelRow ─────────────────────────────────────────────────── export function ModelRow({ model, fullModel, copied, onCopy, testStatus, isCustom, isFree, onDeleteAlias, onTest, isTesting }) { @@ -116,6 +117,7 @@ export default function ModelsCard({ providerId, kindFilter, providerAliasOverri const [testingModelId, setTestingModelId] = useState(null); const [testError, setTestError] = useState(""); const [showAddCustomModel, setShowAddCustomModel] = useState(false); + const [suggestedModels, setSuggestedModels] = useState([]); const providerAlias = providerAliasOverride || getProviderAlias(providerId); const effectiveType = kindFilter || "llm"; @@ -135,6 +137,21 @@ export default function ModelsCard({ providerId, kindFilter, providerAliasOverri useEffect(() => { fetchData(); }, [fetchData]); + useEffect(() => { + const fetcher = getProviderModelsFetcher(providerId, effectiveType); + if (!fetcher) { + setSuggestedModels([]); + return; + } + let cancelled = false; + fetchSuggestedModels(fetcher).then((models) => { + if (!cancelled) setSuggestedModels(models); + }); + return () => { + cancelled = true; + }; + }, [providerId, effectiveType]); + const handleSetAlias = async (modelId, alias) => { const fullModel = `${providerAlias}/${modelId}`; try { @@ -214,6 +231,11 @@ export default function ModelsCard({ providerId, kindFilter, providerAliasOverri ); const displayModels = builtInModels; + const customIds = new Set(myCustomModels.map((m) => m.id)); + const builtInIds = new Set(builtInModels.map((m) => m.id)); + const suggestedNotAdded = suggestedModels.filter( + (m) => !builtInIds.has(m.id) && !customIds.has(m.id) + ); return ( <> @@ -268,6 +290,27 @@ export default function ModelsCard({ providerId, kindFilter, providerAliasOverri add Add Model + + {suggestedNotAdded.length > 0 && ( +
+

Suggested models from provider:

+
+ {suggestedNotAdded.map((model) => ( + + ))} +
+
+ )} diff --git a/src/app/api/providers/suggested-models/filters.js b/src/app/api/providers/suggested-models/filters.js index 8299f46b01b..6e74eeb9a8b 100644 --- a/src/app/api/providers/suggested-models/filters.js +++ b/src/app/api/providers/suggested-models/filters.js @@ -23,4 +23,14 @@ export const FILTERS = { (Array.isArray(models) ? models : []) .filter((m) => m.id?.startsWith("mimo") || m.name?.toLowerCase().includes("mimo")) .map((m) => ({ id: m.id, name: m.name || m.id })), + + "huggingface-hub": (models) => + [...(Array.isArray(models) ? models : [])] + .filter((m) => typeof m.id === "string" && m.id.trim()) + .sort( + (a, b) => + (Number(b.downloads) || 0) - (Number(a.downloads) || 0) || + (Number(b.likes) || 0) - (Number(a.likes) || 0) + ) + .map((m) => ({ id: m.id, name: m.name || m.id })), }; diff --git a/src/shared/constants/providers.js b/src/shared/constants/providers.js index dd116d16129..591c754a19c 100644 --- a/src/shared/constants/providers.js +++ b/src/shared/constants/providers.js @@ -6,7 +6,7 @@ const MEDIA_ENTRY_KEYS = [ "serviceKinds", "ttsConfig", "sttConfig", "embeddingConfig", "imageConfig", "imageToTextConfig", "videoConfig", "musicConfig", "searchViaChat", "searchConfig", "fetchConfig", - "modelsFetcher", "mediaPriority", "hiddenKinds", + "modelsFetcher", "modelsFetchers", "mediaPriority", "hiddenKinds", ]; // Build provider UI object from registry entry @@ -65,6 +65,20 @@ export const THINKING_CONFIG = { export const OAUTH_PROVIDERS = byCategory("oauth"); export const APIKEY_PROVIDERS = byCategory("apikey"); +APIKEY_PROVIDERS.huggingface = { + ...APIKEY_PROVIDERS.huggingface, + modelsFetchers: { + image: { + url: "https://huggingface.co/api/models?inference_provider=hf-inference&pipeline_tag=text-to-image&limit=30", + type: "huggingface-hub", + }, + stt: { + url: "https://huggingface.co/api/models?inference_provider=hf-inference&pipeline_tag=automatic-speech-recognition&limit=30", + type: "huggingface-hub", + }, + }, +}; + // Web Cookie Providers (use browser session cookie instead of API key) export const WEB_COOKIE_PROVIDERS = byCategory("webCookie"); @@ -129,6 +143,13 @@ export function getProviderAlias(providerId) { return provider?.alias || providerId; } +export function getProviderModelsFetcher(providerId, kind) { + const provider = AI_PROVIDERS[providerId]; + if (!provider) return null; + if (kind && provider.modelsFetchers?.[kind]) return provider.modelsFetchers[kind]; + return provider.modelsFetcher || null; +} + // Alias to ID mapping (for quick lookup) export const ALIAS_TO_ID = Object.values(AI_PROVIDERS).reduce((acc, p) => { acc[p.alias] = p.id; diff --git a/src/shared/utils/providerModelsFetcher.js b/src/shared/utils/providerModelsFetcher.js index 3ba1bab6ce6..6e04b484bcd 100644 --- a/src/shared/utils/providerModelsFetcher.js +++ b/src/shared/utils/providerModelsFetcher.js @@ -2,7 +2,7 @@ // Fetches via backend proxy to avoid CORS issues const CACHE_TTL_MS = 10 * 60 * 1000; // 10 minutes -const cache = new Map(); // key: fetcher.url → { data, expiresAt } +const cache = new Map(); // key: `${fetcher.type}:${fetcher.url}` → { data, expiresAt } /** * Fetch suggested models for a provider using its modelsFetcher config. @@ -13,7 +13,8 @@ const cache = new Map(); // key: fetcher.url → { data, expiresAt } export async function fetchSuggestedModels(fetcher) { if (!fetcher?.url || !fetcher?.type) return []; - const cached = cache.get(fetcher.url); + const cacheKey = `${fetcher.type}:${fetcher.url}`; + const cached = cache.get(cacheKey); if (cached && Date.now() < cached.expiresAt) return cached.data; try { @@ -22,7 +23,7 @@ export async function fetchSuggestedModels(fetcher) { if (!res.ok) return []; const json = await res.json(); const data = json.data ?? []; - cache.set(fetcher.url, { data, expiresAt: Date.now() + CACHE_TTL_MS }); + cache.set(cacheKey, { data, expiresAt: Date.now() + CACHE_TTL_MS }); return data; } catch { return []; diff --git a/tests/unit/huggingface-suggested-models.test.js b/tests/unit/huggingface-suggested-models.test.js new file mode 100644 index 00000000000..16d87708681 --- /dev/null +++ b/tests/unit/huggingface-suggested-models.test.js @@ -0,0 +1,38 @@ +import { describe, expect, it } from "vitest"; +import { FILTERS } from "@/app/api/providers/suggested-models/filters.js"; +import { getProviderModelsFetcher } from "@/shared/constants/providers"; + +describe("HuggingFace suggested models", () => { + it("exposes kind-specific fetchers for media providers", () => { + const imageFetcher = getProviderModelsFetcher("huggingface", "image"); + const sttFetcher = getProviderModelsFetcher("huggingface", "stt"); + const imageUrl = new URL(imageFetcher.url); + const sttUrl = new URL(sttFetcher.url); + + expect(imageFetcher.type).toBe("huggingface-hub"); + expect(imageUrl.origin + imageUrl.pathname).toBe("https://huggingface.co/api/models"); + expect(imageUrl.searchParams.get("inference_provider")).toBe("hf-inference"); + expect(imageUrl.searchParams.get("pipeline_tag")).toBe("text-to-image"); + expect(imageUrl.searchParams.get("limit")).toBe("30"); + + expect(sttFetcher.type).toBe("huggingface-hub"); + expect(sttUrl.origin + sttUrl.pathname).toBe("https://huggingface.co/api/models"); + expect(sttUrl.searchParams.get("inference_provider")).toBe("hf-inference"); + expect(sttUrl.searchParams.get("pipeline_tag")).toBe("automatic-speech-recognition"); + expect(sttUrl.searchParams.get("limit")).toBe("30"); + }); + + it("sorts HuggingFace Hub suggestions by popularity", () => { + const models = FILTERS["huggingface-hub"]([ + { id: "org/model-c", downloads: 10, likes: 5 }, + { id: "org/model-a", downloads: 100, likes: 1 }, + { id: "org/model-b", downloads: 100, likes: 10 }, + ]); + + expect(models.map((model) => model.id)).toEqual([ + "org/model-b", + "org/model-a", + "org/model-c", + ]); + }); +}); diff --git a/tests/unit/model-test-routing.test.js b/tests/unit/model-test-routing.test.js index 3eaa57239b0..2621b72317c 100644 --- a/tests/unit/model-test-routing.test.js +++ b/tests/unit/model-test-routing.test.js @@ -147,7 +147,7 @@ describe("model test route kind routing", () => { method: "POST", headers: { "Content-Type": "application/json" }, body: JSON.stringify({ - model: "hf/openai/whisper-small", + model: "hf/openai/whisper-large-v3", kind: "stt", }), });