package helper import ( "net/http" "net/http/httptest" "testing" "github.com/QuantumNous/new-api/common" "github.com/QuantumNous/new-api/constant" "github.com/QuantumNous/new-api/dto" "github.com/QuantumNous/new-api/model" relaycommon "github.com/QuantumNous/new-api/relay/common" "github.com/QuantumNous/new-api/setting/ratio_setting" "github.com/QuantumNous/new-api/types" "github.com/gin-gonic/gin" "github.com/glebarez/sqlite" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" "gorm.io/gorm" ) const testModelUsePriceSwitch = "test-useprice-switch-model" func setupUsePriceSwitchTest(t *testing.T) { t.Helper() db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{}) require.NoError(t, err) sqlDB, err := db.DB() require.NoError(t, err) sqlDB.SetMaxOpenConns(1) origDB := model.DB origUsingSQLite := common.UsingSQLite origRedisEnabled := common.RedisEnabled model.DB = db common.UsingSQLite = true common.RedisEnabled = false require.NoError(t, db.AutoMigrate(&model.UserChannelRatio{})) require.NoError(t, ratio_setting.UpdateModelPriceByJSONString(`{"`+testModelUsePriceSwitch+`":0.5}`)) require.NoError(t, ratio_setting.UpdateModelRatioByJSONString(`{"`+testModelUsePriceSwitch+`":15}`)) require.NoError(t, ratio_setting.UpdateCompletionRatioByJSONString(`{"`+testModelUsePriceSwitch+`":3}`)) require.NoError(t, ratio_setting.UpdateCacheRatioByJSONString(`{"`+testModelUsePriceSwitch+`":0.1}`)) require.NoError(t, ratio_setting.UpdateCreateCacheRatioByJSONString(`{"`+testModelUsePriceSwitch+`":1.25}`)) require.NoError(t, ratio_setting.UpdateImageRatioByJSONString(`{"`+testModelUsePriceSwitch+`":2.0}`)) require.NoError(t, ratio_setting.UpdateAudioRatioByJSONString(`{"`+testModelUsePriceSwitch+`":5.0}`)) require.NoError(t, ratio_setting.UpdateAudioCompletionRatioByJSONString(`{"`+testModelUsePriceSwitch+`":3.0}`)) t.Cleanup(func() { model.DB = origDB common.UsingSQLite = origUsingSQLite common.RedisEnabled = origRedisEnabled require.NoError(t, sqlDB.Close()) require.NoError(t, ratio_setting.UpdateModelPriceByJSONString(`{}`)) require.NoError(t, ratio_setting.UpdateModelRatioByJSONString(`{}`)) require.NoError(t, ratio_setting.UpdateCompletionRatioByJSONString(`{}`)) require.NoError(t, ratio_setting.UpdateCacheRatioByJSONString(`{}`)) require.NoError(t, ratio_setting.UpdateCreateCacheRatioByJSONString(`{}`)) require.NoError(t, ratio_setting.UpdateImageRatioByJSONString(`{}`)) require.NoError(t, ratio_setting.UpdateAudioRatioByJSONString(`{}`)) require.NoError(t, ratio_setting.UpdateAudioCompletionRatioByJSONString(`{}`)) }) } func buildTestContext(t *testing.T) *gin.Context { t.Helper() gin.SetMode(gin.TestMode) w := httptest.NewRecorder() c, _ := gin.CreateTestContext(w) req, err := http.NewRequest(http.MethodPost, "/v1/chat/completions", nil) require.NoError(t, err) c.Request = req return c } func TestGlobalRatiosFilledWhenUsePriceTrue(t *testing.T) { setupUsePriceSwitchTest(t) c := buildTestContext(t) info := &relaycommon.RelayInfo{ OriginModelName: testModelUsePriceSwitch, UserSetting: dto.UserSetting{}, } priceData, err := ModelPriceHelper(c, info, 100, &types.TokenCountMeta{}) require.NoError(t, err) assert.True(t, priceData.UsePrice) assert.Equal(t, 0.5, priceData.ModelPrice) assert.Equal(t, 0.1, priceData.CacheRatio) assert.Equal(t, 1.25, priceData.CacheCreationRatio) assert.Equal(t, 2.0, priceData.ImageRatio) assert.Equal(t, 5.0, priceData.AudioRatio) assert.Equal(t, 3.0, priceData.AudioCompletionRatio) } func TestGlobalRatiosFilledWhenUseRatioMode(t *testing.T) { setupUsePriceSwitchTest(t) c := buildTestContext(t) require.NoError(t, ratio_setting.UpdateModelPriceByJSONString(`{}`)) info := &relaycommon.RelayInfo{ OriginModelName: testModelUsePriceSwitch, UserSetting: dto.UserSetting{}, } priceData, err := ModelPriceHelper(c, info, 100, &types.TokenCountMeta{}) require.NoError(t, err) assert.False(t, priceData.UsePrice) assert.Equal(t, 15.0, priceData.ModelRatio) assert.Equal(t, 3.0, priceData.CompletionRatio) assert.Equal(t, 0.1, priceData.CacheRatio) assert.Equal(t, 2.0, priceData.ImageRatio) } func TestUpdatePriceDataForChannelPricingOnlyAppliesUserChannelRatio(t *testing.T) { setupUsePriceSwitchTest(t) ratio := &model.UserChannelRatio{ UserId: 9, ModelName: testModelUsePriceSwitch, ChannelId: 9901, Ratio: 0.8, } require.NoError(t, ratio.Insert()) c := buildTestContext(t) c.Set(string(constant.ContextKeyUserId), 9) info := &relaycommon.RelayInfo{ OriginModelName: testModelUsePriceSwitch, PriceData: types.PriceData{ ModelPrice: 0.5, UsePrice: true, QuotaToPreConsume: 1000, UserChannelRatio: 1.0, }, } UpdatePriceDataForChannelPricing(c, info, 9901) assert.Equal(t, 0.8, info.PriceData.UserChannelRatio) assert.Equal(t, 800, info.PriceData.QuotaToPreConsume) assert.Equal(t, 0.5, info.PriceData.ModelPrice) }