mirror of
https://github.com/windmill-labs/windmill.git
synced 2026-08-18 16:02:10 +00:00
feat: use provider api to list available AI models in workspace settings (#5947)
* use open router of model lists * draft * allow get in ai proxy * add fetch available models function * use func * fix for anthropic * fix * fetch on mount * fix ai settings * fix * handle azure
This commit is contained in:
@@ -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<RequestBuilder> {
|
||||
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<DB>,
|
||||
Path(ai_path): Path<String>,
|
||||
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<DB>,
|
||||
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)?;
|
||||
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -37,6 +37,62 @@ export const AI_DEFAULT_MODELS: Record<AIProvider, string[]> = {
|
||||
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<string[]> {
|
||||
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
|
||||
|
||||
@@ -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<AIProvider, string[]>
|
||||
)
|
||||
|
||||
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<void> {
|
||||
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]}
|
||||
<div class="mb-4 flex flex-col gap-2">
|
||||
<div class="flex flex-row gap-1">
|
||||
{#key aiProviders[provider].resource_path}
|
||||
<!-- this can be removed once the parent component moves to runes -->
|
||||
<!-- svelte-ignore binding_property_non_reactive -->
|
||||
<ResourcePicker
|
||||
resourceType={provider === 'openai' && usingOpenaiClientCredentialsOauth
|
||||
? 'openai_client_credentials_oauth'
|
||||
: provider}
|
||||
initialValue={aiProviders[provider].resource_path}
|
||||
bind:value={aiProviders[provider].resource_path}
|
||||
on:change={() => {
|
||||
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)
|
||||
<!-- this can be removed once the parent component moves to runes -->
|
||||
<!-- svelte-ignore binding_property_non_reactive -->
|
||||
<ResourcePicker
|
||||
resourceType={provider === 'openai' && usingOpenaiClientCredentialsOauth
|
||||
? 'openai_client_credentials_oauth'
|
||||
: provider}
|
||||
initialValue={aiProviders[provider].resource_path}
|
||||
bind:value={aiProviders[provider].resource_path}
|
||||
on:change={async () => {
|
||||
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)
|
||||
}
|
||||
}}
|
||||
/>
|
||||
<TestAiKey
|
||||
aiProvider={provider}
|
||||
resourcePath={aiProviders[provider].resource_path}
|
||||
@@ -175,12 +215,15 @@
|
||||
<!-- this can be removed once the parent component moves to runes -->
|
||||
<!-- svelte-ignore binding_property_non_reactive -->
|
||||
<MultiSelectWrapper
|
||||
items={AI_DEFAULT_MODELS[provider]}
|
||||
items={availableAiModels[provider]}
|
||||
bind:value={aiProviders[provider].models}
|
||||
placeholder="Select models"
|
||||
allowUserOptions="append"
|
||||
/>
|
||||
</Label>
|
||||
<p class="text-xs">
|
||||
If you don't see the model you want, you can type it manually in the selector.
|
||||
</p>
|
||||
</div>
|
||||
{/if}
|
||||
</div>
|
||||
@@ -195,7 +238,7 @@
|
||||
<Label label="Default chat model">
|
||||
{#key Object.keys(aiProviders).length}
|
||||
<ArgEnum
|
||||
enum_={availableAIModels}
|
||||
enum_={selectedAiModels}
|
||||
bind:value={defaultModel}
|
||||
disabled={false}
|
||||
autofocus={false}
|
||||
@@ -224,7 +267,7 @@
|
||||
{#if codeCompletionModel != undefined}
|
||||
<Label label="Code completion model">
|
||||
<ArgEnum
|
||||
enum_={availableAIModels}
|
||||
enum_={selectedAiModels}
|
||||
bind:value={codeCompletionModel}
|
||||
disabled={false}
|
||||
autofocus={false}
|
||||
@@ -232,7 +275,7 @@
|
||||
valid={true}
|
||||
create={false}
|
||||
/>
|
||||
<p class="text-xs">
|
||||
<p class="text-xs mt-2">
|
||||
We highly recommend using Mistral's Codestral model for code completion.
|
||||
</p>
|
||||
</Label>
|
||||
|
||||
Reference in New Issue
Block a user