diff --git a/packages/types/src/provider-settings.ts b/packages/types/src/provider-settings.ts index 884337767f..2373f8b97c 100644 --- a/packages/types/src/provider-settings.ts +++ b/packages/types/src/provider-settings.ts @@ -32,6 +32,7 @@ export const providerNames = [ "groq", "chutes", "litellm", + "huggingface", ] as const export const providerNamesSchema = z.enum(providerNames) @@ -219,6 +220,11 @@ const groqSchema = apiModelIdProviderModelSchema.extend({ groqApiKey: z.string().optional(), }) +const huggingFaceSchema = baseProviderSettingsSchema.extend({ + huggingFaceApiKey: z.string().optional(), + huggingFaceModelId: z.string().optional(), +}) + const chutesSchema = apiModelIdProviderModelSchema.extend({ chutesApiKey: z.string().optional(), }) @@ -256,6 +262,7 @@ export const providerSettingsSchemaDiscriminated = z.discriminatedUnion("apiProv fakeAiSchema.merge(z.object({ apiProvider: z.literal("fake-ai") })), xaiSchema.merge(z.object({ apiProvider: z.literal("xai") })), groqSchema.merge(z.object({ apiProvider: z.literal("groq") })), + huggingFaceSchema.merge(z.object({ apiProvider: z.literal("huggingface") })), chutesSchema.merge(z.object({ apiProvider: z.literal("chutes") })), litellmSchema.merge(z.object({ apiProvider: z.literal("litellm") })), defaultSchema, @@ -285,6 +292,7 @@ export const providerSettingsSchema = z.object({ ...fakeAiSchema.shape, ...xaiSchema.shape, ...groqSchema.shape, + ...huggingFaceSchema.shape, ...chutesSchema.shape, ...litellmSchema.shape, ...codebaseIndexProviderSchema.shape, @@ -304,6 +312,7 @@ export const MODEL_ID_KEYS: Partial[] = [ "unboundModelId", "requestyModelId", "litellmModelId", + "huggingFaceModelId", ] export const getModelId = (settings: ProviderSettings): string | undefined => { diff --git a/src/api/index.ts b/src/api/index.ts index 4598a711b2..bda390848c 100644 --- a/src/api/index.ts +++ b/src/api/index.ts @@ -26,6 +26,7 @@ import { FakeAIHandler, XAIHandler, GroqHandler, + HuggingFaceHandler, ChutesHandler, LiteLLMHandler, ClaudeCodeHandler, @@ -108,6 +109,8 @@ export function buildApiHandler(configuration: ProviderSettings): ApiHandler { return new XAIHandler(options) case "groq": return new GroqHandler(options) + case "huggingface": + return new HuggingFaceHandler(options) case "chutes": return new ChutesHandler(options) case "litellm": diff --git a/src/api/providers/huggingface.ts b/src/api/providers/huggingface.ts new file mode 100644 index 0000000000..913605bd92 --- /dev/null +++ b/src/api/providers/huggingface.ts @@ -0,0 +1,99 @@ +import OpenAI from "openai" +import { Anthropic } from "@anthropic-ai/sdk" + +import type { ApiHandlerOptions } from "../../shared/api" +import { ApiStream } from "../transform/stream" +import { convertToOpenAiMessages } from "../transform/openai-format" +import type { SingleCompletionHandler, ApiHandlerCreateMessageMetadata } from "../index" +import { DEFAULT_HEADERS } from "./constants" +import { BaseProvider } from "./base-provider" + +export class HuggingFaceHandler extends BaseProvider implements SingleCompletionHandler { + private client: OpenAI + private options: ApiHandlerOptions + + constructor(options: ApiHandlerOptions) { + super() + this.options = options + + if (!this.options.huggingFaceApiKey) { + throw new Error("Hugging Face API key is required") + } + + this.client = new OpenAI({ + baseURL: "https://router.huggingface.co/v1", + apiKey: this.options.huggingFaceApiKey, + defaultHeaders: DEFAULT_HEADERS, + }) + } + + override async *createMessage( + systemPrompt: string, + messages: Anthropic.Messages.MessageParam[], + metadata?: ApiHandlerCreateMessageMetadata, + ): ApiStream { + const modelId = this.options.huggingFaceModelId || "meta-llama/Llama-3.3-70B-Instruct" + const temperature = this.options.modelTemperature ?? 0.7 + + const params: OpenAI.Chat.Completions.ChatCompletionCreateParamsStreaming = { + model: modelId, + temperature, + messages: [{ role: "system", content: systemPrompt }, ...convertToOpenAiMessages(messages)], + stream: true, + stream_options: { include_usage: true }, + } + + const stream = await this.client.chat.completions.create(params) + + for await (const chunk of stream) { + const delta = chunk.choices[0]?.delta + + if (delta?.content) { + yield { + type: "text", + text: delta.content, + } + } + + if (chunk.usage) { + yield { + type: "usage", + inputTokens: chunk.usage.prompt_tokens || 0, + outputTokens: chunk.usage.completion_tokens || 0, + } + } + } + } + + async completePrompt(prompt: string): Promise { + const modelId = this.options.huggingFaceModelId || "meta-llama/Llama-3.3-70B-Instruct" + + try { + const response = await this.client.chat.completions.create({ + model: modelId, + messages: [{ role: "user", content: prompt }], + }) + + return response.choices[0]?.message.content || "" + } catch (error) { + if (error instanceof Error) { + throw new Error(`Hugging Face completion error: ${error.message}`) + } + + throw error + } + } + + override getModel() { + const modelId = this.options.huggingFaceModelId || "meta-llama/Llama-3.3-70B-Instruct" + return { + id: modelId, + info: { + maxTokens: 8192, + contextWindow: 131072, + supportsImages: false, + supportsPromptCache: false, + }, + } + } +} diff --git a/src/api/providers/index.ts b/src/api/providers/index.ts index 89d4c203ad..1cefd0616b 100644 --- a/src/api/providers/index.ts +++ b/src/api/providers/index.ts @@ -9,6 +9,7 @@ export { FakeAIHandler } from "./fake-ai" export { GeminiHandler } from "./gemini" export { GlamaHandler } from "./glama" export { GroqHandler } from "./groq" +export { HuggingFaceHandler } from "./huggingface" export { HumanRelayHandler } from "./human-relay" export { LiteLLMHandler } from "./lite-llm" export { LmStudioHandler } from "./lm-studio" diff --git a/webview-ui/src/components/settings/ApiOptions.tsx b/webview-ui/src/components/settings/ApiOptions.tsx index 6c6c621956..38d2ceebd3 100644 --- a/webview-ui/src/components/settings/ApiOptions.tsx +++ b/webview-ui/src/components/settings/ApiOptions.tsx @@ -59,6 +59,7 @@ import { Gemini, Glama, Groq, + HuggingFace, LMStudio, LiteLLM, Mistral, @@ -487,6 +488,10 @@ const ApiOptions = ({ )} + {selectedProvider === "huggingface" && ( + + )} + {selectedProvider === "chutes" && ( )} diff --git a/webview-ui/src/components/settings/constants.ts b/webview-ui/src/components/settings/constants.ts index 1140e4c0bc..995f591034 100644 --- a/webview-ui/src/components/settings/constants.ts +++ b/webview-ui/src/components/settings/constants.ts @@ -51,6 +51,7 @@ export const PROVIDERS = [ { value: "human-relay", label: "Human Relay" }, { value: "xai", label: "xAI (Grok)" }, { value: "groq", label: "Groq" }, + { value: "huggingface", label: "Hugging Face" }, { value: "chutes", label: "Chutes AI" }, { value: "litellm", label: "LiteLLM" }, ].sort((a, b) => a.label.localeCompare(b.label)) diff --git a/webview-ui/src/components/settings/providers/HuggingFace.tsx b/webview-ui/src/components/settings/providers/HuggingFace.tsx new file mode 100644 index 0000000000..2eb515cbb6 --- /dev/null +++ b/webview-ui/src/components/settings/providers/HuggingFace.tsx @@ -0,0 +1,57 @@ +import { useCallback } from "react" +import { VSCodeTextField } from "@vscode/webview-ui-toolkit/react" + +import type { ProviderSettings } from "@roo-code/types" + +import { useAppTranslation } from "@src/i18n/TranslationContext" +import { VSCodeButtonLink } from "@src/components/common/VSCodeButtonLink" + +import { inputEventTransform } from "../transforms" + +type HuggingFaceProps = { + apiConfiguration: ProviderSettings + setApiConfigurationField: (field: keyof ProviderSettings, value: ProviderSettings[keyof ProviderSettings]) => void +} + +export const HuggingFace = ({ apiConfiguration, setApiConfigurationField }: HuggingFaceProps) => { + const { t } = useAppTranslation() + + const handleInputChange = useCallback( + ( + field: K, + transform: (event: E) => ProviderSettings[K] = inputEventTransform, + ) => + (event: E | Event) => { + setApiConfigurationField(field, transform(event as E)) + }, + [setApiConfigurationField], + ) + + return ( + <> + + + + + + +
+ {t("settings:providers.apiKeyStorageNotice")} +
+ {!apiConfiguration?.huggingFaceApiKey && ( + + {t("settings:providers.getHuggingFaceApiKey")} + + )} + + ) +} diff --git a/webview-ui/src/components/settings/providers/index.ts b/webview-ui/src/components/settings/providers/index.ts index 54974f7200..6c6fdddaee 100644 --- a/webview-ui/src/components/settings/providers/index.ts +++ b/webview-ui/src/components/settings/providers/index.ts @@ -6,6 +6,7 @@ export { DeepSeek } from "./DeepSeek" export { Gemini } from "./Gemini" export { Glama } from "./Glama" export { Groq } from "./Groq" +export { HuggingFace } from "./HuggingFace" export { LMStudio } from "./LMStudio" export { Mistral } from "./Mistral" export { Moonshot } from "./Moonshot" diff --git a/webview-ui/src/components/ui/hooks/useSelectedModel.ts b/webview-ui/src/components/ui/hooks/useSelectedModel.ts index 928ebb42f4..8dceb6e117 100644 --- a/webview-ui/src/components/ui/hooks/useSelectedModel.ts +++ b/webview-ui/src/components/ui/hooks/useSelectedModel.ts @@ -130,6 +130,16 @@ function getSelectedModel({ const info = groqModels[id as keyof typeof groqModels] return { id, info } } + case "huggingface": { + const id = apiConfiguration.huggingFaceModelId ?? "meta-llama/Llama-3.3-70B-Instruct" + const info = { + maxTokens: 8192, + contextWindow: 131072, + supportsImages: false, + supportsPromptCache: false, + } + return { id, info } + } case "chutes": { const id = apiConfiguration.apiModelId ?? chutesDefaultModelId const info = chutesModels[id as keyof typeof chutesModels] diff --git a/webview-ui/src/i18n/locales/en/settings.json b/webview-ui/src/i18n/locales/en/settings.json index 4a826bddab..c8d0e69719 100644 --- a/webview-ui/src/i18n/locales/en/settings.json +++ b/webview-ui/src/i18n/locales/en/settings.json @@ -259,6 +259,9 @@ "geminiApiKey": "Gemini API Key", "getGroqApiKey": "Get Groq API Key", "groqApiKey": "Groq API Key", + "getHuggingFaceApiKey": "Get Hugging Face API Key", + "huggingFaceApiKey": "Hugging Face API Key", + "huggingFaceModelId": "Model ID", "getGeminiApiKey": "Get Gemini API Key", "openAiApiKey": "OpenAI API Key", "apiKey": "API Key",