Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
1 change: 1 addition & 0 deletions constant/context_key.go
Original file line number Diff line number Diff line change
Expand Up @@ -19,6 +19,7 @@ const (
ContextKeyTokenModelLimitEnabled ContextKey = "token_model_limit_enabled"
ContextKeyTokenModelLimit ContextKey = "token_model_limit"
ContextKeyTokenCrossGroupRetry ContextKey = "token_cross_group_retry"
ContextKeyTokenAutoGroups ContextKey = "token_auto_groups"

/* channel related keys */
ContextKeyChannelId ContextKey = "channel_id"
Expand Down
45 changes: 24 additions & 21 deletions controller/model.go
Original file line number Diff line number Diff line change
Expand Up @@ -20,6 +20,7 @@ import (
"github.com/QuantumNous/new-api/relaykit/types"
"github.com/QuantumNous/new-api/service"
"github.com/QuantumNous/new-api/setting/operation_setting"
"github.com/QuantumNous/new-api/setting/ratio_setting"
"github.com/gin-gonic/gin"
"github.com/samber/lo"
)
Expand Down Expand Up @@ -190,7 +191,7 @@ func getModelListGroups(c *gin.Context) (modelListGroups, error) {
return modelListGroups{
userGroup: userGroup,
tokenGroup: tokenGroup,
ownerGroups: service.GetUserAutoGroup(userGroup),
ownerGroups: service.GetRequestAutoGroups(c, userGroup),
}, nil
}

Expand Down Expand Up @@ -228,32 +229,28 @@ func ListModels(c *gin.Context, modelType int) {
}
ownerGroups := groups.ownerGroups
modelLimitEnable := common.GetContextKeyBool(c, constant.ContextKeyTokenModelLimitEnabled)
var tokenModelLimit map[string]bool
if modelLimitEnable {
s, ok := common.GetContextKey(c, constant.ContextKeyTokenModelLimit)
var tokenModelLimit map[string]bool
if ok {
tokenModelLimit = s.(map[string]bool)
} else {
tokenModelLimit, _ = s.(map[string]bool)
}
if tokenModelLimit == nil {
tokenModelLimit = map[string]bool{}
}
for allowModel, _ := range tokenModelLimit {
if !acceptUnsetRatioModel {
if !helper.HasModelBillingConfig(allowModel) {
continue
}
}
models := service.GetGroupsEnabledModels(ownerGroups)
for _, modelName := range models {
if modelLimitEnable {
matchingName := ratio_setting.FormatMatchingModelName(modelName)
if !tokenModelLimit[modelName] && !tokenModelLimit[matchingName] {
continue
}
userModelNames = append(userModelNames, allowModel)
}
} else {
models := service.GetGroupsEnabledModels(ownerGroups)
for _, modelName := range models {
if !acceptUnsetRatioModel {
if !helper.HasModelBillingConfig(modelName) {
continue
}
}
userModelNames = append(userModelNames, modelName)
if !acceptUnsetRatioModel && !helper.HasModelBillingConfig(modelName) {
continue
}
userModelNames = append(userModelNames, modelName)
}

ownerByModel := map[string]string{}
Expand All @@ -276,11 +273,17 @@ func ListModels(c *gin.Context, modelType int) {
Type: "model",
}
}
firstID := ""
lastID := ""
if len(useranthropicModels) > 0 {
firstID = useranthropicModels[0].ID
lastID = useranthropicModels[len(useranthropicModels)-1].ID
}
c.JSON(200, gin.H{
"data": useranthropicModels,
"first_id": useranthropicModels[0].ID,
"first_id": firstID,
"has_more": false,
"last_id": useranthropicModels[len(useranthropicModels)-1].ID,
"last_id": lastID,
})
case constant.ChannelTypeGemini:
userGeminiModels := make([]dto.GeminiModel, len(userOpenAiModels))
Expand Down
70 changes: 69 additions & 1 deletion controller/model_list_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -402,7 +402,13 @@ func TestListModelsTokenLimitIncludesTieredBillingModel(t *testing.T) {
"zz-token-tiered-visible-model": `tier("base", p * 1 + c * 2)`,
"zz-token-tiered-empty-expr-model": "",
})
setupModelListControllerTestDB(t)
db := setupModelListControllerTestDB(t)
require.NoError(t, db.Create(&[]model.Ability{
{Group: "default", Model: "zz-token-tiered-visible-model", ChannelId: 1, Enabled: true},
{Group: "default", Model: "zz-token-tiered-empty-expr-model", ChannelId: 1, Enabled: true},
{Group: "default", Model: "zz-token-tiered-missing-expr-model", ChannelId: 1, Enabled: true},
{Group: "default", Model: "zz-token-unpriced-model", ChannelId: 1, Enabled: true},
}).Error)

recorder := httptest.NewRecorder()
ctx, _ := gin.CreateTestContext(recorder)
Expand All @@ -425,6 +431,68 @@ func TestListModelsTokenLimitIncludesTieredBillingModel(t *testing.T) {
require.NotContains(t, ids, "zz-token-unpriced-model")
}

func TestListModelsTokenLimitUsesResolvedCustomAutoGroups(t *testing.T) {
withSelfUseModeEnabled(t)
originalMax := setting.GetMaxTokenAutoGroups()
originalUsableGroups := setting.UserUsableGroups2JSONString()
originalRatios := ratio_setting.GroupRatio2JSONString()
require.NoError(t, setting.UpdateMaxTokenAutoGroups("5"))
require.NoError(t, setting.UpdateUserUsableGroupsByJSONString(`{"default":"Default","vip":"VIP"}`))
require.NoError(t, ratio_setting.UpdateGroupRatioByJSONString(`{"default":1,"vip":1}`))
t.Cleanup(func() {
require.NoError(t, setting.UpdateMaxTokenAutoGroups(fmt.Sprintf("%d", originalMax)))
require.NoError(t, setting.UpdateUserUsableGroupsByJSONString(originalUsableGroups))
require.NoError(t, ratio_setting.UpdateGroupRatioByJSONString(originalRatios))
})

db := setupModelListControllerTestDB(t)
require.NoError(t, db.Create(&[]model.Ability{
{Group: "vip", Model: "zz-vip-allowed", ChannelId: 1, Enabled: true},
{Group: "vip", Model: "zz-vip-denied", ChannelId: 1, Enabled: true},
{Group: "default", Model: "zz-default-outside-snapshot", ChannelId: 1, Enabled: true},
}).Error)

recorder := httptest.NewRecorder()
ctx, _ := gin.CreateTestContext(recorder)
ctx.Request = httptest.NewRequest(http.MethodGet, "/v1/models", nil)
common.SetContextKey(ctx, constant.ContextKeyUserGroup, "default")
common.SetContextKey(ctx, constant.ContextKeyTokenGroup, "auto")
common.SetContextKey(ctx, constant.ContextKeyTokenAutoGroups, []string{"vip"})
common.SetContextKey(ctx, constant.ContextKeyTokenModelLimitEnabled, true)
common.SetContextKey(ctx, constant.ContextKeyTokenModelLimit, map[string]bool{
"zz-vip-allowed": true,
"zz-default-outside-snapshot": true,
"zz-not-enabled": true,
})

ListModels(ctx, constant.ChannelTypeOpenAI)
ids := decodeListModelsResponse(t, recorder)
require.Equal(t, map[string]struct{}{"zz-vip-allowed": {}}, ids)

require.NoError(t, setting.UpdateUserUsableGroupsByJSONString(`{"default":"Default"}`))
emptyRecorder := httptest.NewRecorder()
emptyCtx, _ := gin.CreateTestContext(emptyRecorder)
emptyCtx.Request = httptest.NewRequest(http.MethodGet, "/v1/models", nil)
common.SetContextKey(emptyCtx, constant.ContextKeyUserGroup, "default")
common.SetContextKey(emptyCtx, constant.ContextKeyTokenGroup, "auto")
common.SetContextKey(emptyCtx, constant.ContextKeyTokenAutoGroups, []string{"vip"})
common.SetContextKey(emptyCtx, constant.ContextKeyTokenModelLimitEnabled, true)
common.SetContextKey(emptyCtx, constant.ContextKeyTokenModelLimit, map[string]bool{"zz-vip-allowed": true})

require.NotPanics(t, func() {
ListModels(emptyCtx, constant.ChannelTypeAnthropic)
})
var anthropicResponse struct {
Data []dto.AnthropicModel `json:"data"`
FirstID string `json:"first_id"`
LastID string `json:"last_id"`
}
require.NoError(t, common.Unmarshal(emptyRecorder.Body.Bytes(), &anthropicResponse))
require.Empty(t, anthropicResponse.Data)
require.Empty(t, anthropicResponse.FirstID)
require.Empty(t, anthropicResponse.LastID)
}

func TestCheckUpdatePasswordRequiresCurrentPassword(t *testing.T) {
db := setupModelListControllerTestDB(t)
hashedPassword, err := common.Password2Hash("CurrentPassword123")
Expand Down
33 changes: 33 additions & 0 deletions controller/model_owned_by_test.go
Original file line number Diff line number Diff line change
@@ -1,11 +1,14 @@
package controller

import (
"fmt"
"net/http/httptest"
"testing"

"github.com/QuantumNous/new-api/common"
"github.com/QuantumNous/new-api/constant"
"github.com/QuantumNous/new-api/setting"
"github.com/QuantumNous/new-api/setting/ratio_setting"
"github.com/gin-gonic/gin"
"github.com/stretchr/testify/require"
)
Expand Down Expand Up @@ -83,3 +86,33 @@ func TestGetModelListGroupsUsesExplicitTokenGroup(t *testing.T) {
require.Equal(t, "vip", groups.tokenGroup)
require.Equal(t, []string{"vip"}, groups.ownerGroups)
}

func TestGetModelListGroupsUsesFilteredTokenAutoGroupsSnapshot(t *testing.T) {
originalMax := setting.GetMaxTokenAutoGroups()
originalUsableGroups := setting.UserUsableGroups2JSONString()
originalRatios := ratio_setting.GroupRatio2JSONString()
require.NoError(t, setting.UpdateMaxTokenAutoGroups("1"))
require.NoError(t, setting.UpdateUserUsableGroupsByJSONString(`{"default":"Default","vip":"VIP"}`))
require.NoError(t, ratio_setting.UpdateGroupRatioByJSONString(`{"default":1,"vip":1}`))
t.Cleanup(func() {
require.NoError(t, setting.UpdateMaxTokenAutoGroups(fmt.Sprintf("%d", originalMax)))
require.NoError(t, setting.UpdateUserUsableGroupsByJSONString(originalUsableGroups))
require.NoError(t, ratio_setting.UpdateGroupRatioByJSONString(originalRatios))
})

gin.SetMode(gin.TestMode)
ctx, _ := gin.CreateTestContext(httptest.NewRecorder())
common.SetContextKey(ctx, constant.ContextKeyUserGroup, "default")
common.SetContextKey(ctx, constant.ContextKeyTokenGroup, "auto")
common.SetContextKey(ctx, constant.ContextKeyTokenAutoGroups, []string{"vip", "default"})

groups, err := getModelListGroups(ctx)
require.NoError(t, err)
require.Equal(t, []string{"vip"}, groups.ownerGroups)

common.SetContextKey(ctx, constant.ContextKeyTokenAutoGroups, []string{"vip"})
require.NoError(t, setting.UpdateUserUsableGroupsByJSONString(`{"default":"Default"}`))
groups, err = getModelListGroups(ctx)
require.NoError(t, err)
require.Empty(t, groups.ownerGroups)
}
Loading
Loading