From 3d5dc36f1d85ccae8d5cb2864764011795b559b5 Mon Sep 17 00:00:00 2001 From: zhibisora <73344387+zhibisora@users.noreply.github.com> Date: Tue, 11 Aug 2026 14:49:20 +0800 Subject: [PATCH] =?UTF-8?q?fix:=20=E4=BF=AE=E5=A4=8D=20Gemini=20=E9=A3=8E?= =?UTF-8?q?=E6=A0=BC=20/v1/models=20=E5=88=97=E8=A1=A8=E8=AF=B7=E6=B1=82?= =?UTF-8?q?=20(#6199)?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit * fix: support Gemini model listing on v1 route * test: cover Gemini query key model listing --- middleware/auth.go | 3 +- router/relay-router.go | 2 +- router/relay_router_test.go | 126 ++++++++++++++++++++++++++++++++++++ 3 files changed, 129 insertions(+), 2 deletions(-) create mode 100644 router/relay_router_test.go diff --git a/middleware/auth.go b/middleware/auth.go index 4e1436f3..9f2a9df9 100644 --- a/middleware/auth.go +++ b/middleware/auth.go @@ -374,7 +374,8 @@ func TokenAuth() func(c *gin.Context) { } } // gemini api 从query中获取key - if strings.HasPrefix(c.Request.URL.Path, "/v1beta/models") || + if c.Request.URL.Path == "/v1/models" || + strings.HasPrefix(c.Request.URL.Path, "/v1beta/models") || strings.HasPrefix(c.Request.URL.Path, "/v1beta/openai/models") || strings.HasPrefix(c.Request.URL.Path, "/v1/models/") { skKey := c.Query("key") diff --git a/router/relay-router.go b/router/relay-router.go index e08ecb14..b230a5a8 100644 --- a/router/relay-router.go +++ b/router/relay-router.go @@ -25,7 +25,7 @@ func SetRelayRouter(router *gin.Engine) { case c.GetHeader("x-api-key") != "" && c.GetHeader("anthropic-version") != "": controller.ListModels(c, constant.ChannelTypeAnthropic) case c.GetHeader("x-goog-api-key") != "" || c.Query("key") != "": // 单独的适配 - controller.RetrieveModel(c, constant.ChannelTypeGemini) + controller.ListModels(c, constant.ChannelTypeGemini) default: controller.ListModels(c, constant.ChannelTypeOpenAI) } diff --git a/router/relay_router_test.go b/router/relay_router_test.go new file mode 100644 index 00000000..579bd7bf --- /dev/null +++ b/router/relay_router_test.go @@ -0,0 +1,126 @@ +package router + +import ( + "fmt" + "net/http" + "net/http/httptest" + "os" + "strings" + "testing" + + "github.com/QuantumNous/new-api/common" + "github.com/QuantumNous/new-api/model" + "github.com/gin-gonic/gin" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func TestListModelsSupportsOpenAIAndGeminiAuthentication(t *testing.T) { + setupRelayRouterTestDB(t) + + user := model.User{ + Username: "models-user", + Status: common.UserStatusEnabled, + Group: "default", + Quota: 100, + } + require.NoError(t, model.DB.Create(&user).Error) + require.NoError(t, model.DB.Create(&model.Token{ + UserId: user.Id, + Key: "modelstestkey", + Status: common.TokenStatusEnabled, + ExpiredTime: -1, + UnlimitedQuota: true, + }).Error) + + engine := gin.New() + SetRelayRouter(engine) + + tests := []struct { + name string + path string + headerName string + expectedObject string + expectedField string + }{ + { + name: "OpenAI bearer token", + path: "/v1/models", + headerName: "Authorization", + expectedObject: "list", + expectedField: "data", + }, + { + name: "Gemini API key header", + path: "/v1/models", + headerName: "x-goog-api-key", + expectedField: "models", + }, + { + name: "Gemini API key query", + path: "/v1/models?key=modelstestkey", + expectedField: "models", + }, + } + + for _, test := range tests { + t.Run(test.name, func(t *testing.T) { + recorder := httptest.NewRecorder() + request := httptest.NewRequest(http.MethodGet, test.path, nil) + if test.headerName != "" { + value := "modelstestkey" + if test.headerName == "Authorization" { + value = "Bearer " + value + } + request.Header.Set(test.headerName, value) + } + + engine.ServeHTTP(recorder, request) + + require.Equal(t, http.StatusOK, recorder.Code) + var payload map[string]any + require.NoError(t, common.Unmarshal(recorder.Body.Bytes(), &payload)) + assert.Contains(t, payload, test.expectedField) + assert.NotContains(t, payload, "error") + if test.expectedObject != "" { + assert.Equal(t, test.expectedObject, payload["object"]) + } + }) + } +} + +func setupRelayRouterTestDB(t *testing.T) { + t.Helper() + + gin.SetMode(gin.TestMode) + originalIsMasterNode := common.IsMasterNode + originalRedisEnabled := common.RedisEnabled + originalSQLitePath := common.SQLitePath + originalMainDatabaseType := common.MainDatabaseType() + originalLogDatabaseType := common.LogDatabaseType() + originalSQLDSN, hadSQLDSN := os.LookupEnv("SQL_DSN") + + common.IsMasterNode = false + common.RedisEnabled = false + common.SQLitePath = fmt.Sprintf("file:%s?mode=memory&cache=shared", strings.ReplaceAll(t.Name(), "/", "_")) + common.SetDatabaseTypes(common.DatabaseTypeSQLite, common.DatabaseTypeSQLite) + require.NoError(t, os.Setenv("SQL_DSN", "local")) + require.NoError(t, model.InitDB()) + model.LOG_DB = model.DB + require.NoError(t, model.DB.AutoMigrate(&model.User{}, &model.Token{}, &model.Ability{})) + + t.Cleanup(func() { + if sqlDB, err := model.DB.DB(); err == nil { + _ = sqlDB.Close() + } + common.IsMasterNode = originalIsMasterNode + common.RedisEnabled = originalRedisEnabled + common.SQLitePath = originalSQLitePath + common.SetDatabaseTypes(originalMainDatabaseType, originalLogDatabaseType) + if hadSQLDSN { + require.NoError(t, os.Setenv("SQL_DSN", originalSQLDSN)) + } else { + require.NoError(t, os.Unsetenv("SQL_DSN")) + } + }) +}