mirror of
https://github.com/coder/coder.git
synced 2026-09-24 15:04:27 +08:00
feat(agents): add chat model pricing metadata (#22959)
## Summary - add chat model pricing metadata to the agents admin form and SDK metadata - split pricing into its own section and show default pricing as placeholders - apply default pricing when admins leave pricing fields blank
This commit is contained in:
@@ -553,7 +553,8 @@ func normalizedEnumValue(value string, allowed ...string) *string {
|
||||
return nil
|
||||
}
|
||||
|
||||
// MergeMissingCallConfig fills unset call config values from defaults.
|
||||
// MergeMissingCallConfig fills unset call config values from a provider or
|
||||
// profile default config.
|
||||
func MergeMissingCallConfig(
|
||||
dst *codersdk.ChatModelCallConfig,
|
||||
defaults codersdk.ChatModelCallConfig,
|
||||
@@ -576,9 +577,39 @@ func MergeMissingCallConfig(
|
||||
if dst.FrequencyPenalty == nil {
|
||||
dst.FrequencyPenalty = defaults.FrequencyPenalty
|
||||
}
|
||||
MergeMissingModelCostConfig(&dst.Cost, defaults.Cost)
|
||||
MergeMissingProviderOptions(&dst.ProviderOptions, defaults.ProviderOptions)
|
||||
}
|
||||
|
||||
// MergeMissingModelCostConfig fills unset pricing metadata from defaults.
|
||||
func MergeMissingModelCostConfig(
|
||||
dst **codersdk.ModelCostConfig,
|
||||
defaults *codersdk.ModelCostConfig,
|
||||
) {
|
||||
if defaults == nil {
|
||||
return
|
||||
}
|
||||
if *dst == nil {
|
||||
copied := *defaults
|
||||
*dst = &copied
|
||||
return
|
||||
}
|
||||
|
||||
current := *dst
|
||||
if current.InputPricePerMillionTokens == nil {
|
||||
current.InputPricePerMillionTokens = defaults.InputPricePerMillionTokens
|
||||
}
|
||||
if current.OutputPricePerMillionTokens == nil {
|
||||
current.OutputPricePerMillionTokens = defaults.OutputPricePerMillionTokens
|
||||
}
|
||||
if current.CacheReadPricePerMillionTokens == nil {
|
||||
current.CacheReadPricePerMillionTokens = defaults.CacheReadPricePerMillionTokens
|
||||
}
|
||||
if current.CacheWritePricePerMillionTokens == nil {
|
||||
current.CacheWritePricePerMillionTokens = defaults.CacheWritePricePerMillionTokens
|
||||
}
|
||||
}
|
||||
|
||||
// MergeMissingProviderOptions fills unset provider option fields from defaults.
|
||||
func MergeMissingProviderOptions(
|
||||
dst **codersdk.ChatModelProviderOptions,
|
||||
|
||||
@@ -142,16 +142,25 @@ func TestMergeMissingCallConfig_FillsUnsetFields(t *testing.T) {
|
||||
|
||||
dst := codersdk.ChatModelCallConfig{
|
||||
Temperature: float64Ptr(0.2),
|
||||
Cost: &codersdk.ModelCostConfig{
|
||||
OutputPricePerMillionTokens: float64Ptr(0.7),
|
||||
},
|
||||
ProviderOptions: &codersdk.ChatModelProviderOptions{
|
||||
OpenAI: &codersdk.ChatModelOpenAIProviderOptions{
|
||||
User: stringPtr("alice"),
|
||||
},
|
||||
},
|
||||
}
|
||||
defaults := codersdk.ChatModelCallConfig{
|
||||
defaultCallConfig := codersdk.ChatModelCallConfig{
|
||||
MaxOutputTokens: int64Ptr(512),
|
||||
Temperature: float64Ptr(0.9),
|
||||
TopP: float64Ptr(0.8),
|
||||
Cost: &codersdk.ModelCostConfig{
|
||||
InputPricePerMillionTokens: float64Ptr(0.15),
|
||||
OutputPricePerMillionTokens: float64Ptr(0.9),
|
||||
CacheReadPricePerMillionTokens: float64Ptr(0.03),
|
||||
CacheWritePricePerMillionTokens: float64Ptr(0.3),
|
||||
},
|
||||
ProviderOptions: &codersdk.ChatModelProviderOptions{
|
||||
OpenAI: &codersdk.ChatModelOpenAIProviderOptions{
|
||||
User: stringPtr("bob"),
|
||||
@@ -160,7 +169,7 @@ func TestMergeMissingCallConfig_FillsUnsetFields(t *testing.T) {
|
||||
},
|
||||
}
|
||||
|
||||
chatprovider.MergeMissingCallConfig(&dst, defaults)
|
||||
chatprovider.MergeMissingCallConfig(&dst, defaultCallConfig)
|
||||
|
||||
require.NotNil(t, dst.MaxOutputTokens)
|
||||
require.EqualValues(t, 512, *dst.MaxOutputTokens)
|
||||
@@ -168,6 +177,15 @@ func TestMergeMissingCallConfig_FillsUnsetFields(t *testing.T) {
|
||||
require.Equal(t, 0.2, *dst.Temperature)
|
||||
require.NotNil(t, dst.TopP)
|
||||
require.Equal(t, 0.8, *dst.TopP)
|
||||
require.NotNil(t, dst.Cost)
|
||||
require.NotNil(t, dst.Cost.InputPricePerMillionTokens)
|
||||
require.Equal(t, 0.15, *dst.Cost.InputPricePerMillionTokens)
|
||||
require.NotNil(t, dst.Cost.OutputPricePerMillionTokens)
|
||||
require.Equal(t, 0.7, *dst.Cost.OutputPricePerMillionTokens)
|
||||
require.NotNil(t, dst.Cost.CacheReadPricePerMillionTokens)
|
||||
require.Equal(t, 0.03, *dst.Cost.CacheReadPricePerMillionTokens)
|
||||
require.NotNil(t, dst.Cost.CacheWritePricePerMillionTokens)
|
||||
require.Equal(t, 0.3, *dst.Cost.CacheWritePricePerMillionTokens)
|
||||
require.NotNil(t, dst.ProviderOptions)
|
||||
require.NotNil(t, dst.ProviderOptions.OpenAI)
|
||||
require.Equal(t, "alice", *dst.ProviderOptions.OpenAI.User)
|
||||
|
||||
@@ -3095,6 +3095,10 @@ func marshalChatModelCallConfig(
|
||||
return json.RawMessage("{}"), nil
|
||||
}
|
||||
|
||||
if err := validateChatModelCallConfig(modelConfig); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
encoded, err := json.Marshal(modelConfig)
|
||||
if err != nil {
|
||||
return nil, xerrors.Errorf("encode model config: %w", err)
|
||||
@@ -3102,6 +3106,44 @@ func marshalChatModelCallConfig(
|
||||
return encoded, nil
|
||||
}
|
||||
|
||||
func validateChatModelCallConfig(modelConfig *codersdk.ChatModelCallConfig) error {
|
||||
if modelConfig == nil {
|
||||
return nil
|
||||
}
|
||||
|
||||
costConfig := codersdk.ModelCostConfig{}
|
||||
if modelConfig.Cost != nil {
|
||||
costConfig = *modelConfig.Cost
|
||||
}
|
||||
|
||||
pricingFields := []struct {
|
||||
name string
|
||||
value *float64
|
||||
}{
|
||||
{name: "cost.input_price_per_million_tokens", value: costConfig.InputPricePerMillionTokens},
|
||||
{name: "cost.output_price_per_million_tokens", value: costConfig.OutputPricePerMillionTokens},
|
||||
{name: "cost.cache_read_price_per_million_tokens", value: costConfig.CacheReadPricePerMillionTokens},
|
||||
{name: "cost.cache_write_price_per_million_tokens", value: costConfig.CacheWritePricePerMillionTokens},
|
||||
}
|
||||
for _, field := range pricingFields {
|
||||
if err := validateNonNegativeFloat64Field(field.name, field.value); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func validateNonNegativeFloat64Field(name string, value *float64) error {
|
||||
if value == nil {
|
||||
return nil
|
||||
}
|
||||
if *value < 0 {
|
||||
return xerrors.Errorf("%s must be greater than or equal to zero", name)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func unmarshalChatModelCallConfig(
|
||||
raw json.RawMessage,
|
||||
) *codersdk.ChatModelCallConfig {
|
||||
@@ -3130,9 +3172,21 @@ func isZeroChatModelCallConfig(config *codersdk.ChatModelCallConfig) bool {
|
||||
config.TopK == nil &&
|
||||
config.PresencePenalty == nil &&
|
||||
config.FrequencyPenalty == nil &&
|
||||
isZeroModelCostConfig(config.Cost) &&
|
||||
isZeroChatModelProviderOptions(config.ProviderOptions)
|
||||
}
|
||||
|
||||
func isZeroModelCostConfig(cost *codersdk.ModelCostConfig) bool {
|
||||
if cost == nil {
|
||||
return true
|
||||
}
|
||||
|
||||
return cost.InputPricePerMillionTokens == nil &&
|
||||
cost.OutputPricePerMillionTokens == nil &&
|
||||
cost.CacheReadPricePerMillionTokens == nil &&
|
||||
cost.CacheWritePricePerMillionTokens == nil
|
||||
}
|
||||
|
||||
func isZeroChatModelProviderOptions(options *codersdk.ChatModelProviderOptions) bool {
|
||||
if options == nil {
|
||||
return true
|
||||
|
||||
@@ -22,6 +22,7 @@ import (
|
||||
"github.com/coder/coder/v2/coderd/database/dbfake"
|
||||
"github.com/coder/coder/v2/coderd/externalauth"
|
||||
coderdpubsub "github.com/coder/coder/v2/coderd/pubsub"
|
||||
"github.com/coder/coder/v2/coderd/util/ptr"
|
||||
"github.com/coder/coder/v2/codersdk"
|
||||
"github.com/coder/coder/v2/testutil"
|
||||
"github.com/coder/websocket"
|
||||
@@ -902,6 +903,48 @@ func TestListChatModelConfigs(t *testing.T) {
|
||||
require.True(t, found)
|
||||
})
|
||||
|
||||
t.Run("DeserializesLegacyPricingJSON", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
ctx := testutil.Context(t, testutil.WaitLong)
|
||||
client, db := newChatClientWithDatabase(t)
|
||||
firstUser := coderdtest.CreateFirstUser(t, client)
|
||||
|
||||
_, err := client.CreateChatProvider(ctx, codersdk.CreateChatProviderConfigRequest{
|
||||
Provider: "openai",
|
||||
APIKey: "test-api-key",
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
legacyOptions := json.RawMessage(`{"input_price_per_million_tokens":0.15,"output_price_per_million_tokens":0.6,"cache_read_price_per_million_tokens":0.03,"cache_write_price_per_million_tokens":0.3}`)
|
||||
storedConfig, err := db.InsertChatModelConfig(dbauthz.AsSystemRestricted(ctx), database.InsertChatModelConfigParams{
|
||||
Provider: "openai",
|
||||
Model: "gpt-4o-mini-legacy",
|
||||
DisplayName: "GPT-4o Mini Legacy",
|
||||
CreatedBy: uuid.NullUUID{UUID: firstUser.UserID, Valid: true},
|
||||
UpdatedBy: uuid.NullUUID{UUID: firstUser.UserID, Valid: true},
|
||||
Enabled: true,
|
||||
IsDefault: false,
|
||||
ContextLimit: 4096,
|
||||
CompressionThreshold: 80,
|
||||
Options: legacyOptions,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
configs, err := client.ListChatModelConfigs(ctx)
|
||||
require.NoError(t, err)
|
||||
require.Len(t, configs, 1)
|
||||
require.Equal(t, storedConfig.ID, configs[0].ID)
|
||||
requireChatModelPricing(t, configs[0].ModelConfig, &codersdk.ChatModelCallConfig{
|
||||
Cost: &codersdk.ModelCostConfig{
|
||||
InputPricePerMillionTokens: ptr.Ref(0.15),
|
||||
OutputPricePerMillionTokens: ptr.Ref(0.6),
|
||||
CacheReadPricePerMillionTokens: ptr.Ref(0.03),
|
||||
CacheWritePricePerMillionTokens: ptr.Ref(0.3),
|
||||
},
|
||||
})
|
||||
})
|
||||
|
||||
t.Run("SuccessForOrganizationMember", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
@@ -946,11 +989,20 @@ func TestCreateChatModelConfig(t *testing.T) {
|
||||
|
||||
contextLimit := int64(4096)
|
||||
isDefault := true
|
||||
pricing := &codersdk.ChatModelCallConfig{
|
||||
Cost: &codersdk.ModelCostConfig{
|
||||
InputPricePerMillionTokens: ptr.Ref(0.15),
|
||||
OutputPricePerMillionTokens: ptr.Ref(0.6),
|
||||
CacheReadPricePerMillionTokens: ptr.Ref(0.03),
|
||||
CacheWritePricePerMillionTokens: ptr.Ref(0.3),
|
||||
},
|
||||
}
|
||||
modelConfig, err := client.CreateChatModelConfig(ctx, codersdk.CreateChatModelConfigRequest{
|
||||
Provider: "openai",
|
||||
Model: "gpt-4o-mini",
|
||||
ContextLimit: &contextLimit,
|
||||
IsDefault: &isDefault,
|
||||
ModelConfig: pricing,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.NotEqual(t, uuid.Nil, modelConfig.ID)
|
||||
@@ -958,6 +1010,45 @@ func TestCreateChatModelConfig(t *testing.T) {
|
||||
require.Equal(t, "gpt-4o-mini", modelConfig.Model)
|
||||
require.EqualValues(t, 4096, modelConfig.ContextLimit)
|
||||
require.True(t, modelConfig.IsDefault)
|
||||
requireChatModelPricing(t, modelConfig.ModelConfig, pricing)
|
||||
|
||||
configs, err := client.ListChatModelConfigs(ctx)
|
||||
require.NoError(t, err)
|
||||
require.Len(t, configs, 1)
|
||||
requireChatModelPricing(t, configs[0].ModelConfig, pricing)
|
||||
})
|
||||
|
||||
t.Run("RejectsNegativePricing", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
ctx := testutil.Context(t, testutil.WaitLong)
|
||||
client := newChatClient(t)
|
||||
_ = coderdtest.CreateFirstUser(t, client)
|
||||
|
||||
_, err := client.CreateChatProvider(ctx, codersdk.CreateChatProviderConfigRequest{
|
||||
Provider: "openai",
|
||||
APIKey: "test-api-key",
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
contextLimit := int64(4096)
|
||||
_, err = client.CreateChatModelConfig(ctx, codersdk.CreateChatModelConfigRequest{
|
||||
Provider: "openai",
|
||||
Model: "gpt-4o-mini",
|
||||
ContextLimit: &contextLimit,
|
||||
ModelConfig: &codersdk.ChatModelCallConfig{
|
||||
Cost: &codersdk.ModelCostConfig{
|
||||
InputPricePerMillionTokens: ptr.Ref(-0.01),
|
||||
},
|
||||
},
|
||||
})
|
||||
sdkErr := requireSDKError(t, err, http.StatusBadRequest)
|
||||
require.Equal(t, "Invalid model config.", sdkErr.Message)
|
||||
require.Equal(
|
||||
t,
|
||||
"cost.input_price_per_million_tokens must be greater than or equal to zero",
|
||||
sdkErr.Detail,
|
||||
)
|
||||
})
|
||||
|
||||
t.Run("MissingContextLimit", func(t *testing.T) {
|
||||
@@ -1028,14 +1119,53 @@ func TestUpdateChatModelConfig(t *testing.T) {
|
||||
modelConfig := createChatModelConfig(t, client)
|
||||
|
||||
contextLimit := int64(8192)
|
||||
pricing := &codersdk.ChatModelCallConfig{
|
||||
Cost: &codersdk.ModelCostConfig{
|
||||
InputPricePerMillionTokens: ptr.Ref(0.2),
|
||||
OutputPricePerMillionTokens: ptr.Ref(0.8),
|
||||
CacheReadPricePerMillionTokens: ptr.Ref(0.04),
|
||||
CacheWritePricePerMillionTokens: ptr.Ref(0.4),
|
||||
},
|
||||
}
|
||||
updated, err := client.UpdateChatModelConfig(ctx, modelConfig.ID, codersdk.UpdateChatModelConfigRequest{
|
||||
DisplayName: "GPT-4o Mini Updated",
|
||||
ContextLimit: &contextLimit,
|
||||
ModelConfig: pricing,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, modelConfig.ID, updated.ID)
|
||||
require.Equal(t, "GPT-4o Mini Updated", updated.DisplayName)
|
||||
require.EqualValues(t, 8192, updated.ContextLimit)
|
||||
requireChatModelPricing(t, updated.ModelConfig, pricing)
|
||||
|
||||
configs, err := client.ListChatModelConfigs(ctx)
|
||||
require.NoError(t, err)
|
||||
require.Len(t, configs, 1)
|
||||
requireChatModelPricing(t, configs[0].ModelConfig, pricing)
|
||||
})
|
||||
|
||||
t.Run("RejectsNegativePricing", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
ctx := testutil.Context(t, testutil.WaitLong)
|
||||
client := newChatClient(t)
|
||||
_ = coderdtest.CreateFirstUser(t, client)
|
||||
modelConfig := createChatModelConfig(t, client)
|
||||
|
||||
_, err := client.UpdateChatModelConfig(ctx, modelConfig.ID, codersdk.UpdateChatModelConfigRequest{
|
||||
ModelConfig: &codersdk.ChatModelCallConfig{
|
||||
Cost: &codersdk.ModelCostConfig{
|
||||
OutputPricePerMillionTokens: ptr.Ref(-1.0),
|
||||
},
|
||||
},
|
||||
})
|
||||
sdkErr := requireSDKError(t, err, http.StatusBadRequest)
|
||||
require.Equal(t, "Invalid model config.", sdkErr.Message)
|
||||
require.Equal(
|
||||
t,
|
||||
"cost.output_price_per_million_tokens must be greater than or equal to zero",
|
||||
sdkErr.Detail,
|
||||
)
|
||||
})
|
||||
|
||||
t.Run("NotFound", func(t *testing.T) {
|
||||
@@ -3303,6 +3433,28 @@ func TestGetChatFile(t *testing.T) {
|
||||
})
|
||||
}
|
||||
|
||||
func requireChatModelPricing(
|
||||
t *testing.T,
|
||||
actual *codersdk.ChatModelCallConfig,
|
||||
expected *codersdk.ChatModelCallConfig,
|
||||
) {
|
||||
t.Helper()
|
||||
require.NotNil(t, actual)
|
||||
require.NotNil(t, expected)
|
||||
|
||||
require.NotNil(t, actual.Cost)
|
||||
require.NotNil(t, expected.Cost)
|
||||
require.NotNil(t, actual.Cost.InputPricePerMillionTokens)
|
||||
require.NotNil(t, actual.Cost.OutputPricePerMillionTokens)
|
||||
require.NotNil(t, actual.Cost.CacheReadPricePerMillionTokens)
|
||||
require.NotNil(t, actual.Cost.CacheWritePricePerMillionTokens)
|
||||
|
||||
require.Equal(t, *expected.Cost.InputPricePerMillionTokens, *actual.Cost.InputPricePerMillionTokens)
|
||||
require.Equal(t, *expected.Cost.OutputPricePerMillionTokens, *actual.Cost.OutputPricePerMillionTokens)
|
||||
require.Equal(t, *expected.Cost.CacheReadPricePerMillionTokens, *actual.Cost.CacheReadPricePerMillionTokens)
|
||||
require.Equal(t, *expected.Cost.CacheWritePricePerMillionTokens, *actual.Cost.CacheWritePricePerMillionTokens)
|
||||
}
|
||||
|
||||
func createChatModelConfig(t *testing.T, client *codersdk.Client) codersdk.ChatModelConfig {
|
||||
t.Helper()
|
||||
|
||||
|
||||
Reference in New Issue
Block a user