feat: support upstream model fetch for advanced custom channels (#5971)
* feat: support upstream model fetch for advanced custom channels * fix: add advanced custom routes as separate groups * fix: select advanced custom route entry before adding --------- Co-authored-by: CaIon <i@caion.me>
This commit is contained in:
+135
-33
@@ -16,6 +16,7 @@ import (
|
||||
"github.com/QuantumNous/new-api/model"
|
||||
relaychannel "github.com/QuantumNous/new-api/relay/channel"
|
||||
"github.com/QuantumNous/new-api/relay/channel/ollama"
|
||||
relaycommon "github.com/QuantumNous/new-api/relay/common"
|
||||
"github.com/QuantumNous/new-api/service"
|
||||
"github.com/QuantumNous/new-api/service/authz"
|
||||
|
||||
@@ -201,22 +202,29 @@ func buildFetchModelsHeaders(channel *model.Channel, key string) (http.Header, e
|
||||
headers = GetAuthHeader(key)
|
||||
}
|
||||
|
||||
headerOverride := channel.GetHeaderOverride()
|
||||
for k, v := range headerOverride {
|
||||
if relaychannel.IsHeaderPassthroughRuleKey(k) {
|
||||
continue
|
||||
}
|
||||
str, ok := v.(string)
|
||||
if !ok {
|
||||
return nil, fmt.Errorf("invalid header override for key %s", k)
|
||||
}
|
||||
if strings.Contains(str, "{api_key}") {
|
||||
str = strings.ReplaceAll(str, "{api_key}", key)
|
||||
}
|
||||
headers.Set(k, str)
|
||||
if err := applyFetchModelsHeaderOverrides(channel, key, headers); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return headers, nil
|
||||
}
|
||||
|
||||
func applyFetchModelsHeaderOverrides(channel *model.Channel, key string, headers http.Header) error {
|
||||
info := &relaycommon.RelayInfo{
|
||||
IsChannelTest: true,
|
||||
ChannelMeta: &relaycommon.ChannelMeta{
|
||||
ApiKey: key,
|
||||
HeadersOverride: channel.GetHeaderOverride(),
|
||||
},
|
||||
}
|
||||
overrides, err := relaychannel.ResolveHeaderOverride(info, nil)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
for name, value := range overrides {
|
||||
headers.Set(name, value)
|
||||
}
|
||||
|
||||
return headers, nil
|
||||
return nil
|
||||
}
|
||||
|
||||
func FetchUpstreamModels(c *gin.Context) {
|
||||
@@ -464,6 +472,10 @@ func validateTwoFactorAuth(twoFA *model.TwoFA, code string) bool {
|
||||
|
||||
// validateChannel 通用的渠道校验函数
|
||||
func validateChannel(channel *model.Channel, isAdd bool) error {
|
||||
if channel == nil {
|
||||
return fmt.Errorf("channel cannot be empty")
|
||||
}
|
||||
|
||||
// 校验 channel settings
|
||||
if err := channel.ValidateSettings(); err != nil {
|
||||
return fmt.Errorf("渠道额外设置[channel setting] 格式错误:%s", err.Error())
|
||||
@@ -471,7 +483,7 @@ func validateChannel(channel *model.Channel, isAdd bool) error {
|
||||
|
||||
// 如果是添加操作,检查 channel 和 key 是否为空
|
||||
if isAdd {
|
||||
if channel == nil || channel.Key == "" {
|
||||
if channel.Key == "" {
|
||||
return fmt.Errorf("channel cannot be empty")
|
||||
}
|
||||
|
||||
@@ -1155,13 +1167,87 @@ func equalStringPtr(a, b *string) bool {
|
||||
return *a == *b
|
||||
}
|
||||
|
||||
func FetchModels(c *gin.Context) {
|
||||
var req struct {
|
||||
BaseURL string `json:"base_url"`
|
||||
Type int `json:"type"`
|
||||
Key string `json:"key"`
|
||||
type fetchModelsRequest struct {
|
||||
ChannelID int `json:"channel_id"`
|
||||
BaseURL *string `json:"base_url"`
|
||||
Type int `json:"type"`
|
||||
Key string `json:"key"`
|
||||
AdvancedCustom *string `json:"advanced_custom"`
|
||||
HeaderOverride *string `json:"header_override"`
|
||||
Proxy *string `json:"proxy"`
|
||||
}
|
||||
|
||||
func buildAdvancedCustomModelPreviewChannel(req fetchModelsRequest) (*model.Channel, error) {
|
||||
var channel *model.Channel
|
||||
if req.ChannelID > 0 {
|
||||
savedChannel, err := model.GetChannelById(req.ChannelID, true)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if savedChannel.Type != constant.ChannelTypeAdvancedCustom {
|
||||
return nil, fmt.Errorf("channel %d is not an advanced custom channel", req.ChannelID)
|
||||
}
|
||||
channel = savedChannel
|
||||
} else {
|
||||
key := strings.TrimSpace(req.Key)
|
||||
if key != "" {
|
||||
key = strings.Split(key, "\n")[0]
|
||||
}
|
||||
channel = &model.Channel{
|
||||
Type: req.Type,
|
||||
Key: key,
|
||||
}
|
||||
}
|
||||
|
||||
if channel.Type != constant.ChannelTypeAdvancedCustom {
|
||||
return nil, fmt.Errorf("channel type must be advanced custom")
|
||||
}
|
||||
if req.BaseURL != nil {
|
||||
baseURL := strings.TrimSpace(*req.BaseURL)
|
||||
channel.BaseURL = &baseURL
|
||||
}
|
||||
|
||||
settings := channel.GetOtherSettings()
|
||||
if req.AdvancedCustom != nil {
|
||||
rawConfig := strings.TrimSpace(*req.AdvancedCustom)
|
||||
if rawConfig == "" {
|
||||
return nil, fmt.Errorf("advanced_custom is required")
|
||||
}
|
||||
var config dto.AdvancedCustomConfig
|
||||
if err := common.UnmarshalJsonStr(rawConfig, &config); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
settings.AdvancedCustom = &config
|
||||
} else if req.ChannelID <= 0 {
|
||||
return nil, fmt.Errorf("advanced_custom is required")
|
||||
}
|
||||
channel.SetOtherSettings(settings)
|
||||
|
||||
if req.HeaderOverride != nil {
|
||||
rawHeaderOverride := strings.TrimSpace(*req.HeaderOverride)
|
||||
if rawHeaderOverride != "" {
|
||||
var headerOverride map[string]any
|
||||
if err := common.UnmarshalJsonStr(rawHeaderOverride, &headerOverride); err != nil {
|
||||
return nil, fmt.Errorf("header_override must be a JSON object: %w", err)
|
||||
}
|
||||
}
|
||||
channel.HeaderOverride = &rawHeaderOverride
|
||||
}
|
||||
if req.Proxy != nil {
|
||||
channelSettings := channel.GetSetting()
|
||||
channelSettings.Proxy = strings.TrimSpace(*req.Proxy)
|
||||
channel.SetSetting(channelSettings)
|
||||
}
|
||||
|
||||
if err := validateChannel(channel, false); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return channel, nil
|
||||
}
|
||||
|
||||
func FetchModels(c *gin.Context) {
|
||||
var req fetchModelsRequest
|
||||
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
c.JSON(http.StatusBadRequest, gin.H{
|
||||
"success": false,
|
||||
@@ -1170,21 +1256,37 @@ func FetchModels(c *gin.Context) {
|
||||
return
|
||||
}
|
||||
|
||||
baseURL := req.BaseURL
|
||||
if baseURL == "" {
|
||||
baseURL = constant.ChannelBaseURLs[req.Type]
|
||||
var channel *model.Channel
|
||||
if req.Type == constant.ChannelTypeAdvancedCustom || req.ChannelID > 0 {
|
||||
var err error
|
||||
channel, err = buildAdvancedCustomModelPreviewChannel(req)
|
||||
if err != nil {
|
||||
c.JSON(http.StatusOK, gin.H{
|
||||
"success": false,
|
||||
"message": err.Error(),
|
||||
})
|
||||
return
|
||||
}
|
||||
} else {
|
||||
baseURL := ""
|
||||
if req.BaseURL != nil {
|
||||
baseURL = strings.TrimSpace(*req.BaseURL)
|
||||
}
|
||||
if baseURL == "" {
|
||||
baseURL = constant.ChannelBaseURLs[req.Type]
|
||||
}
|
||||
|
||||
key := strings.TrimSpace(req.Key)
|
||||
if req.Type != constant.ChannelTypeCodex {
|
||||
key = strings.Split(key, "\n")[0]
|
||||
}
|
||||
channel = &model.Channel{
|
||||
Type: req.Type,
|
||||
Key: key,
|
||||
BaseURL: &baseURL,
|
||||
}
|
||||
}
|
||||
|
||||
key := strings.TrimSpace(req.Key)
|
||||
if req.Type != constant.ChannelTypeCodex {
|
||||
key = strings.Split(key, "\n")[0]
|
||||
}
|
||||
|
||||
channel := &model.Channel{
|
||||
Type: req.Type,
|
||||
Key: key,
|
||||
BaseURL: &baseURL,
|
||||
}
|
||||
models, err := fetchChannelUpstreamModelIDs(channel)
|
||||
if err != nil {
|
||||
c.JSON(http.StatusOK, gin.H{
|
||||
|
||||
@@ -2,8 +2,11 @@ package controller
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"net/http"
|
||||
"net/url"
|
||||
"regexp"
|
||||
"slices"
|
||||
"strings"
|
||||
@@ -14,9 +17,13 @@ import (
|
||||
"github.com/QuantumNous/new-api/constant"
|
||||
"github.com/QuantumNous/new-api/dto"
|
||||
"github.com/QuantumNous/new-api/model"
|
||||
"github.com/QuantumNous/new-api/relay/channel/advancedcustom"
|
||||
"github.com/QuantumNous/new-api/relay/channel/gemini"
|
||||
"github.com/QuantumNous/new-api/relay/channel/ollama"
|
||||
relaycommon "github.com/QuantumNous/new-api/relay/common"
|
||||
relayconstant "github.com/QuantumNous/new-api/relay/constant"
|
||||
"github.com/QuantumNous/new-api/service"
|
||||
"github.com/QuantumNous/new-api/types"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/samber/lo"
|
||||
@@ -255,6 +262,76 @@ func getUpstreamModelUpdateMinCheckIntervalSeconds() int64 {
|
||||
return interval
|
||||
}
|
||||
|
||||
func parseOpenAIModelIDs(body []byte) ([]string, error) {
|
||||
var result struct {
|
||||
Data *[]OpenAIModel `json:"data"`
|
||||
}
|
||||
if err := common.Unmarshal(body, &result); err != nil {
|
||||
return nil, fmt.Errorf("invalid OpenAI Models response: %w", err)
|
||||
}
|
||||
if result.Data == nil {
|
||||
return nil, fmt.Errorf("invalid OpenAI Models response: data is required")
|
||||
}
|
||||
ids := normalizeModelNames(lo.Map(*result.Data, func(item OpenAIModel, _ int) string {
|
||||
return item.ID
|
||||
}))
|
||||
if len(ids) == 0 {
|
||||
return nil, fmt.Errorf("OpenAI Models response contains no valid model IDs")
|
||||
}
|
||||
return ids, nil
|
||||
}
|
||||
|
||||
func sanitizeFetchModelsError(err error, key string) error {
|
||||
if err == nil {
|
||||
return nil
|
||||
}
|
||||
|
||||
// net/http includes the complete request URL in url.Error. Discovery routes
|
||||
// may put the API key in a custom query name or value, so never return that
|
||||
// wrapper to an API client.
|
||||
var urlErr *url.Error
|
||||
if errors.As(err, &urlErr) && urlErr.Err != nil {
|
||||
err = urlErr.Err
|
||||
}
|
||||
|
||||
message := err.Error()
|
||||
key = strings.TrimSpace(key)
|
||||
if key != "" {
|
||||
message = strings.ReplaceAll(message, key, "[REDACTED]")
|
||||
message = strings.ReplaceAll(message, url.QueryEscape(key), "[REDACTED]")
|
||||
message = strings.ReplaceAll(message, url.PathEscape(key), "[REDACTED]")
|
||||
}
|
||||
return errors.New(message)
|
||||
}
|
||||
|
||||
func getFetchModelsResponseBody(method string, requestURL string, channel *model.Channel, headers http.Header) ([]byte, error) {
|
||||
request, err := http.NewRequest(method, requestURL, nil)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
for name, values := range headers {
|
||||
for _, value := range values {
|
||||
request.Header.Add(name, value)
|
||||
}
|
||||
if strings.EqualFold(name, "Host") {
|
||||
request.Host = headers.Get(name)
|
||||
}
|
||||
}
|
||||
client, err := service.NewProxyHttpClient(channel.GetSetting().Proxy)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
response, err := client.Do(request)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer response.Body.Close()
|
||||
if response.StatusCode != http.StatusOK {
|
||||
return nil, fmt.Errorf("status code: %d", response.StatusCode)
|
||||
}
|
||||
return io.ReadAll(response.Body)
|
||||
}
|
||||
|
||||
func fetchChannelUpstreamModelIDs(channel *model.Channel) ([]string, error) {
|
||||
baseURL := constant.ChannelBaseURLs[channel.Type]
|
||||
if channel.GetBaseURL() != "" {
|
||||
@@ -285,6 +362,10 @@ func fetchChannelUpstreamModelIDs(channel *model.Channel) ([]string, error) {
|
||||
return normalizeModelNames(models), nil
|
||||
}
|
||||
|
||||
if channel.Type == constant.ChannelTypeAdvancedCustom {
|
||||
return fetchAdvancedCustomUpstreamModelIDs(channel, baseURL)
|
||||
}
|
||||
|
||||
if channel.Type == constant.ChannelTypeCodex {
|
||||
return service.FetchCodexChannelModels(channel)
|
||||
}
|
||||
@@ -323,29 +404,62 @@ func fetchChannelUpstreamModelIDs(channel *model.Channel) ([]string, error) {
|
||||
|
||||
headers, err := buildFetchModelsHeaders(channel, key)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
return nil, sanitizeFetchModelsError(err, key)
|
||||
}
|
||||
|
||||
body, err := GetResponseBody(http.MethodGet, url, channel, headers)
|
||||
body, err := getFetchModelsResponseBody(http.MethodGet, url, channel, headers)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
return nil, sanitizeFetchModelsError(err, key)
|
||||
}
|
||||
|
||||
var result OpenAIModelsResponse
|
||||
if err := common.Unmarshal(body, &result); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
ids := lo.Map(result.Data, func(item OpenAIModel, _ int) string {
|
||||
if channel.Type == constant.ChannelTypeGemini {
|
||||
return strings.TrimPrefix(item.ID, "models/")
|
||||
}
|
||||
return item.ID
|
||||
})
|
||||
|
||||
return normalizeModelNames(ids), nil
|
||||
}
|
||||
|
||||
func fetchAdvancedCustomUpstreamModelIDs(channel *model.Channel, baseURL string) ([]string, error) {
|
||||
key, _, apiErr := channel.GetNextEnabledKey()
|
||||
if apiErr != nil {
|
||||
return nil, fmt.Errorf("获取渠道密钥失败: %w", apiErr)
|
||||
}
|
||||
key = strings.TrimSpace(key)
|
||||
|
||||
info := &relaycommon.RelayInfo{
|
||||
RelayFormat: types.RelayFormatOpenAI,
|
||||
RelayMode: relayconstant.RelayModeUnknown,
|
||||
RequestURLPath: dto.AdvancedCustomModelListPath,
|
||||
ChannelMeta: &relaycommon.ChannelMeta{
|
||||
ChannelType: constant.ChannelTypeAdvancedCustom,
|
||||
ChannelBaseUrl: baseURL,
|
||||
ApiKey: key,
|
||||
ChannelOtherSettings: channel.GetOtherSettings(),
|
||||
},
|
||||
}
|
||||
|
||||
adaptor := &advancedcustom.Adaptor{}
|
||||
url, headers, err := adaptor.BuildModelListRequest(info)
|
||||
if err != nil {
|
||||
return nil, sanitizeFetchModelsError(err, key)
|
||||
}
|
||||
if err := applyFetchModelsHeaderOverrides(channel, key, headers); err != nil {
|
||||
return nil, sanitizeFetchModelsError(err, key)
|
||||
}
|
||||
|
||||
body, err := getFetchModelsResponseBody(http.MethodGet, url, channel, headers)
|
||||
if err != nil {
|
||||
return nil, sanitizeFetchModelsError(err, key)
|
||||
}
|
||||
return parseOpenAIModelIDs(body)
|
||||
}
|
||||
|
||||
func updateChannelUpstreamModelSettings(channel *model.Channel, settings dto.ChannelOtherSettings, updateModels bool) error {
|
||||
channel.SetOtherSettings(settings)
|
||||
updates := map[string]interface{}{
|
||||
|
||||
@@ -1,9 +1,11 @@
|
||||
package controller
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"errors"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"strings"
|
||||
"net/url"
|
||||
"testing"
|
||||
|
||||
"github.com/QuantumNous/new-api/common"
|
||||
@@ -14,6 +16,340 @@ import (
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
func newAdvancedCustomModelListChannel(baseURL string, key string, upstreamPath string, auth *dto.AdvancedCustomRouteAuth) *model.Channel {
|
||||
config := &dto.AdvancedCustomConfig{
|
||||
Routes: []dto.AdvancedCustomRoute{
|
||||
{
|
||||
IncomingPath: dto.AdvancedCustomModelListPath,
|
||||
UpstreamPath: upstreamPath,
|
||||
Converter: "none",
|
||||
Auth: auth,
|
||||
},
|
||||
},
|
||||
}
|
||||
channel := &model.Channel{
|
||||
Type: constant.ChannelTypeAdvancedCustom,
|
||||
Key: key,
|
||||
BaseURL: &baseURL,
|
||||
}
|
||||
channel.SetOtherSettings(dto.ChannelOtherSettings{AdvancedCustom: config})
|
||||
return channel
|
||||
}
|
||||
|
||||
func TestParseOpenAIModelIDsStrictResponseContract(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
body string
|
||||
want []string
|
||||
wantError string
|
||||
}{
|
||||
{name: "malformed JSON", body: `{"data":`, wantError: "invalid OpenAI Models response"},
|
||||
{name: "missing data", body: `{"object":"list"}`, wantError: "data is required"},
|
||||
{name: "null data", body: `{"data":null}`, wantError: "data is required"},
|
||||
{name: "empty data", body: `{"data":[]}`, wantError: "no valid model IDs"},
|
||||
{name: "all IDs empty", body: `{"data":[{"id":""},{"id":" "}]}`, wantError: "no valid model IDs"},
|
||||
{
|
||||
name: "filters empty IDs and normalizes valid IDs",
|
||||
body: `{"data":[{"id":" gpt-4.1 "},{"id":""},{"id":"gpt-4.1"},{"id":"o3"}]}`,
|
||||
want: []string{"gpt-4.1", "o3"},
|
||||
},
|
||||
}
|
||||
|
||||
for _, test := range tests {
|
||||
t.Run(test.name, func(t *testing.T) {
|
||||
models, err := parseOpenAIModelIDs([]byte(test.body))
|
||||
if test.wantError != "" {
|
||||
require.ErrorContains(t, err, test.wantError)
|
||||
require.Nil(t, models)
|
||||
return
|
||||
}
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, test.want, models)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestFetchAdvancedCustomModelsAppliesHeaderOverrideAfterRouteAuth(t *testing.T) {
|
||||
type receivedRequest struct {
|
||||
Headers http.Header
|
||||
Host string
|
||||
}
|
||||
received := make(chan receivedRequest, 1)
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
received <- receivedRequest{Headers: r.Header.Clone(), Host: r.Host}
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
_, _ = w.Write([]byte(`{"data":[{"id":"gpt-4.1"}]}`))
|
||||
}))
|
||||
defer server.Close()
|
||||
|
||||
channel := newAdvancedCustomModelListChannel(server.URL, "secret-key", "/provider/models", &dto.AdvancedCustomRouteAuth{
|
||||
Type: dto.AdvancedCustomAuthTypeHeader,
|
||||
Name: "X-Route-Key",
|
||||
Value: "route-{api_key}",
|
||||
})
|
||||
headerOverride := `{
|
||||
"X-Route-Key":"global-{api_key}",
|
||||
"X-Static":"static-value",
|
||||
"X-Client":"{client_header:X-Client}",
|
||||
"Host":"models.example.test",
|
||||
"*":""
|
||||
}`
|
||||
channel.HeaderOverride = &headerOverride
|
||||
|
||||
models, err := fetchChannelUpstreamModelIDs(channel)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, []string{"gpt-4.1"}, models)
|
||||
|
||||
request := <-received
|
||||
require.Equal(t, "global-secret-key", request.Headers.Get("X-Route-Key"))
|
||||
require.Equal(t, "static-value", request.Headers.Get("X-Static"))
|
||||
require.Empty(t, request.Headers.Get("X-Client"))
|
||||
require.Equal(t, "models.example.test", request.Host)
|
||||
}
|
||||
|
||||
func TestFetchAdvancedCustomModelsUsesEnabledSavedMultiKey(t *testing.T) {
|
||||
authorization := make(chan string, 1)
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
authorization <- r.Header.Get("Authorization")
|
||||
_, _ = w.Write([]byte(`{"data":[{"id":"gpt-4.1-mini"}]}`))
|
||||
}))
|
||||
defer server.Close()
|
||||
|
||||
channel := newAdvancedCustomModelListChannel(server.URL, "disabled-key\nenabled-key", "/v1/models", nil)
|
||||
channel.ChannelInfo = model.ChannelInfo{
|
||||
IsMultiKey: true,
|
||||
MultiKeyStatusList: map[int]int{
|
||||
0: common.ChannelStatusManuallyDisabled,
|
||||
1: common.ChannelStatusEnabled,
|
||||
},
|
||||
}
|
||||
|
||||
models, err := fetchChannelUpstreamModelIDs(channel)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, []string{"gpt-4.1-mini"}, models)
|
||||
require.Equal(t, "Bearer enabled-key", <-authorization)
|
||||
}
|
||||
|
||||
func TestFetchAdvancedCustomModelsRejectsNonOKResponse(t *testing.T) {
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
|
||||
w.WriteHeader(http.StatusBadGateway)
|
||||
_, _ = w.Write([]byte(`{"data":[{"id":"must-not-be-used"}]}`))
|
||||
}))
|
||||
defer server.Close()
|
||||
|
||||
channel := newAdvancedCustomModelListChannel(server.URL, "secret-key", "/v1/models", nil)
|
||||
models, err := fetchChannelUpstreamModelIDs(channel)
|
||||
require.ErrorContains(t, err, "status code: 502")
|
||||
require.Nil(t, models)
|
||||
}
|
||||
|
||||
func TestFetchAdvancedCustomModelsRedactsQueryKeyFromTransportErrors(t *testing.T) {
|
||||
const secret = "secret key/+"
|
||||
server := httptest.NewServer(http.HandlerFunc(func(http.ResponseWriter, *http.Request) {}))
|
||||
baseURL := server.URL
|
||||
server.Close()
|
||||
|
||||
channel := newAdvancedCustomModelListChannel(baseURL, secret, "/v1/models", &dto.AdvancedCustomRouteAuth{
|
||||
Type: dto.AdvancedCustomAuthTypeQuery,
|
||||
Name: "custom-token",
|
||||
Value: "prefix-{api_key}",
|
||||
})
|
||||
|
||||
_, err := fetchChannelUpstreamModelIDs(channel)
|
||||
require.Error(t, err)
|
||||
require.NotContains(t, err.Error(), secret)
|
||||
require.NotContains(t, err.Error(), "custom-token")
|
||||
require.NotContains(t, err.Error(), "prefix-")
|
||||
|
||||
direct := sanitizeFetchModelsError(&url.Error{
|
||||
Op: http.MethodGet,
|
||||
URL: baseURL + "/v1/models?custom-token=prefix-" + url.QueryEscape(secret),
|
||||
Err: errors.New("connection refused"),
|
||||
}, secret)
|
||||
require.EqualError(t, direct, "connection refused")
|
||||
}
|
||||
|
||||
func TestFetchOrdinaryOpenAIModelsKeepsExistingEmptyDataBehavior(t *testing.T) {
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
|
||||
_, _ = w.Write([]byte(`{"object":"list"}`))
|
||||
}))
|
||||
defer server.Close()
|
||||
|
||||
baseURL := server.URL
|
||||
channel := &model.Channel{
|
||||
Type: constant.ChannelTypeOpenAI,
|
||||
Key: "ordinary-key",
|
||||
BaseURL: &baseURL,
|
||||
}
|
||||
models, err := fetchChannelUpstreamModelIDs(channel)
|
||||
require.NoError(t, err)
|
||||
require.Empty(t, models)
|
||||
}
|
||||
|
||||
func TestFetchModelsAdvancedCustomCreatePreview(t *testing.T) {
|
||||
receivedAuthorization := make(chan string, 1)
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
receivedAuthorization <- r.Header.Get("Authorization")
|
||||
_, _ = w.Write([]byte(`{"data":[{"id":"preview-model"}]}`))
|
||||
}))
|
||||
defer server.Close()
|
||||
|
||||
config := dto.AdvancedCustomConfig{Routes: []dto.AdvancedCustomRoute{{
|
||||
IncomingPath: dto.AdvancedCustomModelListPath,
|
||||
UpstreamPath: "/preview/models",
|
||||
Converter: "none",
|
||||
}}}
|
||||
configBytes, err := common.Marshal(config)
|
||||
require.NoError(t, err)
|
||||
rawConfig := string(configBytes)
|
||||
baseURL := server.URL
|
||||
emptyProxy := ""
|
||||
req := fetchModelsRequest{
|
||||
BaseURL: &baseURL,
|
||||
Type: constant.ChannelTypeAdvancedCustom,
|
||||
Key: "create-preview-key",
|
||||
AdvancedCustom: &rawConfig,
|
||||
Proxy: &emptyProxy,
|
||||
}
|
||||
body, err := common.Marshal(req)
|
||||
require.NoError(t, err)
|
||||
|
||||
recorder := httptest.NewRecorder()
|
||||
ctx, _ := gin.CreateTestContext(recorder)
|
||||
ctx.Request = httptest.NewRequest(http.MethodPost, "/api/channel/fetch_models", bytes.NewReader(body))
|
||||
ctx.Request.Header.Set("Content-Type", "application/json")
|
||||
FetchModels(ctx)
|
||||
|
||||
var response struct {
|
||||
Success bool `json:"success"`
|
||||
Message string `json:"message"`
|
||||
Data []string `json:"data"`
|
||||
}
|
||||
require.NoError(t, common.Unmarshal(recorder.Body.Bytes(), &response))
|
||||
require.True(t, response.Success, response.Message)
|
||||
require.Equal(t, []string{"preview-model"}, response.Data)
|
||||
require.Equal(t, "Bearer create-preview-key", <-receivedAuthorization)
|
||||
}
|
||||
|
||||
func TestFetchModelsAdvancedCustomEditPreviewUsesSavedKeyAndExplicitClears(t *testing.T) {
|
||||
db := setupModelListControllerTestDB(t)
|
||||
receivedHeaders := make(chan http.Header, 1)
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
receivedHeaders <- r.Header.Clone()
|
||||
_, _ = w.Write([]byte(`{"data":[{"id":"edited-preview-model"}]}`))
|
||||
}))
|
||||
defer server.Close()
|
||||
|
||||
savedChannel := newAdvancedCustomModelListChannel("http://127.0.0.1:1", "disabled-saved-key\nenabled-saved-key", "/saved/models", nil)
|
||||
savedChannel.Name = "saved advanced channel"
|
||||
savedChannel.Models = "old-model"
|
||||
savedChannel.ChannelInfo = model.ChannelInfo{
|
||||
IsMultiKey: true,
|
||||
MultiKeyStatusList: map[int]int{
|
||||
0: common.ChannelStatusManuallyDisabled,
|
||||
1: common.ChannelStatusEnabled,
|
||||
},
|
||||
}
|
||||
savedHeaderOverride := `{"X-Saved":"must-not-be-sent"}`
|
||||
savedChannel.HeaderOverride = &savedHeaderOverride
|
||||
savedChannel.SetSetting(dto.ChannelSettings{Proxy: "http://127.0.0.1:1"})
|
||||
require.NoError(t, db.Create(savedChannel).Error)
|
||||
|
||||
preserved, err := buildAdvancedCustomModelPreviewChannel(fetchModelsRequest{ChannelID: savedChannel.Id})
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, "http://127.0.0.1:1", preserved.GetBaseURL())
|
||||
require.Equal(t, savedHeaderOverride, *preserved.HeaderOverride)
|
||||
require.Equal(t, "http://127.0.0.1:1", preserved.GetSetting().Proxy)
|
||||
|
||||
previewConfig := dto.AdvancedCustomConfig{Routes: []dto.AdvancedCustomRoute{{
|
||||
IncomingPath: dto.AdvancedCustomModelListPath,
|
||||
UpstreamPath: "/edited/models",
|
||||
Converter: "none",
|
||||
}}}
|
||||
configBytes, err := common.Marshal(previewConfig)
|
||||
require.NoError(t, err)
|
||||
rawConfig := string(configBytes)
|
||||
baseURL := server.URL
|
||||
explicitEmpty := ""
|
||||
req := fetchModelsRequest{
|
||||
ChannelID: savedChannel.Id,
|
||||
BaseURL: &baseURL,
|
||||
Type: constant.ChannelTypeAdvancedCustom,
|
||||
Key: "request-key-must-be-ignored",
|
||||
AdvancedCustom: &rawConfig,
|
||||
HeaderOverride: &explicitEmpty,
|
||||
Proxy: &explicitEmpty,
|
||||
}
|
||||
cleared, err := buildAdvancedCustomModelPreviewChannel(fetchModelsRequest{
|
||||
ChannelID: savedChannel.Id,
|
||||
BaseURL: &explicitEmpty,
|
||||
AdvancedCustom: &rawConfig,
|
||||
HeaderOverride: &explicitEmpty,
|
||||
Proxy: &explicitEmpty,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, cleared.BaseURL)
|
||||
require.Empty(t, *cleared.BaseURL)
|
||||
require.NotNil(t, cleared.HeaderOverride)
|
||||
require.Empty(t, *cleared.HeaderOverride)
|
||||
require.Empty(t, cleared.GetSetting().Proxy)
|
||||
|
||||
body, err := common.Marshal(req)
|
||||
require.NoError(t, err)
|
||||
|
||||
recorder := httptest.NewRecorder()
|
||||
ctx, _ := gin.CreateTestContext(recorder)
|
||||
ctx.Request = httptest.NewRequest(http.MethodPost, "/api/channel/fetch_models", bytes.NewReader(body))
|
||||
ctx.Request.Header.Set("Content-Type", "application/json")
|
||||
FetchModels(ctx)
|
||||
|
||||
var response struct {
|
||||
Success bool `json:"success"`
|
||||
Message string `json:"message"`
|
||||
Data []string `json:"data"`
|
||||
}
|
||||
require.NoError(t, common.Unmarshal(recorder.Body.Bytes(), &response))
|
||||
require.True(t, response.Success, response.Message)
|
||||
require.Equal(t, []string{"edited-preview-model"}, response.Data)
|
||||
require.NotContains(t, recorder.Body.String(), "enabled-saved-key")
|
||||
require.NotContains(t, recorder.Body.String(), "request-key-must-be-ignored")
|
||||
|
||||
headers := <-receivedHeaders
|
||||
require.Equal(t, "Bearer enabled-saved-key", headers.Get("Authorization"))
|
||||
require.Empty(t, headers.Get("X-Saved"))
|
||||
}
|
||||
|
||||
func TestFailedAdvancedCustomDetectionDoesNotStageFullRemoval(t *testing.T) {
|
||||
db := setupModelListControllerTestDB(t)
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
|
||||
_, _ = w.Write([]byte(`{"data":[]}`))
|
||||
}))
|
||||
defer server.Close()
|
||||
|
||||
channel := newAdvancedCustomModelListChannel(server.URL, "secret-key", "/v1/models", nil)
|
||||
channel.Name = "empty discovery response"
|
||||
channel.Models = "gpt-4.1,o3"
|
||||
settings := channel.GetOtherSettings()
|
||||
settings.UpstreamModelUpdateCheckEnabled = true
|
||||
settings.UpstreamModelUpdateAutoSyncEnabled = true
|
||||
channel.SetOtherSettings(settings)
|
||||
require.NoError(t, db.Create(channel).Error)
|
||||
|
||||
modelsChanged, autoAdded, err := checkAndPersistChannelUpstreamModelUpdates(channel, &settings, true, true)
|
||||
require.ErrorContains(t, err, "no valid model IDs")
|
||||
require.False(t, modelsChanged)
|
||||
require.Zero(t, autoAdded)
|
||||
require.Empty(t, settings.UpstreamModelUpdateLastDetectedModels)
|
||||
require.Empty(t, settings.UpstreamModelUpdateLastRemovedModels)
|
||||
|
||||
reloaded, err := model.GetChannelById(channel.Id, true)
|
||||
require.NoError(t, err)
|
||||
persistedSettings := reloaded.GetOtherSettings()
|
||||
require.Empty(t, persistedSettings.UpstreamModelUpdateLastDetectedModels)
|
||||
require.Empty(t, persistedSettings.UpstreamModelUpdateLastRemovedModels)
|
||||
require.Equal(t, "gpt-4.1,o3", reloaded.Models)
|
||||
}
|
||||
|
||||
func TestFetchModelsUsesSharedChannelFetchBehavior(t *testing.T) {
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
if r.URL.Path != "/v1/models" {
|
||||
@@ -39,7 +375,7 @@ func TestFetchModelsUsesSharedChannelFetchBehavior(t *testing.T) {
|
||||
|
||||
recorder := httptest.NewRecorder()
|
||||
ctx, _ := gin.CreateTestContext(recorder)
|
||||
ctx.Request = httptest.NewRequest(http.MethodPost, "/api/channel/fetch_models", strings.NewReader(string(body)))
|
||||
ctx.Request = httptest.NewRequest(http.MethodPost, "/api/channel/fetch_models", bytes.NewReader(body))
|
||||
ctx.Request.Header.Set("Content-Type", "application/json")
|
||||
|
||||
FetchModels(ctx)
|
||||
|
||||
Reference in New Issue
Block a user