diff --git a/controller/channel.go b/controller/channel.go index c59e492a5a02..ec23cb0c5705 100644 --- a/controller/channel.go +++ b/controller/channel.go @@ -224,6 +224,15 @@ func FetchUpstreamModels(c *gin.Context) { common.ApiError(c, err) return } + + if typeStr := c.Query("type"); typeStr != "" { + if t, err := strconv.Atoi(typeStr); err == nil { + channel.Type = t + } + } + if baseURL := c.Query("base_url"); baseURL != "" { + channel.BaseURL = &baseURL + } ids, err := fetchChannelUpstreamModelIDs(channel) if err != nil { diff --git a/web/default/src/features/channels/api.ts b/web/default/src/features/channels/api.ts index 6e92519f98b5..f94cb1b5046c 100644 --- a/web/default/src/features/channels/api.ts +++ b/web/default/src/features/channels/api.ts @@ -214,11 +214,15 @@ export async function updateChannelBalance( * Fetch available models from upstream provider */ export async function fetchUpstreamModels( - id: number + id: number, + overrides?: { type?: number; base_url?: string } ): Promise { + const params: Record = {} + if (overrides?.type != null) params.type = String(overrides.type) + if (overrides?.base_url) params.base_url = overrides.base_url const res = await api.get( `/api/channel/fetch_models/${id}`, - channelActionConfig() + channelActionConfig({ params: Object.keys(params).length > 0 ? params : undefined }) ) return res.data } diff --git a/web/default/src/features/channels/components/drawers/channel-mutate-drawer.tsx b/web/default/src/features/channels/components/drawers/channel-mutate-drawer.tsx index 6b26cd171505..57c71c398994 100644 --- a/web/default/src/features/channels/components/drawers/channel-mutate-drawer.tsx +++ b/web/default/src/features/channels/components/drawers/channel-mutate-drawer.tsx @@ -107,6 +107,7 @@ import { } from '@/features/auth/secure-verification' import { fetchModels, + fetchUpstreamModels, getAllModels, getChannel, getChannelKey, @@ -287,6 +288,9 @@ export function ChannelMutateDrawer({ const initialModelsRef = useRef([]) const initialModelMappingRef = useRef('') const initialStatusCodeMappingRef = useRef('') + const initialTypeRef = useRef(0) + const initialBaseUrlRef = useRef('') + const initialKeyRef = useRef('') const [statusCodeRiskOpen, setStatusCodeRiskOpen] = useState(false) const [statusCodeRiskDetailItems, setStatusCodeRiskDetailItems] = useState< string[] @@ -592,14 +596,20 @@ export function ChannelMutateDrawer({ initialModelMappingRef.current = channelData.data.model_mapping || '' initialStatusCodeMappingRef.current = channelData.data.status_code_mapping || '' + initialTypeRef.current = channelData.data.type ?? 0 + initialBaseUrlRef.current = channelData.data.base_url || '' + initialKeyRef.current = channelKey ?? '' } else if (!isEditing) { form.reset(CHANNEL_FORM_DEFAULT_VALUES) setAdvancedSettingsOpen(false) initialModelsRef.current = [] initialModelMappingRef.current = '' initialStatusCodeMappingRef.current = '' + initialTypeRef.current = 0 + initialBaseUrlRef.current = '' + initialKeyRef.current = '' } - }, [isEditing, channelData, form]) + }, [isEditing, channelData, form, channelKey]) // Handle type change - set default values for specific types useEffect(() => { @@ -769,6 +779,38 @@ export function ChannelMutateDrawer({ throw new Error(response.message || 'No models fetched from upstream') }, [form]) + const editModeFetcher = useCallback(async (): Promise => { + const formKey = form.getValues('key') + if (formKey?.trim()) { + const response = await fetchModels({ + type: form.getValues('type'), + key: formKey, + base_url: form.getValues('base_url') || '', + }) + if (response.success && response.data) { + return response.data + } + throw new Error(response.message || 'No models fetched from upstream') + } + const overrides: { type?: number; base_url?: string } = {} + const currentTypeVal = form.getValues('type') + const currentBaseUrlVal = form.getValues('base_url') || '' + if (currentTypeVal !== initialTypeRef.current) { + overrides.type = currentTypeVal + } + if (currentBaseUrlVal !== initialBaseUrlRef.current) { + overrides.base_url = currentBaseUrlVal + } + const response = await fetchUpstreamModels( + channelId!, + Object.keys(overrides).length > 0 ? overrides : undefined + ) + if (response.success && response.data) { + return response.data + } + throw new Error(response.message || 'No models fetched from upstream') + }, [form, channelId]) + // Handle model operations const handleFillRelatedModels = useCallback(() => { if (!basicModels.length) { @@ -3419,7 +3461,7 @@ export function ChannelMutateDrawer({ }} redirectModels={redirectModelList} redirectSourceModels={redirectModelKeyList} - customFetcher={!isEditing ? createModeFetcher : undefined} + customFetcher={isEditing ? editModeFetcher : createModeFetcher} channelName={!isEditing ? currentName?.trim() : undefined} existingModelsOverride={ !isEditing