From 921ede741d539558d4c595f3ff88f6163f092b40 Mon Sep 17 00:00:00 2001 From: "Thomas G. Lopes" <26071571+TGlide@users.noreply.github.com> Date: Wed, 23 Jul 2025 16:53:13 +0100 Subject: [PATCH] fetch hf models and providers --- src/api/huggingface-models.ts | 17 ++ src/core/webview/webviewMessageHandler.ts | 16 ++ src/services/huggingface-models.ts | 171 ++++++++++++++++++ src/shared/ExtensionMessage.ts | 23 +++ src/shared/WebviewMessage.ts | 1 + .../settings/providers/HuggingFace.tsx | 151 +++++++++++++++- 6 files changed, 371 insertions(+), 8 deletions(-) create mode 100644 src/api/huggingface-models.ts create mode 100644 src/services/huggingface-models.ts diff --git a/src/api/huggingface-models.ts b/src/api/huggingface-models.ts new file mode 100644 index 0000000000..ec1915d0e3 --- /dev/null +++ b/src/api/huggingface-models.ts @@ -0,0 +1,17 @@ +import { fetchHuggingFaceModels, type HuggingFaceModel } from "../services/huggingface-models" + +export interface HuggingFaceModelsResponse { + models: HuggingFaceModel[] + cached: boolean + timestamp: number +} + +export async function getHuggingFaceModels(): Promise { + const models = await fetchHuggingFaceModels() + + return { + models, + cached: false, // We could enhance this to track if data came from cache + timestamp: Date.now(), + } +} diff --git a/src/core/webview/webviewMessageHandler.ts b/src/core/webview/webviewMessageHandler.ts index 780d40df89..ebe95530f2 100644 --- a/src/core/webview/webviewMessageHandler.ts +++ b/src/core/webview/webviewMessageHandler.ts @@ -674,6 +674,22 @@ export const webviewMessageHandler = async ( // TODO: Cache like we do for OpenRouter, etc? provider.postMessageToWebview({ type: "vsCodeLmModels", vsCodeLmModels }) break + case "requestHuggingFaceModels": + try { + const { getHuggingFaceModels } = await import("../../api/huggingface-models") + const huggingFaceModelsResponse = await getHuggingFaceModels() + provider.postMessageToWebview({ + type: "huggingFaceModels", + huggingFaceModels: huggingFaceModelsResponse.models, + }) + } catch (error) { + console.error("Failed to fetch Hugging Face models:", error) + provider.postMessageToWebview({ + type: "huggingFaceModels", + huggingFaceModels: [], + }) + } + break case "openImage": openImage(message.text!, { values: message.values }) break diff --git a/src/services/huggingface-models.ts b/src/services/huggingface-models.ts new file mode 100644 index 0000000000..9c0bc406f9 --- /dev/null +++ b/src/services/huggingface-models.ts @@ -0,0 +1,171 @@ +export interface HuggingFaceModel { + _id: string + id: string + inferenceProviderMapping: InferenceProviderMapping[] + trendingScore: number + config: ModelConfig + tags: string[] + pipeline_tag: "text-generation" | "image-text-to-text" + library_name?: string +} + +export interface InferenceProviderMapping { + provider: string + providerId: string + status: "live" | "staging" | "error" + task: "conversational" +} + +export interface ModelConfig { + architectures: string[] + model_type: string + tokenizer_config?: { + chat_template?: string | Array<{ name: string; template: string }> + model_max_length?: number + } +} + +interface HuggingFaceApiParams { + pipeline_tag?: "text-generation" | "image-text-to-text" + filter: string + inference_provider: string + limit: number + expand: string[] +} + +const DEFAULT_PARAMS: HuggingFaceApiParams = { + filter: "conversational", + inference_provider: "all", + limit: 100, + expand: [ + "inferenceProviderMapping", + "config", + "library_name", + "pipeline_tag", + "tags", + "mask_token", + "trendingScore", + ], +} + +const BASE_URL = "https://huggingface.co/api/models" +const CACHE_DURATION = 1000 * 60 * 60 // 1 hour + +interface CacheEntry { + data: HuggingFaceModel[] + timestamp: number + status: "success" | "partial" | "error" +} + +let cache: CacheEntry | null = null + +function buildApiUrl(params: HuggingFaceApiParams): string { + const url = new URL(BASE_URL) + + // Add simple params + Object.entries(params).forEach(([key, value]) => { + if (!Array.isArray(value)) { + url.searchParams.append(key, String(value)) + } + }) + + // Handle array params specially + params.expand.forEach((item) => { + url.searchParams.append("expand[]", item) + }) + + return url.toString() +} + +const headers: HeadersInit = { + "Upgrade-Insecure-Requests": "1", + "Sec-Fetch-Dest": "document", + "Sec-Fetch-Mode": "navigate", + "Sec-Fetch-Site": "none", + "Sec-Fetch-User": "?1", + Priority: "u=0, i", + Pragma: "no-cache", + "Cache-Control": "no-cache", +} + +const requestInit: RequestInit = { + credentials: "include", + headers, + method: "GET", + mode: "cors", +} + +export async function fetchHuggingFaceModels(): Promise { + const now = Date.now() + + // Check cache + if (cache && now - cache.timestamp < CACHE_DURATION) { + console.log("Using cached Hugging Face models") + return cache.data + } + + try { + console.log("Fetching Hugging Face models from API...") + + // Fetch both text-generation and image-text-to-text models in parallel + const [textGenResponse, imgTextResponse] = await Promise.allSettled([ + fetch(buildApiUrl({ ...DEFAULT_PARAMS, pipeline_tag: "text-generation" }), requestInit), + fetch(buildApiUrl({ ...DEFAULT_PARAMS, pipeline_tag: "image-text-to-text" }), requestInit), + ]) + + let textGenModels: HuggingFaceModel[] = [] + let imgTextModels: HuggingFaceModel[] = [] + let hasErrors = false + + // Process text-generation models + if (textGenResponse.status === "fulfilled" && textGenResponse.value.ok) { + textGenModels = await textGenResponse.value.json() + } else { + console.error("Failed to fetch text-generation models:", textGenResponse) + hasErrors = true + } + + // Process image-text-to-text models + if (imgTextResponse.status === "fulfilled" && imgTextResponse.value.ok) { + imgTextModels = await imgTextResponse.value.json() + } else { + console.error("Failed to fetch image-text-to-text models:", imgTextResponse) + hasErrors = true + } + + // Combine and filter models + const allModels = [...textGenModels, ...imgTextModels] + .filter((model) => model.inferenceProviderMapping.length > 0) + .sort((a, b) => a.id.toLowerCase().localeCompare(b.id.toLowerCase())) + + // Update cache + cache = { + data: allModels, + timestamp: now, + status: hasErrors ? "partial" : "success", + } + + console.log(`Fetched ${allModels.length} Hugging Face models (status: ${cache.status})`) + return allModels + } catch (error) { + console.error("Error fetching Hugging Face models:", error) + + // Return cached data if available + if (cache) { + console.log("Using stale cached data due to fetch error") + cache.status = "error" + return cache.data + } + + // No cache available, return empty array + return [] + } +} + +export function getCachedModels(): HuggingFaceModel[] | null { + return cache?.data || null +} + +export function clearCache(): void { + cache = null +} diff --git a/src/shared/ExtensionMessage.ts b/src/shared/ExtensionMessage.ts index 4f2aa2da15..ef45c7ebe3 100644 --- a/src/shared/ExtensionMessage.ts +++ b/src/shared/ExtensionMessage.ts @@ -67,6 +67,7 @@ export interface ExtensionMessage { | "ollamaModels" | "lmStudioModels" | "vsCodeLmModels" + | "huggingFaceModels" | "vsCodeLmApiAvailable" | "updatePrompt" | "systemPrompt" @@ -135,6 +136,28 @@ export interface ExtensionMessage { ollamaModels?: string[] lmStudioModels?: string[] vsCodeLmModels?: { vendor?: string; family?: string; version?: string; id?: string }[] + huggingFaceModels?: Array<{ + _id: string + id: string + inferenceProviderMapping: Array<{ + provider: string + providerId: string + status: "live" | "staging" | "error" + task: "conversational" + }> + trendingScore: number + config: { + architectures: string[] + model_type: string + tokenizer_config?: { + chat_template?: string | Array<{ name: string; template: string }> + model_max_length?: number + } + } + tags: string[] + pipeline_tag: "text-generation" | "image-text-to-text" + library_name?: string + }> mcpServers?: McpServer[] commits?: GitCommit[] listApiConfig?: ProviderSettingsEntry[] diff --git a/src/shared/WebviewMessage.ts b/src/shared/WebviewMessage.ts index 1f56829f7b..b0529c5a27 100644 --- a/src/shared/WebviewMessage.ts +++ b/src/shared/WebviewMessage.ts @@ -67,6 +67,7 @@ export interface WebviewMessage { | "requestOllamaModels" | "requestLmStudioModels" | "requestVsCodeLmModels" + | "requestHuggingFaceModels" | "openImage" | "saveImage" | "openFile" diff --git a/webview-ui/src/components/settings/providers/HuggingFace.tsx b/webview-ui/src/components/settings/providers/HuggingFace.tsx index 2eb515cbb6..d1f7aa29fd 100644 --- a/webview-ui/src/components/settings/providers/HuggingFace.tsx +++ b/webview-ui/src/components/settings/providers/HuggingFace.tsx @@ -1,13 +1,40 @@ -import { useCallback } from "react" +import { useCallback, useState, useEffect, useMemo } from "react" +import { useEvent } from "react-use" import { VSCodeTextField } from "@vscode/webview-ui-toolkit/react" import type { ProviderSettings } from "@roo-code/types" +import { ExtensionMessage } from "@roo/ExtensionMessage" +import { vscode } from "@src/utils/vscode" import { useAppTranslation } from "@src/i18n/TranslationContext" import { VSCodeButtonLink } from "@src/components/common/VSCodeButtonLink" +import { SearchableSelect, type SearchableSelectOption } from "@src/components/ui" import { inputEventTransform } from "../transforms" +type HuggingFaceModel = { + _id: string + id: string + inferenceProviderMapping: Array<{ + provider: string + providerId: string + status: "live" | "staging" | "error" + task: "conversational" + }> + trendingScore: number + config: { + architectures: string[] + model_type: string + tokenizer_config?: { + chat_template?: string | Array<{ name: string; template: string }> + model_max_length?: number + } + } + tags: string[] + pipeline_tag: "text-generation" | "image-text-to-text" + library_name?: string +} + type HuggingFaceProps = { apiConfiguration: ProviderSettings setApiConfigurationField: (field: keyof ProviderSettings, value: ProviderSettings[keyof ProviderSettings]) => void @@ -15,6 +42,9 @@ type HuggingFaceProps = { export const HuggingFace = ({ apiConfiguration, setApiConfigurationField }: HuggingFaceProps) => { const { t } = useAppTranslation() + const [models, setModels] = useState([]) + const [loading, setLoading] = useState(false) + const [selectedProvider, setSelectedProvider] = useState("") const handleInputChange = useCallback( ( @@ -27,6 +57,71 @@ export const HuggingFace = ({ apiConfiguration, setApiConfigurationField }: Hugg [setApiConfigurationField], ) + // Fetch models when component mounts + useEffect(() => { + setLoading(true) + vscode.postMessage({ type: "requestHuggingFaceModels" }) + }, []) + + // Handle messages from extension + const onMessage = useCallback((event: MessageEvent) => { + const message: ExtensionMessage = event.data + + switch (message.type) { + case "huggingFaceModels": + setModels(message.huggingFaceModels || []) + setLoading(false) + break + } + }, []) + + useEvent("message", onMessage) + + // Get current model and its providers + const currentModel = models.find((m) => m.id === apiConfiguration?.huggingFaceModelId) + const availableProviders = useMemo( + () => currentModel?.inferenceProviderMapping || [], + [currentModel?.inferenceProviderMapping], + ) + + // Set default provider when model changes + useEffect(() => { + if (currentModel && availableProviders.length > 0) { + const currentProvider = availableProviders.find((p) => p.provider === selectedProvider) + if (!currentProvider) { + // Set to first available provider or "auto" + setSelectedProvider("auto") + } + } + }, [currentModel, availableProviders, selectedProvider]) + + const handleModelSelect = (modelId: string) => { + setApiConfigurationField("huggingFaceModelId", modelId) + // Reset provider selection when model changes + setSelectedProvider("auto") + } + + const handleProviderSelect = (provider: string) => { + setSelectedProvider(provider) + // You could store this in a separate field if needed + } + + // Format provider name for display + const formatProviderName = (provider: string) => { + const nameMap: Record = { + sambanova: "SambaNova", + "fireworks-ai": "Fireworks", + together: "Together AI", + nebius: "Nebius AI Studio", + hyperbolic: "Hyperbolic", + novita: "Novita", + cohere: "Cohere", + "hf-inference": "HF Inference API", + replicate: "Replicate", + } + return nameMap[provider] || provider.charAt(0).toUpperCase() + provider.slice(1) + } + return ( <> - - - + +
+ + + ({ + value: model.id, + label: model.id, + }), + )} + placeholder="Select a model..." + searchPlaceholder="Search models..." + emptyMessage="No models found" + disabled={loading} + /> +
+ + {currentModel && availableProviders.length > 0 && ( +
+ + ({ + value: mapping.provider, + label: `${formatProviderName(mapping.provider)} (${mapping.status})`, + }), + ), + ]} + placeholder="Select a provider..." + searchPlaceholder="Search providers..." + emptyMessage="No providers found" + /> +
+ )} +
{t("settings:providers.apiKeyStorageNotice")}
+ {!apiConfiguration?.huggingFaceApiKey && ( {t("settings:providers.getHuggingFaceApiKey")}