mirror of
https://github.com/coder/coder.git
synced 2026-09-24 15:04:27 +08:00
feat: remove native chat cost tracking in favor of AI Gateway cost data (#27330)
## Stack Context This stack makes AI Gateway data and budgets the source of truth for AI spend controls. 1. Re-back the per-chat cost endpoint with AI Gateway data (#27328, merged). 2. Remove native chat usage limits (#27329, merged). 3. **This PR, now based on `main`:** remove native chat cost tracking and its dedicated admin UI. ## Summary Removes native per-message price calculation, model pricing fields, cost persistence, aggregate cost queries, and admin cost API types. It also deletes the Analytics and Spend pages plus their legacy redirects. The AI Gateway-backed per-chat cost row and compact budget indicators remain. The spend documentation is renamed to `spend-management.md` and updated for the remaining surfaces, group budget APIs, CSV export, upgrade handling for native pricing and cost history, and the absence of a deployment-wide spend dashboard. The per-chat cost API documents that data follows AI Gateway retention and reports zero after all matching requests are purged. No schema is dropped in this release. `chat_messages.total_cost_micros` remains nullable and unwritten so replicas from the previous release can continue inserting messages during rolling upgrades. #27600 tracks removal after the compatibility window. > Mux prepared this PR on Mike's behalf.
This commit is contained in:
@@ -1,71 +0,0 @@
|
||||
package chatcost
|
||||
|
||||
import (
|
||||
"github.com/shopspring/decimal"
|
||||
|
||||
"github.com/coder/coder/v2/coderd/util/ptr"
|
||||
"github.com/coder/coder/v2/codersdk"
|
||||
)
|
||||
|
||||
// Returns cost in micros -- millionths of a dollar, rounded up to the next
|
||||
// whole microdollar.
|
||||
// Returns nil when pricing is not configured or when all priced usage fields
|
||||
// are nil, allowing callers to distinguish "zero cost" from "unpriced".
|
||||
func CalculateTotalCostMicros(
|
||||
usage codersdk.ChatMessageUsage,
|
||||
cost *codersdk.ModelCostConfig,
|
||||
) *int64 {
|
||||
if cost == nil {
|
||||
return nil
|
||||
}
|
||||
|
||||
// A cost config with no prices set means pricing is effectively
|
||||
// unconfigured — return nil (unpriced) rather than zero.
|
||||
if cost.InputPricePerMillionTokens == nil &&
|
||||
cost.OutputPricePerMillionTokens == nil &&
|
||||
cost.CacheReadPricePerMillionTokens == nil &&
|
||||
cost.CacheWritePricePerMillionTokens == nil {
|
||||
return nil
|
||||
}
|
||||
|
||||
if usage.InputTokens == nil &&
|
||||
usage.OutputTokens == nil &&
|
||||
usage.ReasoningTokens == nil &&
|
||||
usage.CacheCreationTokens == nil &&
|
||||
usage.CacheReadTokens == nil {
|
||||
return nil
|
||||
}
|
||||
|
||||
// OutputTokens already includes reasoning tokens per provider
|
||||
// semantics (e.g. OpenAI's completion_tokens encompasses
|
||||
// reasoning_tokens). Adding ReasoningTokens here would
|
||||
// double-count.
|
||||
|
||||
// Preserve nil when usage exists only in categories without configured
|
||||
// pricing, so callers can distinguish "unpriced" from "priced at zero".
|
||||
hasMatchingPrice := (usage.InputTokens != nil && cost.InputPricePerMillionTokens != nil) ||
|
||||
(usage.OutputTokens != nil && cost.OutputPricePerMillionTokens != nil) ||
|
||||
(usage.CacheReadTokens != nil && cost.CacheReadPricePerMillionTokens != nil) ||
|
||||
(usage.CacheCreationTokens != nil && cost.CacheWritePricePerMillionTokens != nil)
|
||||
if !hasMatchingPrice {
|
||||
return nil
|
||||
}
|
||||
|
||||
inputMicros := calcCost(usage.InputTokens, cost.InputPricePerMillionTokens)
|
||||
outputMicros := calcCost(usage.OutputTokens, cost.OutputPricePerMillionTokens)
|
||||
cacheReadMicros := calcCost(usage.CacheReadTokens, cost.CacheReadPricePerMillionTokens)
|
||||
cacheWriteMicros := calcCost(usage.CacheCreationTokens, cost.CacheWritePricePerMillionTokens)
|
||||
|
||||
total := inputMicros.
|
||||
Add(outputMicros).
|
||||
Add(cacheReadMicros).
|
||||
Add(cacheWriteMicros)
|
||||
rounded := total.Ceil().IntPart()
|
||||
return &rounded
|
||||
}
|
||||
|
||||
// calcCost returns the cost in fractional microdollars (millionths of a USD)
|
||||
// for the given token count at the specified per-million-token price.
|
||||
func calcCost(tokens *int64, pricePerMillion *decimal.Decimal) decimal.Decimal {
|
||||
return decimal.NewFromInt(ptr.NilToEmpty(tokens)).Mul(ptr.NilToEmpty(pricePerMillion))
|
||||
}
|
||||
@@ -1,163 +0,0 @@
|
||||
package chatcost_test
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"github.com/shopspring/decimal"
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
"github.com/coder/coder/v2/coderd/util/ptr"
|
||||
"github.com/coder/coder/v2/coderd/x/chatd/chatcost"
|
||||
"github.com/coder/coder/v2/codersdk"
|
||||
)
|
||||
|
||||
func TestCalculateTotalCostMicros(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
usage codersdk.ChatMessageUsage
|
||||
cost *codersdk.ModelCostConfig
|
||||
want *int64
|
||||
}{
|
||||
{
|
||||
name: "nil cost returns nil",
|
||||
usage: codersdk.ChatMessageUsage{InputTokens: ptr.Ref[int64](1000)},
|
||||
cost: nil,
|
||||
want: nil,
|
||||
},
|
||||
{
|
||||
name: "all priced usage fields nil returns nil",
|
||||
usage: codersdk.ChatMessageUsage{
|
||||
TotalTokens: ptr.Ref[int64](1234),
|
||||
ContextLimit: ptr.Ref[int64](8192),
|
||||
},
|
||||
cost: &codersdk.ModelCostConfig{
|
||||
InputPricePerMillionTokens: ptr.Ref(decimal.RequireFromString("3")),
|
||||
},
|
||||
want: nil,
|
||||
},
|
||||
{
|
||||
name: "sub-micro total rounds up to 1",
|
||||
usage: codersdk.ChatMessageUsage{InputTokens: ptr.Ref[int64](1)},
|
||||
cost: &codersdk.ModelCostConfig{
|
||||
InputPricePerMillionTokens: ptr.Ref(decimal.RequireFromString("0.01")),
|
||||
},
|
||||
want: ptr.Ref[int64](1),
|
||||
},
|
||||
{
|
||||
name: "simple input only",
|
||||
usage: codersdk.ChatMessageUsage{InputTokens: ptr.Ref[int64](1000)},
|
||||
cost: &codersdk.ModelCostConfig{
|
||||
InputPricePerMillionTokens: ptr.Ref(decimal.RequireFromString("3")),
|
||||
},
|
||||
want: ptr.Ref[int64](3000),
|
||||
},
|
||||
{
|
||||
name: "simple output only",
|
||||
usage: codersdk.ChatMessageUsage{OutputTokens: ptr.Ref[int64](500)},
|
||||
cost: &codersdk.ModelCostConfig{
|
||||
OutputPricePerMillionTokens: ptr.Ref(decimal.RequireFromString("15")),
|
||||
},
|
||||
want: ptr.Ref[int64](7500),
|
||||
},
|
||||
{
|
||||
name: "reasoning tokens included in output total",
|
||||
usage: codersdk.ChatMessageUsage{
|
||||
OutputTokens: ptr.Ref[int64](500),
|
||||
ReasoningTokens: ptr.Ref[int64](200),
|
||||
},
|
||||
cost: &codersdk.ModelCostConfig{
|
||||
OutputPricePerMillionTokens: ptr.Ref(decimal.RequireFromString("15")),
|
||||
},
|
||||
want: ptr.Ref[int64](7500),
|
||||
},
|
||||
{
|
||||
name: "cache read tokens",
|
||||
usage: codersdk.ChatMessageUsage{CacheReadTokens: ptr.Ref[int64](10000)},
|
||||
cost: &codersdk.ModelCostConfig{
|
||||
CacheReadPricePerMillionTokens: ptr.Ref(decimal.RequireFromString("0.3")),
|
||||
},
|
||||
want: ptr.Ref[int64](3000),
|
||||
},
|
||||
{
|
||||
name: "cache creation tokens",
|
||||
usage: codersdk.ChatMessageUsage{CacheCreationTokens: ptr.Ref[int64](5000)},
|
||||
cost: &codersdk.ModelCostConfig{
|
||||
CacheWritePricePerMillionTokens: ptr.Ref(decimal.RequireFromString("3.75")),
|
||||
},
|
||||
want: ptr.Ref[int64](18750),
|
||||
},
|
||||
{
|
||||
name: "full mixed usage totals all components exactly",
|
||||
usage: codersdk.ChatMessageUsage{
|
||||
InputTokens: ptr.Ref[int64](101),
|
||||
OutputTokens: ptr.Ref[int64](201),
|
||||
ReasoningTokens: ptr.Ref[int64](52),
|
||||
CacheReadTokens: ptr.Ref[int64](1005),
|
||||
CacheCreationTokens: ptr.Ref[int64](33),
|
||||
TotalTokens: ptr.Ref[int64](1391),
|
||||
ContextLimit: ptr.Ref[int64](4096),
|
||||
},
|
||||
cost: &codersdk.ModelCostConfig{
|
||||
InputPricePerMillionTokens: ptr.Ref(decimal.RequireFromString("1.23")),
|
||||
OutputPricePerMillionTokens: ptr.Ref(decimal.RequireFromString("4.56")),
|
||||
CacheReadPricePerMillionTokens: ptr.Ref(decimal.RequireFromString("0.7")),
|
||||
CacheWritePricePerMillionTokens: ptr.Ref(decimal.RequireFromString("7.89")),
|
||||
},
|
||||
want: ptr.Ref[int64](2005),
|
||||
},
|
||||
{
|
||||
name: "partial pricing only input contributes",
|
||||
usage: codersdk.ChatMessageUsage{
|
||||
InputTokens: ptr.Ref[int64](1234),
|
||||
OutputTokens: ptr.Ref[int64](999),
|
||||
ReasoningTokens: ptr.Ref[int64](111),
|
||||
CacheReadTokens: ptr.Ref[int64](500),
|
||||
CacheCreationTokens: ptr.Ref[int64](250),
|
||||
},
|
||||
cost: &codersdk.ModelCostConfig{
|
||||
InputPricePerMillionTokens: ptr.Ref(decimal.RequireFromString("2.5")),
|
||||
},
|
||||
want: ptr.Ref[int64](3085),
|
||||
},
|
||||
{
|
||||
name: "zero tokens with pricing returns zero pointer",
|
||||
usage: codersdk.ChatMessageUsage{InputTokens: ptr.Ref[int64](0)},
|
||||
cost: &codersdk.ModelCostConfig{
|
||||
InputPricePerMillionTokens: ptr.Ref(decimal.RequireFromString("3")),
|
||||
},
|
||||
want: ptr.Ref[int64](0),
|
||||
},
|
||||
{
|
||||
name: "usage only in unpriced categories returns nil",
|
||||
usage: codersdk.ChatMessageUsage{InputTokens: ptr.Ref[int64](1000)},
|
||||
cost: &codersdk.ModelCostConfig{
|
||||
OutputPricePerMillionTokens: ptr.Ref(decimal.RequireFromString("15")),
|
||||
},
|
||||
want: nil,
|
||||
},
|
||||
{
|
||||
name: "non nil usage with empty cost config returns nil",
|
||||
usage: codersdk.ChatMessageUsage{InputTokens: ptr.Ref[int64](42)},
|
||||
cost: &codersdk.ModelCostConfig{},
|
||||
want: nil,
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
tt := tt
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
got := chatcost.CalculateTotalCostMicros(tt.usage, tt.cost)
|
||||
|
||||
if tt.want == nil {
|
||||
require.Nil(t, got)
|
||||
} else {
|
||||
require.NotNil(t, got)
|
||||
require.Equal(t, *tt.want, *got)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -2929,7 +2929,6 @@ type chatMessage struct {
|
||||
cacheCreationTokens int64
|
||||
cacheReadTokens int64
|
||||
contextLimit int64
|
||||
totalCostMicros int64
|
||||
runtimeMs int64
|
||||
}
|
||||
|
||||
@@ -2973,7 +2972,6 @@ func appendMessageFields(
|
||||
params.CacheReadTokens = append(params.CacheReadTokens, msg.cacheReadTokens)
|
||||
params.ContextLimit = append(params.ContextLimit, msg.contextLimit)
|
||||
params.Compressed = append(params.Compressed, msg.compressed)
|
||||
params.TotalCostMicros = append(params.TotalCostMicros, msg.totalCostMicros)
|
||||
params.RuntimeMs = append(params.RuntimeMs, msg.runtimeMs)
|
||||
}
|
||||
|
||||
|
||||
@@ -33,7 +33,6 @@ type Message struct {
|
||||
CacheCreationTokens sql.NullInt64
|
||||
CacheReadTokens sql.NullInt64
|
||||
ContextLimit sql.NullInt64
|
||||
TotalCostMicros sql.NullInt64
|
||||
RuntimeMs sql.NullInt64
|
||||
}
|
||||
|
||||
@@ -62,7 +61,6 @@ func toInsertParams(chatID uuid.UUID, messages []Message) database.InsertChatMes
|
||||
CacheReadTokens: make([]int64, n),
|
||||
ContextLimit: make([]int64, n),
|
||||
Compressed: make([]bool, n),
|
||||
TotalCostMicros: make([]int64, n),
|
||||
RuntimeMs: make([]int64, n),
|
||||
}
|
||||
for i, m := range messages {
|
||||
@@ -89,7 +87,6 @@ func toInsertParams(chatID uuid.UUID, messages []Message) database.InsertChatMes
|
||||
params.CacheReadTokens[i] = nullInt64Or(m.CacheReadTokens, 0)
|
||||
params.ContextLimit[i] = nullInt64Or(m.ContextLimit, 0)
|
||||
params.Compressed[i] = m.Compressed
|
||||
params.TotalCostMicros[i] = nullInt64Or(m.TotalCostMicros, 0)
|
||||
params.RuntimeMs[i] = nullInt64Or(m.RuntimeMs, 0)
|
||||
}
|
||||
return params
|
||||
|
||||
@@ -751,7 +751,6 @@ func (s *taskStarter) generateAssistant(
|
||||
outcome.Step.Content = chathooks.ApplyAdmittedToolCalls(outcome.Step.Content, preflight)
|
||||
messages, err := buildCommitStepMessages(buildCommitStepMessagesInput{
|
||||
modelConfigID: prepared.ModelConfigID,
|
||||
modelCallConfig: prepared.ModelConfig,
|
||||
step: stepDataFromPersisted(outcome.Step),
|
||||
toolNameToConfigID: prepared.ToolNameToConfigID,
|
||||
logger: s.opts.Logger,
|
||||
@@ -858,7 +857,6 @@ func (s *taskStarter) executeLocalTools(
|
||||
chathooks.RestoreToolCallOrder(outcome.Step.Content, decision.localToolCalls)
|
||||
messages, err := buildCommitStepMessages(buildCommitStepMessagesInput{
|
||||
modelConfigID: prepared.ModelConfigID,
|
||||
modelCallConfig: prepared.ModelConfig,
|
||||
step: stepDataFromPersisted(outcome.Step),
|
||||
toolNameToConfigID: prepared.ToolNameToConfigID,
|
||||
logger: s.opts.Logger,
|
||||
|
||||
@@ -16,7 +16,6 @@ import (
|
||||
|
||||
"cdr.dev/slog/v3"
|
||||
"github.com/coder/coder/v2/coderd/database"
|
||||
"github.com/coder/coder/v2/coderd/x/chatd/chatcost"
|
||||
"github.com/coder/coder/v2/coderd/x/chatd/chatloop"
|
||||
"github.com/coder/coder/v2/coderd/x/chatd/chatprompt"
|
||||
"github.com/coder/coder/v2/coderd/x/chatd/chatstate"
|
||||
@@ -29,7 +28,6 @@ const interruptedToolResultErrorMessage = "tool call was interrupted before it p
|
||||
|
||||
type buildCommitStepMessagesInput struct {
|
||||
modelConfigID uuid.UUID
|
||||
modelCallConfig codersdk.ChatModelCallConfig
|
||||
step stepData
|
||||
toolNameToConfigID map[string]uuid.UUID
|
||||
logger slog.Logger
|
||||
@@ -60,7 +58,7 @@ func buildCommitStepMessages(input buildCommitStepMessagesInput) (stepMessagesFo
|
||||
if err != nil {
|
||||
return stepMessagesForCommit{}, xerrors.Errorf("marshal assistant content: %w", err)
|
||||
}
|
||||
messages = append(messages, assistantMessage(input.modelConfigID, contentVersion, assistantContent, input.step, input.modelCallConfig))
|
||||
messages = append(messages, assistantMessage(input.modelConfigID, contentVersion, assistantContent, input.step))
|
||||
}
|
||||
|
||||
for _, toolResult := range toolResults {
|
||||
@@ -186,7 +184,6 @@ func assistantMessage(
|
||||
contentVersion int16,
|
||||
content pqtype.NullRawMessage,
|
||||
step stepData,
|
||||
modelCallConfig codersdk.ChatModelCallConfig,
|
||||
) chatstate.Message {
|
||||
msg := baseMessage(database.ChatMessageRoleAssistant, database.ChatMessageVisibilityBoth, modelConfigID, contentVersion, content)
|
||||
if step.Usage != (fantasy.Usage{}) {
|
||||
@@ -196,16 +193,6 @@ func assistantMessage(
|
||||
msg.ReasoningTokens = nullInt64IfNonZero(step.Usage.ReasoningTokens)
|
||||
msg.CacheCreationTokens = nullInt64IfNonZero(step.Usage.CacheCreationTokens)
|
||||
msg.CacheReadTokens = nullInt64IfNonZero(step.Usage.CacheReadTokens)
|
||||
usage := codersdk.ChatMessageUsage{
|
||||
InputTokens: int64PtrIfNonZero(step.Usage.InputTokens),
|
||||
OutputTokens: int64PtrIfNonZero(step.Usage.OutputTokens),
|
||||
ReasoningTokens: int64PtrIfNonZero(step.Usage.ReasoningTokens),
|
||||
CacheCreationTokens: int64PtrIfNonZero(step.Usage.CacheCreationTokens),
|
||||
CacheReadTokens: int64PtrIfNonZero(step.Usage.CacheReadTokens),
|
||||
}
|
||||
if totalCost := chatcost.CalculateTotalCostMicros(usage, modelCallConfig.Cost); totalCost != nil {
|
||||
msg.TotalCostMicros = sql.NullInt64{Int64: *totalCost, Valid: true}
|
||||
}
|
||||
}
|
||||
msg.ContextLimit = step.ContextLimit
|
||||
if step.Runtime > 0 {
|
||||
@@ -237,13 +224,6 @@ func nullInt64IfNonZero(value int64) sql.NullInt64 {
|
||||
return sql.NullInt64{Int64: value, Valid: true}
|
||||
}
|
||||
|
||||
func int64PtrIfNonZero(value int64) *int64 {
|
||||
if value == 0 {
|
||||
return nil
|
||||
}
|
||||
return &value
|
||||
}
|
||||
|
||||
func visibleMessageIndexes(messages []chatstate.Message) []int {
|
||||
indexes := make([]int, 0, len(messages))
|
||||
for i, msg := range messages {
|
||||
|
||||
@@ -10,7 +10,6 @@ import (
|
||||
|
||||
"charm.land/fantasy"
|
||||
"github.com/google/uuid"
|
||||
"github.com/shopspring/decimal"
|
||||
"github.com/sqlc-dev/pqtype"
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
@@ -127,21 +126,13 @@ func TestBuildCommitStepMessages_ProviderExecutedResultsStayAssistantContent(t *
|
||||
require.True(t, parts[1].ProviderExecuted)
|
||||
}
|
||||
|
||||
func TestBuildCommitStepMessages_UsageCostRuntime(t *testing.T) {
|
||||
func TestBuildCommitStepMessages_UsageRuntime(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
inputPrice := decimal.NewFromFloat(2.5)
|
||||
outputPrice := decimal.NewFromFloat(7.5)
|
||||
got, err := buildCommitStepMessages(buildCommitStepMessagesInput{
|
||||
modelConfigID: uuid.New(),
|
||||
contentVersion: chatprompt.CurrentContentVersion,
|
||||
logger: slog.Make(),
|
||||
modelCallConfig: codersdk.ChatModelCallConfig{
|
||||
Cost: &codersdk.ModelCostConfig{
|
||||
InputPricePerMillionTokens: &inputPrice,
|
||||
OutputPricePerMillionTokens: &outputPrice,
|
||||
},
|
||||
},
|
||||
step: stepData{
|
||||
Content: []fantasy.Content{fantasy.TextContent{Text: "usage"}},
|
||||
Usage: fantasy.Usage{InputTokens: 100, OutputTokens: 20, TotalTokens: 120, ReasoningTokens: 3, CacheCreationTokens: 4, CacheReadTokens: 5},
|
||||
@@ -160,8 +151,6 @@ func TestBuildCommitStepMessages_UsageCostRuntime(t *testing.T) {
|
||||
require.Equal(t, sql.NullInt64{Int64: 5, Valid: true}, msg.CacheReadTokens)
|
||||
require.Equal(t, sql.NullInt64{Int64: 4096, Valid: true}, msg.ContextLimit)
|
||||
require.Equal(t, sql.NullInt64{Int64: 1500, Valid: true}, msg.RuntimeMs)
|
||||
require.True(t, msg.TotalCostMicros.Valid)
|
||||
require.Greater(t, msg.TotalCostMicros.Int64, int64(0))
|
||||
}
|
||||
|
||||
func TestBuildCommitStepMessages_ToolTimestampsAndMCPConfigIDs(t *testing.T) {
|
||||
|
||||
Reference in New Issue
Block a user