No puede seleccionar más de 25 temas Los temas deben comenzar con una letra o número, pueden incluir guiones ('-') y pueden tener hasta 35 caracteres de largo.
 
 
 

207 líneas
8.0 KiB

  1. package ratio_setting
  2. import (
  3. "testing"
  4. "github.com/QuantumNous/new-api/types"
  5. "github.com/stretchr/testify/require"
  6. )
  7. func validPricingConfig() *types.PricingConfig {
  8. return &types.PricingConfig{
  9. SchemaVersion: types.PricingSchemaVersion,
  10. Scope: types.PricingScopeModel,
  11. Dimensions: []types.PricingDimension{
  12. {Key: "resolution", Source: "request.resolution", Type: "string", Default: nil},
  13. {Key: "duration", Source: "request.duration", Type: "number", Default: nil},
  14. },
  15. Table: []types.PricingRow{
  16. {"resolution": "720P", "duration": float64(5), "price": float64(0.10), "source": types.PricingRowSourceManual},
  17. {"resolution": "720P", "duration": float64(10), "price": float64(0.20), "source": types.PricingRowSourceGenerated},
  18. },
  19. Fallback: types.PricingFallback{Strategy: types.PricingFallbackReject},
  20. }
  21. }
  22. func TestValidatePricingConfigValidDefaults(t *testing.T) {
  23. cfg := validPricingConfig()
  24. require.NoError(t, ValidatePricingConfig("hailuo-test", cfg))
  25. require.Equal(t, types.BillingUnitPerCall, cfg.BillingUnit)
  26. require.Equal(t, types.PreconsumeStrategyExact, cfg.PreconsumeStrategy)
  27. }
  28. func TestValidatePricingConfigSchemaRules(t *testing.T) {
  29. tests := []struct {
  30. name string
  31. mutate func(*types.PricingConfig)
  32. wantErr string
  33. }{
  34. {"schema_version", func(c *types.PricingConfig) { c.SchemaVersion = 2 }, "schema_version"},
  35. {"scope", func(c *types.PricingConfig) { c.Scope = "channel" }, "scope"},
  36. {"dimensions", func(c *types.PricingConfig) { c.Dimensions = nil }, "dimensions"},
  37. {"table", func(c *types.PricingConfig) { c.Table = nil }, "table"},
  38. {"fallback", func(c *types.PricingConfig) { c.Fallback.Strategy = "nearest" }, "fallback.strategy"},
  39. {"default_price", func(c *types.PricingConfig) {
  40. c.Fallback = types.PricingFallback{Strategy: types.PricingFallbackDefault, DefaultPrice: -1}
  41. }, "default_price"},
  42. {"per_call_strategy", func(c *types.PricingConfig) {
  43. c.BillingUnit = types.BillingUnitPerCall
  44. c.PreconsumeStrategy = types.PreconsumeStrategyMinimum
  45. }, "per_call"},
  46. {"usage_strategy", func(c *types.PricingConfig) {
  47. c.BillingUnit = types.BillingUnitPer1MTokens
  48. c.PreconsumeStrategy = types.PreconsumeStrategyExact
  49. }, "per_1m_tokens"},
  50. {"usage_allowlist", func(c *types.PricingConfig) {
  51. c.BillingUnit = types.BillingUnitPer1MTokens
  52. c.PreconsumeStrategy = types.PreconsumeStrategyMinimum
  53. }, "does not support"},
  54. }
  55. for _, tt := range tests {
  56. t.Run(tt.name, func(t *testing.T) {
  57. cfg := validPricingConfig()
  58. tt.mutate(cfg)
  59. err := ValidatePricingConfig("hailuo-test", cfg)
  60. require.Error(t, err)
  61. require.Contains(t, err.Error(), tt.wantErr)
  62. })
  63. }
  64. }
  65. func TestValidatePricingConfigUsageAllowlist(t *testing.T) {
  66. cfg := validPricingConfig()
  67. cfg.BillingUnit = types.BillingUnitPer1MTokens
  68. cfg.PreconsumeStrategy = types.PreconsumeStrategyMinimum
  69. require.NoError(t, ValidatePricingConfig("seedance-2", cfg))
  70. }
  71. func TestSupportsAnyMatrixUsageBillingModelIncludesTianyiYunMiniAlias(t *testing.T) {
  72. require.True(t, SupportsAnyMatrixUsageBillingModel("Doubao-Seedance-2.0-mini"))
  73. }
  74. func TestValidatePricingConfigUsageAllowlistKlingAipingModels(t *testing.T) {
  75. models := []string{
  76. "kling-v1",
  77. "kling-v1-6",
  78. "kling-v2-6",
  79. "kling-v3",
  80. "kling-video-o1",
  81. "kling-v3-omni",
  82. }
  83. for _, model := range models {
  84. t.Run(model, func(t *testing.T) {
  85. cfg := validPricingConfig()
  86. cfg.BillingUnit = types.BillingUnitPer1MTokens
  87. cfg.PreconsumeStrategy = types.PreconsumeStrategyMinimum
  88. require.NoError(t, ValidatePricingConfig(model, cfg))
  89. })
  90. }
  91. }
  92. func TestValidatePricingConfigAllowsDerivedVideoInput(t *testing.T) {
  93. cfg := &types.PricingConfig{
  94. SchemaVersion: types.PricingSchemaVersion,
  95. Scope: types.PricingScopeModel,
  96. BillingUnit: types.BillingUnitPer1MTokens,
  97. PreconsumeStrategy: types.PreconsumeStrategyMinimum,
  98. Dimensions: []types.PricingDimension{
  99. {Key: "video_input", Source: "derived.video_input", Type: "boolean", Default: nil},
  100. },
  101. Table: []types.PricingRow{
  102. {"video_input": false, "price": float64(46), "source": types.PricingRowSourceManual},
  103. {"video_input": true, "price": float64(28), "source": types.PricingRowSourceManual},
  104. },
  105. Fallback: types.PricingFallback{Strategy: types.PricingFallbackReject},
  106. }
  107. require.NoError(t, ValidatePricingConfig("seedance-2", cfg))
  108. }
  109. func TestValidatePricingConfigDimensionRules(t *testing.T) {
  110. tests := []struct {
  111. name string
  112. dim types.PricingDimension
  113. wantErr string
  114. }{
  115. {"bad_key", types.PricingDimension{Key: "Bad", Source: "request.x", Type: "string", Default: nil}, "invalid"},
  116. {"reserved", types.PricingDimension{Key: "price", Source: "request.x", Type: "string", Default: nil}, "reserved"},
  117. {"bad_type", types.PricingDimension{Key: "quality", Source: "request.x", Type: "object", Default: nil}, "type"},
  118. {"bad_source", types.PricingDimension{Key: "quality", Source: "response.usage.total_tokens", Type: "number", Default: nil}, "source"},
  119. {"optional_default_missing", types.PricingDimension{Key: "quality", Source: "request.quality", Type: "string", Optional: true, Default: nil}, "default required"},
  120. {"optional_default_type", types.PricingDimension{Key: "quality", Source: "request.quality", Type: "number", Optional: true, Default: "fast"}, "default type"},
  121. {"required_default", types.PricingDimension{Key: "quality", Source: "request.quality", Type: "string", Default: "fast"}, "default must be null"},
  122. }
  123. for _, tt := range tests {
  124. t.Run(tt.name, func(t *testing.T) {
  125. cfg := validPricingConfig()
  126. cfg.Dimensions = []types.PricingDimension{tt.dim}
  127. cfg.Table = []types.PricingRow{{tt.dim.Key: "x", "price": float64(1), "source": types.PricingRowSourceManual}}
  128. if tt.dim.Type == "number" {
  129. cfg.Table[0][tt.dim.Key] = float64(1)
  130. }
  131. err := ValidatePricingConfig("hailuo-test", cfg)
  132. require.Error(t, err)
  133. require.Contains(t, err.Error(), tt.wantErr)
  134. })
  135. }
  136. }
  137. func TestValidatePricingConfigRowRules(t *testing.T) {
  138. tests := []struct {
  139. name string
  140. mutate func(*types.PricingConfig)
  141. wantErr string
  142. }{
  143. {"missing_dim", func(c *types.PricingConfig) { delete(c.Table[0], "duration") }, "missing dimension"},
  144. {"extra_field", func(c *types.PricingConfig) { c.Table[0]["price_formula"] = "x" }, "unexpected field"},
  145. {"bad_dim_type", func(c *types.PricingConfig) { c.Table[0]["duration"] = "five" }, "value type"},
  146. {"bad_price", func(c *types.PricingConfig) { c.Table[0]["price"] = -1.0 }, "price"},
  147. {"bad_source", func(c *types.PricingConfig) { c.Table[0]["source"] = "api" }, "source"},
  148. {"duplicate", func(c *types.PricingConfig) {
  149. c.Table = append(c.Table, types.PricingRow(types.CloneMapAny(c.Table[0])))
  150. }, "duplicate"},
  151. {"ambiguous", func(c *types.PricingConfig) {
  152. c.Dimensions = []types.PricingDimension{
  153. {Key: "res", Source: "request.res", Type: "string", Default: nil},
  154. {Key: "dur", Source: "request.dur", Type: "number", Default: nil},
  155. }
  156. c.Table = []types.PricingRow{
  157. {"res": "*", "dur": float64(5), "price": float64(0.10), "source": types.PricingRowSourceManual},
  158. {"res": "720P", "dur": "*", "price": float64(0.12), "source": types.PricingRowSourceManual},
  159. }
  160. }, "ambiguous"},
  161. }
  162. for _, tt := range tests {
  163. t.Run(tt.name, func(t *testing.T) {
  164. cfg := validPricingConfig()
  165. tt.mutate(cfg)
  166. err := ValidatePricingConfig("hailuo-test", cfg)
  167. require.Error(t, err)
  168. require.Contains(t, err.Error(), tt.wantErr)
  169. })
  170. }
  171. }
  172. func TestModelPricingRulesRoundTripAndClone(t *testing.T) {
  173. require.NoError(t, UpdateModelPricingRulesByJSONString(`{}`))
  174. cfg := validPricingConfig()
  175. require.NoError(t, SetPricingConfig("hailuo-test", cfg))
  176. cfg.Table[0]["price"] = float64(999)
  177. got := GetPricingConfig("hailuo-test")
  178. require.NotNil(t, got)
  179. require.Equal(t, float64(0.10), got.Table[0]["price"])
  180. jsonStr := ModelPricingRules2JSONString()
  181. require.Contains(t, jsonStr, "hailuo-test")
  182. require.NoError(t, UpdateModelPricingRulesByJSONString(jsonStr))
  183. require.NotNil(t, GetPricingConfig("hailuo-test"))
  184. require.NoError(t, DeletePricingConfig("hailuo-test"))
  185. require.Nil(t, GetPricingConfig("hailuo-test"))
  186. }