diff --git a/backend/windmill-api/src/ai.rs b/backend/windmill-api/src/ai.rs index fc60c65c3c..9f13d32cf5 100644 --- a/backend/windmill-api/src/ai.rs +++ b/backend/windmill-api/src/ai.rs @@ -10,6 +10,7 @@ use reqwest::{Client, RequestBuilder}; use serde::{Deserialize, Serialize}; use serde_json::value::RawValue; use std::collections::HashMap; +use windmill_common::ai_providers::{AIProvider, ProviderConfig, ProviderModel}; use windmill_audit::{audit_oss::audit_log, ActionKind}; use windmill_common::error::{to_anyhow, Error, Result}; @@ -261,90 +262,6 @@ impl ExpiringAIRequestConfig { } } -#[derive(Serialize, Deserialize, Debug, Eq, PartialEq, Hash, Clone)] -#[serde(rename_all = "lowercase")] -pub enum AIProvider { - OpenAI, - #[serde(rename = "azure_openai")] - AzureOpenAI, - Anthropic, - Mistral, - DeepSeek, - GoogleAI, - Groq, - OpenRouter, - TogetherAI, - CustomAI, -} - -impl AIProvider { - pub async fn get_base_url(&self, resource_base_url: Option, db: &DB) -> Result { - match self { - AIProvider::OpenAI => { - let azure_base_path = sqlx::query_scalar!( - "SELECT value - FROM global_settings - WHERE name = 'openai_azure_base_path'", - ) - .fetch_optional(db) - .await?; - - let azure_base_path = if let Some(azure_base_path) = azure_base_path { - Some( - serde_json::from_value::(azure_base_path).map_err(|e| { - Error::internal_err(format!("validating openai azure base path {e:#}")) - })?, - ) - } else { - OPENAI_AZURE_BASE_PATH.clone() - }; - - Ok(azure_base_path.unwrap_or(OPENAI_BASE_URL.to_string())) - } - AIProvider::DeepSeek => Ok("https://api.deepseek.com/v1".to_string()), - AIProvider::GoogleAI => { - Ok("https://generativelanguage.googleapis.com/v1beta/openai".to_string()) - } - AIProvider::Groq => Ok("https://api.groq.com/openai/v1".to_string()), - AIProvider::OpenRouter => Ok("https://openrouter.ai/api/v1".to_string()), - AIProvider::TogetherAI => Ok("https://api.together.xyz/v1".to_string()), - AIProvider::Anthropic => Ok("https://api.anthropic.com/v1".to_string()), - AIProvider::Mistral => Ok("https://api.mistral.ai/v1".to_string()), - p @ (AIProvider::CustomAI | AIProvider::AzureOpenAI) => { - if let Some(base_url) = resource_base_url { - Ok(base_url) - } else { - Err(Error::BadRequest(format!( - "{:?} provider requires a base URL in the resource", - p - ))) - } - } - } - } -} - -impl TryFrom<&str> for AIProvider { - type Error = Error; - fn try_from(s: &str) -> Result { - let s = serde_json::from_value::(serde_json::Value::String(s.to_string())) - .map_err(|e| Error::BadRequest(format!("Invalid AI provider: {}", e)))?; - Ok(s) - } -} - -#[derive(Serialize, Deserialize, Debug)] -pub struct ProviderConfig { - pub resource_path: String, - pub models: Vec, -} - -#[derive(Serialize, Deserialize, Debug)] -pub struct ProviderModel { - pub model: String, - pub provider: AIProvider, -} - #[derive(Serialize, Deserialize, Debug)] pub struct AIConfig { #[serde(skip_serializing_if = "Option::is_none")] diff --git a/backend/windmill-common/src/ai_providers.rs b/backend/windmill-common/src/ai_providers.rs new file mode 100644 index 0000000000..577f613f21 --- /dev/null +++ b/backend/windmill-common/src/ai_providers.rs @@ -0,0 +1,102 @@ +/* + * This file contains shared AI provider utilities used by both the API and worker. + */ + +use crate::db::DB; +use crate::error::{Error, Result}; +use serde::{Deserialize, Serialize}; + +lazy_static::lazy_static! { + static ref OPENAI_AZURE_BASE_PATH: Option = std::env::var("OPENAI_AZURE_BASE_PATH").ok(); +} + +#[derive(Serialize, Deserialize, Debug, Eq, PartialEq, Hash, Clone)] +#[serde(rename_all = "lowercase")] +pub enum AIProvider { + OpenAI, + #[serde(rename = "azure_openai")] + AzureOpenAI, + Anthropic, + Mistral, + DeepSeek, + GoogleAI, + Groq, + OpenRouter, + TogetherAI, + CustomAI, +} + +impl AIProvider { + /// Get the base URL for the AI provider + pub async fn get_base_url(&self, resource_base_url: Option, db: &DB) -> Result { + match self { + AIProvider::OpenAI => { + // Check for Azure base path override + let azure_base_path = sqlx::query_scalar!( + "SELECT value + FROM global_settings + WHERE name = 'openai_azure_base_path'", + ) + .fetch_optional(db) + .await?; + + let azure_base_path = if let Some(azure_base_path) = azure_base_path { + Some( + serde_json::from_value::(azure_base_path).map_err(|e| { + Error::internal_err(format!("validating openai azure base path {e:#}")) + })?, + ) + } else { + OPENAI_AZURE_BASE_PATH.clone() + }; + + Ok(azure_base_path.unwrap_or("https://api.openai.com/v1".to_string())) + } + AIProvider::DeepSeek => Ok("https://api.deepseek.com/v1".to_string()), + AIProvider::GoogleAI => { + Ok("https://generativelanguage.googleapis.com/v1beta/openai".to_string()) + } + AIProvider::Groq => Ok("https://api.groq.com/openai/v1".to_string()), + AIProvider::OpenRouter => Ok("https://openrouter.ai/api/v1".to_string()), + AIProvider::TogetherAI => Ok("https://api.together.xyz/v1".to_string()), + AIProvider::Anthropic => Ok("https://api.anthropic.com/v1".to_string()), + AIProvider::Mistral => Ok("https://api.mistral.ai/v1".to_string()), + p @ (AIProvider::CustomAI | AIProvider::AzureOpenAI) => { + if let Some(base_url) = resource_base_url { + Ok(base_url) + } else { + Err(Error::BadRequest(format!( + "{:?} provider requires a base URL in the resource", + p + ))) + } + } + } + } + + /// Check if this provider is Anthropic (needs special handling) + pub fn is_anthropic(&self) -> bool { + matches!(self, AIProvider::Anthropic) + } +} + +impl TryFrom<&str> for AIProvider { + type Error = Error; + fn try_from(s: &str) -> Result { + let s = serde_json::from_value::(serde_json::Value::String(s.to_string())) + .map_err(|e| Error::BadRequest(format!("Invalid AI provider: {}", e)))?; + Ok(s) + } +} + +#[derive(Serialize, Deserialize, Debug)] +pub struct ProviderConfig { + pub resource_path: String, + pub models: Vec, +} + +#[derive(Serialize, Deserialize, Debug)] +pub struct ProviderModel { + pub model: String, + pub provider: AIProvider, +} diff --git a/backend/windmill-common/src/lib.rs b/backend/windmill-common/src/lib.rs index 9dd0b09809..4980e81700 100644 --- a/backend/windmill-common/src/lib.rs +++ b/backend/windmill-common/src/lib.rs @@ -26,6 +26,7 @@ use scripts::ScriptLang; use sqlx::{Acquire, Postgres}; pub mod agent_workers; +pub mod ai_providers; pub mod apps; pub mod assets; pub mod auth; diff --git a/backend/windmill-worker/src/ai_executor.rs b/backend/windmill-worker/src/ai_executor.rs index 35dc6b4ee8..a2ad22ec84 100644 --- a/backend/windmill-worker/src/ai_executor.rs +++ b/backend/windmill-worker/src/ai_executor.rs @@ -6,6 +6,7 @@ use std::{collections::HashMap, sync::Arc}; #[cfg(feature = "benchmark")] use windmill_common::bench::BenchmarkIter; use windmill_common::{ + ai_providers::AIProvider, auth::get_job_perms, cache, client::AuthedClient, @@ -155,7 +156,7 @@ struct Tool { #[derive(Deserialize, Debug)] struct AIAgentArgs { - provider: Provider, + provider: ProviderWithResource, system_prompt: Option, user_message: String, temperature: Option, @@ -167,35 +168,30 @@ struct AIAgentArgs { struct ProviderResource { #[serde(alias = "apiKey")] api_key: String, + #[serde(alias = "baseUrl")] + base_url: Option, } #[derive(Deserialize, Debug)] -#[serde(tag = "kind")] -enum Provider { - OpenAI { resource: ProviderResource, model: String }, - Anthropic { resource: ProviderResource, model: String }, +struct ProviderWithResource { + kind: AIProvider, + resource: ProviderResource, + model: String, } -impl Provider { +impl ProviderWithResource { fn get_api_key(&self) -> &str { - match self { - Provider::OpenAI { resource, .. } => &resource.api_key, - Provider::Anthropic { resource, .. } => &resource.api_key, - } + &self.resource.api_key } fn get_model(&self) -> &str { - match self { - Provider::OpenAI { model, .. } => model, - Provider::Anthropic { model, .. } => model, - } + &self.model } - fn get_base_url(&self) -> &str { - match self { - Provider::OpenAI { .. } => "https://api.openai.com/v1", - Provider::Anthropic { .. } => "https://api.anthropic.com/v1", - } + async fn get_base_url(&self, db: &DB) -> Result { + self.kind + .get_base_url(self.resource.base_url.clone(), db) + .await } } @@ -765,7 +761,7 @@ async fn run_agent( let mut content = None; - let base_url = args.provider.get_base_url(); + let base_url = args.provider.get_base_url(db).await?; let api_key = args.provider.get_api_key(); let mut tool_defs: Option> = if tools.is_empty() { @@ -780,7 +776,10 @@ async fn run_agent( .and_then(|schema| schema.properties.as_ref()) .map(|props| !props.is_empty()) .unwrap_or(false); - let is_anthropic = matches!(args.provider, Provider::Anthropic { .. }); + let provider_is_anthropic = args.provider.kind.is_anthropic(); + let is_openrouter_anthropic = args.provider.kind == AIProvider::OpenRouter + && args.provider.model.starts_with("anthropic/"); + let is_anthropic = provider_is_anthropic || is_openrouter_anthropic; let mut response_format: Option = None; let mut used_structured_output_tool = false; let mut structured_output_tool_name: Option = None; diff --git a/frontend/src/lib/components/AIProviderPicker.svelte b/frontend/src/lib/components/AIProviderPicker.svelte new file mode 100644 index 0000000000..8979de8915 --- /dev/null +++ b/frontend/src/lib/components/AIProviderPicker.svelte @@ -0,0 +1,197 @@ + + +
+ +
+ + {#snippet children({ item })} + {#each providerOptions as option} + + {/each} + {/snippet} + +
+ + +
+
+

resource

+ resourceValueToPath(value?.resource), + (v) => { + if (value) { + value.resource = pathToResourceValue(v) + } + } + } + resourceType={value?.kind} + disabled={disabled || !value?.kind} + placeholder="Select resource" + selectFirst={true} + /> +
+ + +
+

model

+ -
+
{@render children?.({ item, disabled })}
diff --git a/frontend/src/lib/components/copilot/MetadataGen.svelte b/frontend/src/lib/components/copilot/MetadataGen.svelte index 91b6c43cbf..f31558327e 100644 --- a/frontend/src/lib/components/copilot/MetadataGen.svelte +++ b/frontend/src/lib/components/copilot/MetadataGen.svelte @@ -324,7 +324,7 @@ Generate a tool name for the script below: on:blur={() => (focused = false)} /> {#if promptConfigName === 'agentToolFunctionName' && !validateToolName(content ?? '')} -
+
Invalid tool name, should only contain letters, numbers and underscores
{/if} diff --git a/frontend/src/lib/components/copilot/lib.ts b/frontend/src/lib/components/copilot/lib.ts index 5cacb53894..dacd221db0 100644 --- a/frontend/src/lib/components/copilot/lib.ts +++ b/frontend/src/lib/components/copilot/lib.ts @@ -27,6 +27,11 @@ import type { Stream } from 'openai/core/streaming.mjs' export const SUPPORTED_LANGUAGES = new Set(Object.keys(GEN_CONFIG.prompts)) +interface AIProviderDetails { + label: string + defaultModels: string[] +} + const OPENAI_MODELS = [ 'gpt-5', 'gpt-5-mini', @@ -38,18 +43,47 @@ const OPENAI_MODELS = [ 'o3-mini' ] -// need at least one model for each provider except customai -export const AI_DEFAULT_MODELS: Record = { - openai: OPENAI_MODELS, - azure_openai: OPENAI_MODELS, - anthropic: ['claude-sonnet-4-0', 'claude-sonnet-4-0/thinking', 'claude-3-5-haiku-latest'], - mistral: ['codestral-latest'], - deepseek: ['deepseek-chat', 'deepseek-reasoner'], - googleai: ['gemini-2.0-flash', 'gemini-1.5-flash', 'gemini-1.5-pro'], - groq: ['llama-3.3-70b-versatile', 'llama-3.1-8b-instant'], - openrouter: ['meta-llama/llama-3.2-3b-instruct:free'], - togetherai: ['meta-llama/Llama-3.3-70B-Instruct-Turbo'], - customai: [] +export const AI_PROVIDERS: Record = { + openai: { + label: 'OpenAI', + defaultModels: OPENAI_MODELS + }, + azure_openai: { + label: 'Azure OpenAI', + defaultModels: OPENAI_MODELS + }, + anthropic: { + label: 'Anthropic', + defaultModels: ['claude-sonnet-4-0', 'claude-sonnet-4-0/thinking', 'claude-3-5-haiku-latest'] + }, + mistral: { + label: 'Mistral', + defaultModels: ['codestral-latest'] + }, + deepseek: { + label: 'DeepSeek', + defaultModels: ['deepseek-chat', 'deepseek-reasoner'] + }, + googleai: { + label: 'Google AI', + defaultModels: ['gemini-2.0-flash', 'gemini-1.5-flash', 'gemini-1.5-pro'] + }, + groq: { + label: 'Groq', + defaultModels: ['llama-3.3-70b-versatile', 'llama-3.1-8b-instant'] + }, + openrouter: { + label: 'OpenRouter', + defaultModels: ['meta-llama/llama-3.2-3b-instruct:free'] + }, + togetherai: { + label: 'Together AI', + defaultModels: ['meta-llama/Llama-3.3-70B-Instruct-Turbo'] + }, + customai: { + label: 'Custom AI', + defaultModels: [] + } } export interface ModelResponse { @@ -67,9 +101,11 @@ export interface ModelResponse { export async function fetchAvailableModels( resourcePath: string, workspace: string, - provider: AIProvider + provider: AIProvider, + signal?: AbortSignal ): Promise { const models = await fetch(`${location.origin}${OpenAPI.BASE}/w/${workspace}/ai/proxy/models`, { + signal, headers: { 'X-Resource-Path': resourcePath, 'X-Provider': provider, @@ -82,6 +118,16 @@ export async function fetchAvailableModels( } const data = (await models.json()) as { data: ModelResponse[] } if (data.data.length > 0) { + const sortFunc = (provider: AIProvider) => (a: string, b: string) => { + // First prioritize models in defaultModels array + const defaultModels = AI_PROVIDERS[provider]?.defaultModels || [] + const aInDefault = defaultModels.includes(a) + const bInDefault = defaultModels.includes(b) + + if (aInDefault && !bInDefault) return -1 + if (!aInDefault && bInDefault) return 1 + return 0 + } switch (provider) { case 'openai': return data.data @@ -89,6 +135,7 @@ export async function fetchAvailableModels( (m) => m.id.startsWith('gpt-') || m.id.startsWith('o') || m.id.startsWith('codex') ) .map((m) => m.id) + .sort(sortFunc(provider)) case 'azure_openai': return data.data .filter( @@ -98,10 +145,11 @@ export async function fetchAvailableModels( (m.capabilities.completion || m.capabilities.chat_completion) ) .map((m) => m.id) + .sort(sortFunc(provider)) case 'googleai': - return data.data.map((m) => m.id.split('/')[1]) + return data.data.map((m) => m.id.split('/')[1]).sort(sortFunc(provider)) default: - return data.data.map((m) => m.id) + return data.data.map((m) => m.id).sort(sortFunc(provider)) } } @@ -291,7 +339,7 @@ export async function testKey({ if (!apiKey && !resourcePath) { throw new Error('API key or resource path is required') } - const modelToTest = model ?? AI_DEFAULT_MODELS[aiProvider][0] + const modelToTest = model ?? AI_PROVIDERS[aiProvider].defaultModels[0] if (!modelToTest) { throw new Error('Missing a model to test') diff --git a/frontend/src/lib/components/flows/flowInfers.ts b/frontend/src/lib/components/flows/flowInfers.ts index ec84a83c6d..3b4f36aebe 100644 --- a/frontend/src/lib/components/flows/flowInfers.ts +++ b/frontend/src/lib/components/flows/flowInfers.ts @@ -60,41 +60,7 @@ export async function loadSchemaFromModule(module: FlowModule): Promise<{ properties: { provider: { type: 'object', - oneOf: [ - { - type: 'object', - title: 'OpenAI', - properties: { - kind: { type: 'string', enum: ['OpenAI'] }, - resource: { - type: 'object', - format: 'resource-openai' - }, - - model: { - type: 'string', - enum: ['gpt-5', 'gpt-5-mini', 'gpt-5-nano', 'gpt-4.1', 'gpt-4o', 'gpt-4o-mini'] - } - }, - required: ['kind', 'resource', 'model'] - }, - { - type: 'object', - title: 'Anthropic', - properties: { - kind: { type: 'string', enum: ['Anthropic'] }, - resource: { - type: 'object', - format: 'resource-anthropic' - }, - model: { - type: 'string', - enum: ['claude-sonnet-4-0', 'claude-3-7-sonnet-latest', 'claude-3-5-haiku-latest'] - } - }, - required: ['kind', 'resource', 'model'] - } - ] + format: 'ai-provider' }, user_message: { type: 'string' diff --git a/frontend/src/lib/components/workspaceSettings/AISettings.svelte b/frontend/src/lib/components/workspaceSettings/AISettings.svelte index b4cbd5eb09..0b34bf0b54 100644 --- a/frontend/src/lib/components/workspaceSettings/AISettings.svelte +++ b/frontend/src/lib/components/workspaceSettings/AISettings.svelte @@ -2,7 +2,7 @@ import { WorkspaceService, type AIConfig, type AIProvider } from '$lib/gen' import { setCopilotInfo, workspaceStore } from '$lib/stores' import { sendUserToast } from '$lib/toast' - import { AI_DEFAULT_MODELS, fetchAvailableModels } from '../copilot/lib' + import { AI_PROVIDERS, fetchAvailableModels } from '../copilot/lib' import TestAiKey from '../copilot/TestAIKey.svelte' import Description from '../Description.svelte' import Label from '../Label.svelte' @@ -19,19 +19,6 @@ import ToggleButton from '../common/toggleButton-v2/ToggleButton.svelte' import autosize from '$lib/autosize' - const aiProviderLabels: [AIProvider, string][] = [ - ['openai', 'OpenAI'], - ['azure_openai', 'Azure OpenAI'], - ['anthropic', 'Anthropic'], - ['mistral', 'Mistral'], - ['deepseek', 'DeepSeek'], - ['googleai', 'Google AI'], - ['groq', 'Groq'], - ['openrouter', 'OpenRouter'], - ['togetherai', 'Together AI'], - ['customai', 'Custom AI'] - ] - const MAX_CUSTOM_PROMPT_LENGTH = 5000 let { @@ -51,7 +38,7 @@ let fetchedAiModels = $state(false) let availableAiModels = $state( Object.fromEntries( - aiProviderLabels.map(([provider]) => [provider, AI_DEFAULT_MODELS[provider]]) + Object.keys(AI_PROVIDERS).map((provider) => [provider, AI_PROVIDERS[provider].defaultModels]) ) as Record ) @@ -88,7 +75,7 @@ availableAiModels[provider] = models } catch (e) { console.error('failed to fetch models for provider', provider, e) - availableAiModels[provider] = AI_DEFAULT_MODELS[provider] + availableAiModels[provider] = AI_PROVIDERS[provider].defaultModels } } fetchedAiModels = true @@ -142,7 +129,7 @@ availableAiModels[provider] = models } catch (e) { console.error('failed to fetch models for provider', provider, e) - availableAiModels[provider] = AI_DEFAULT_MODELS[provider] + availableAiModels[provider] = AI_PROVIDERS[provider].defaultModels } } @@ -169,12 +156,12 @@

AI Providers

- {#each aiProviderLabels as [provider, label]} + {#each Object.entries(AI_PROVIDERS) as [provider, details]}
{ @@ -240,12 +227,12 @@ () => aiProviders[provider].resource_path, (v) => { aiProviders[provider].resource_path = v - onAiProviderChange(provider) + onAiProviderChange(provider as AIProvider) } } /> diff --git a/frontend/src/lib/utils.ts b/frontend/src/lib/utils.ts index 15ae837934..55d021392b 100644 --- a/frontend/src/lib/utils.ts +++ b/frontend/src/lib/utils.ts @@ -593,6 +593,7 @@ export type InputCat = | 'oneOf' | 'dynamic' | 'json-schema' + | 'ai-provider' export namespace DynamicInput { const DYN_FORMAT_PREFIX = ['dynmultiselect-', 'dynselect-'] @@ -657,6 +658,8 @@ export function setInputCat( return 'list' } else if (type == 'object' && format?.startsWith('resource')) { return 'resource-object' + } else if (type == 'object' && format == 'ai-provider') { + return 'ai-provider' } else if (type == 'object' && DynamicInput.isDynInputFormat(format)) { return 'dynamic' } else if (!type || type == 'object' || type == 'array') { diff --git a/frontend/src/routes/(root)/(logged)/user/(user)/create_workspace/+page.svelte b/frontend/src/routes/(root)/(logged)/user/(user)/create_workspace/+page.svelte index d489169927..37e9f733a5 100644 --- a/frontend/src/routes/(root)/(logged)/user/(user)/create_workspace/+page.svelte +++ b/frontend/src/routes/(root)/(logged)/user/(user)/create_workspace/+page.svelte @@ -26,7 +26,7 @@ import { isCloudHosted } from '$lib/cloud' import ToggleButtonGroup from '$lib/components/common/toggleButton-v2/ToggleButtonGroup.svelte' import ToggleButton from '$lib/components/common/toggleButton-v2/ToggleButton.svelte' - import { AI_DEFAULT_MODELS } from '$lib/components/copilot/lib' + import { AI_PROVIDERS } from '$lib/components/copilot/lib' const rd = $page.url.searchParams.get('rd') @@ -115,15 +115,15 @@ providers: { [selected]: { resource_path: path, - models: [AI_DEFAULT_MODELS[selected][0]] + models: [AI_PROVIDERS[selected].defaultModels[0]] } }, default_model: { - model: AI_DEFAULT_MODELS[selected][0], + model: AI_PROVIDERS[selected].defaultModels[0], provider: selected }, code_completion_model: codeCompletionEnabled - ? { model: AI_DEFAULT_MODELS[selected][0], provider: selected } + ? { model: AI_PROVIDERS[selected].defaultModels[0], provider: selected } : undefined } : {} @@ -291,7 +291,7 @@ apiKey={aiKey} disabled={!aiKey} aiProvider={selected} - model={AI_DEFAULT_MODELS[selected][0]} + model={AI_PROVIDERS[selected].defaultModels[0]} />
{#if aiKey}