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:
Michael Suchacz
2026-03-12 07:37:33 +01:00
committed by GitHub
parent 3325b86903
commit fba00a6b3a
17 changed files with 958 additions and 165 deletions
+32 -1
View File
@@ -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,
+20 -2
View File
@@ -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)
+54
View File
@@ -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
+152
View File
@@ -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()