diff --git a/backend/windmill-api/src/ai.rs b/backend/windmill-api/src/ai.rs index f72e2b4dc5..44a6b2986e 100644 --- a/backend/windmill-api/src/ai.rs +++ b/backend/windmill-api/src/ai.rs @@ -4,7 +4,7 @@ use crate::{ }; use axum::{body::Bytes, extract::Path, response::IntoResponse, routing::post, Extension, Router}; -use http::HeaderMap; +use http::{HeaderMap, Method}; use quick_cache::sync::Cache; use reqwest::{Client, RequestBuilder}; use serde::{Deserialize, Serialize}; @@ -141,6 +141,8 @@ impl AIRequestConfig { self, provider: &AIProvider, path: &str, + method: Method, + headers: HeaderMap, body: Bytes, ) -> Result { let body = if let Some(user) = self.user { @@ -153,8 +155,9 @@ impl AIRequestConfig { let is_azure = matches!(provider, AIProvider::OpenAI) && base_url != OPENAI_BASE_URL || matches!(provider, AIProvider::AzureOpenAI); + let is_anthropic = matches!(provider, AIProvider::Anthropic); - let url = if is_azure { + let url = if is_azure && method != Method::GET { if base_url.ends_with("/deployments") { let model = Self::get_azure_model(&body)?; format!("{}/{}/{}", base_url, model, path) @@ -171,9 +174,16 @@ impl AIRequestConfig { tracing::debug!("AI request URL: {}", url); let mut request = HTTP_CLIENT - .post(url) - .header("content-type", "application/json") - .body(body); + .request(method, url) + .header("content-type", "application/json"); + + for (header_name, header_value) in headers.iter() { + if header_name.to_string().starts_with("anthropic-") { + request = request.header(header_name, header_value); + } + } + + request = request.body(body); if is_azure { request = request.query(&[("api-version", AZURE_API_VERSION)]) @@ -181,9 +191,12 @@ impl AIRequestConfig { if let Some(api_key) = self.api_key { if is_azure { - request = request.header("api-key", api_key) + request = request.header("api-key", api_key.clone()) } else { - request = request.header("authorization", format!("Bearer {}", api_key)) + request = request.header("authorization", format!("Bearer {}", api_key.clone())) + } + if is_anthropic { + request = request.header("X-API-Key", api_key); } } @@ -339,17 +352,18 @@ pub struct AIConfig { } pub fn global_service() -> Router { - Router::new().route("/proxy/*ai", post(global_proxy)) + Router::new().route("/proxy/*ai", post(global_proxy).get(global_proxy)) } pub fn workspaced_service() -> Router { - Router::new().route("/proxy/*ai", post(proxy)) + Router::new().route("/proxy/*ai", post(proxy).get(proxy)) } async fn global_proxy( authed: ApiAuthed, Extension(db): Extension, Path(ai_path): Path, + method: Method, headers: HeaderMap, body: Bytes, ) -> impl IntoResponse { @@ -374,7 +388,7 @@ async fn global_proxy( let url = format!("{}/{}", base_url, ai_path); let request = HTTP_CLIENT - .post(url) + .request(method, url) .header("content-type", "application/json") .header("Authorization", format!("Bearer {}", api_key)) .body(body); @@ -410,6 +424,7 @@ async fn proxy( authed: ApiAuthed, Extension(db): Extension, Path((w_id, ai_path)): Path<(String, String)>, + method: Method, headers: HeaderMap, body: Bytes, ) -> impl IntoResponse { @@ -492,7 +507,7 @@ async fn proxy( } }; - let request = request_config.prepare_request(&provider, &ai_path, body)?; + let request = request_config.prepare_request(&provider, &ai_path, method, headers, body)?; let response = request.send().await.map_err(to_anyhow)?; diff --git a/frontend/src/lib/components/ResourcePicker.svelte b/frontend/src/lib/components/ResourcePicker.svelte index 6169fc10e7..69e8a4bf9e 100644 --- a/frontend/src/lib/components/ResourcePicker.svelte +++ b/frontend/src/lib/components/ResourcePicker.svelte @@ -10,8 +10,8 @@ import { Pen, Plus, RotateCw } from 'lucide-svelte' import { sendUserToast } from '$lib/toast' import { isDbType } from './apps/components/display/dbtable/utils' - import { createDispatcherIfMounted } from '$lib/createDispatcherIfMounted' import Select from './Select.svelte' + import { createDispatcherIfMounted } from '$lib/createDispatcherIfMounted' const dispatch = createEventDispatcher() const dispatchIfMounted = createDispatcherIfMounted(dispatch) diff --git a/frontend/src/lib/components/copilot/lib.ts b/frontend/src/lib/components/copilot/lib.ts index d525e58c2b..dd7a8784db 100644 --- a/frontend/src/lib/components/copilot/lib.ts +++ b/frontend/src/lib/components/copilot/lib.ts @@ -37,6 +37,62 @@ export const AI_DEFAULT_MODELS: Record = { customai: [] } +export interface ModelResponse { + id: string + object: string + created: number + owned_by: string + lifecycle_status: string + capabilities: { + completion: boolean + chat_completion: boolean + } +} + +export async function fetchAvailableModels( + resourcePath: string, + workspace: string, + provider: AIProvider +): Promise { + const models = await fetch(`${location.origin}${OpenAPI.BASE}/w/${workspace}/ai/proxy/models`, { + headers: { + 'X-Resource-Path': resourcePath, + 'X-Provider': provider, + ...(provider === 'anthropic' ? { 'anthropic-version': '2023-06-01' } : {}) + } + }) + if (!models.ok) { + console.error('Failed to fetch models for provider', provider, models) + throw new Error(`Failed to fetch models for provider ${provider}`) + } + const data = (await models.json()) as { data: ModelResponse[] } + if (data.data.length > 0) { + switch (provider) { + case 'openai': + return data.data + .filter( + (m) => m.id.startsWith('gpt-') || m.id.startsWith('o') || m.id.startsWith('codex') + ) + .map((m) => m.id) + case 'azure_openai': + return data.data + .filter( + (m) => + (m.id.startsWith('gpt-') || m.id.startsWith('o') || m.id.startsWith('codex')) && + m.lifecycle_status !== 'deprecated' && + (m.capabilities.completion || m.capabilities.chat_completion) + ) + .map((m) => m.id) + case 'googleai': + return data.data.map((m) => m.id.split('/')[1]) + default: + return data.data.map((m) => m.id) + } + } + + return data?.data.map((m) => m.id) ?? [] +} + function getModelMaxTokens(model: string) { if (model.startsWith('gpt-4.1')) { return 32768 diff --git a/frontend/src/lib/components/workspaceSettings/AISettings.svelte b/frontend/src/lib/components/workspaceSettings/AISettings.svelte index 47445bb32f..0d9ad0bef8 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 } from '../copilot/lib' + import { AI_DEFAULT_MODELS, fetchAvailableModels } from '../copilot/lib' import TestAiKey from '../copilot/TestAIKey.svelte' import Description from '../Description.svelte' import Label from '../Label.svelte' @@ -37,7 +37,14 @@ usingOpenaiClientCredentialsOauth: boolean } = $props() - let availableAIModels = $derived(Object.values(aiProviders).flatMap((p) => p.models)) + let fetchedAiModels = $state(false) + let availableAiModels = $state( + Object.fromEntries( + aiProviderLabels.map(([provider]) => [provider, AI_DEFAULT_MODELS[provider]]) + ) as Record + ) + + let selectedAiModels = $derived(Object.values(aiProviders).flatMap((p) => p.models)) let modelProviderMap = $derived( Object.fromEntries( Object.entries(aiProviders).flatMap(([provider, config]) => @@ -52,27 +59,48 @@ } }) + $effect(() => { + ;(async () => { + if (fetchedAiModels) { + return + } + for (const provider of Object.keys(aiProviders)) { + try { + const models = await fetchAvailableModels( + aiProviders[provider].resource_path, + $workspaceStore!, + provider as AIProvider + ) + availableAiModels[provider] = models + } catch (e) { + console.error('failed to fetch models for provider', provider, e) + availableAiModels[provider] = AI_DEFAULT_MODELS[provider] + } + } + fetchedAiModels = true + })() + }) + async function editCopilotConfig(): Promise { if (Object.keys(aiProviders ?? {}).length > 0) { - const code_completion_model = codeCompletionModel - ? { model: codeCompletionModel, provider: modelProviderMap[codeCompletionModel] } - : undefined - const default_model = defaultModel - ? { model: defaultModel, provider: modelProviderMap[defaultModel] } - : undefined - await WorkspaceService.editCopilotConfig({ - workspace: $workspaceStore!, - requestBody: { - providers: aiProviders, - code_completion_model, - default_model - } - }) - setCopilotInfo({ + const code_completion_model = + codeCompletionModel && modelProviderMap[codeCompletionModel] + ? { model: codeCompletionModel, provider: modelProviderMap[codeCompletionModel] } + : undefined + const default_model = + defaultModel && modelProviderMap[defaultModel] + ? { model: defaultModel, provider: modelProviderMap[defaultModel] } + : undefined + const config: AIConfig = { providers: aiProviders, code_completion_model, default_model + } + await WorkspaceService.editCopilotConfig({ + workspace: $workspaceStore!, + requestBody: config }) + setCopilotInfo(config) } else { await WorkspaceService.editCopilotConfig({ workspace: $workspaceStore!, @@ -111,12 +139,12 @@ [provider]: { resource_path: '', models: - AI_DEFAULT_MODELS[provider].length > 0 ? [AI_DEFAULT_MODELS[provider][0]] : [] + availableAiModels[provider].length > 0 ? [availableAiModels[provider][0]] : [] } } - if (AI_DEFAULT_MODELS[provider].length > 0 && !defaultModel) { - defaultModel = AI_DEFAULT_MODELS[provider][0] + if (availableAiModels[provider].length > 0 && !defaultModel) { + defaultModel = availableAiModels[provider][0] } } else { aiProviders = Object.fromEntries( @@ -144,26 +172,38 @@ {#if aiProviders[provider]}
- {#key aiProviders[provider].resource_path} - - - { - if ( - aiProviders[provider]?.resource_path && - aiProviders[provider]?.models.length === 0 && - AI_DEFAULT_MODELS[provider].length > 0 - ) { - aiProviders[provider].models = AI_DEFAULT_MODELS[provider].slice(0, 1) + + + { + if (aiProviders[provider].resource_path) { + try { + const models = await fetchAvailableModels( + aiProviders[provider].resource_path, + $workspaceStore!, + provider as AIProvider + ) + availableAiModels[provider] = models + } catch (e) { + console.error('failed to fetch models for provider', provider, e) + availableAiModels[provider] = AI_DEFAULT_MODELS[provider] } - }} - /> - {/key} + } + + if ( + aiProviders[provider]?.resource_path && + aiProviders[provider]?.models.length === 0 && + availableAiModels[provider].length > 0 + ) { + aiProviders[provider].models = availableAiModels[provider].slice(0, 1) + } + }} + /> +

+ If you don't see the model you want, you can type it manually in the selector. +

{/if}
@@ -195,7 +238,7 @@