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:
centdix
2025-06-17 10:48:47 +02:00
committed by GitHub
parent ad2de83354
commit 7490e883d7
4 changed files with 169 additions and 55 deletions
+26 -11
View File
@@ -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>