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:
@@ -1349,13 +1349,6 @@ func New(options *Options) *API {
|
||||
r.Post("/", api.postChats)
|
||||
r.Get("/models", api.listChatModels)
|
||||
r.Get("/watch", api.watchChats)
|
||||
r.Route("/cost", func(r chi.Router) {
|
||||
r.Get("/users", api.chatCostUsers)
|
||||
r.Route("/{user}", func(r chi.Router) {
|
||||
r.Use(httpmw.ExtractUserParam(options.Database))
|
||||
r.Get("/summary", api.chatCostSummary)
|
||||
})
|
||||
})
|
||||
r.Route("/files", func(r chi.Router) {
|
||||
r.Use(httpmw.RateLimit(options.FilesRateLimit, time.Minute))
|
||||
r.Post("/", api.postChatFile)
|
||||
|
||||
@@ -1992,13 +1992,6 @@ func (q *querier) CountConnectionLogs(ctx context.Context, arg database.CountCon
|
||||
return q.db.CountAuthorizedConnectionLogs(ctx, arg, prep)
|
||||
}
|
||||
|
||||
func (q *querier) CountEnabledModelsWithoutPricing(ctx context.Context) (int64, error) {
|
||||
if err := q.authorizeContext(ctx, policy.ActionRead, rbac.ResourceDeploymentConfig); err != nil {
|
||||
return 0, err
|
||||
}
|
||||
return q.db.CountEnabledModelsWithoutPricing(ctx)
|
||||
}
|
||||
|
||||
func (q *querier) CountInProgressPrebuilds(ctx context.Context) ([]database.CountInProgressPrebuildsRow, error) {
|
||||
if err := q.authorizeContext(ctx, policy.ActionRead, rbac.ResourceWorkspace.All()); err != nil {
|
||||
return nil, err
|
||||
@@ -3136,41 +3129,6 @@ func (q *querier) GetChatComputerUseProvider(ctx context.Context) (string, error
|
||||
return q.db.GetChatComputerUseProvider(ctx)
|
||||
}
|
||||
|
||||
func (q *querier) GetChatCostPerChat(ctx context.Context, arg database.GetChatCostPerChatParams) ([]database.GetChatCostPerChatRow, error) {
|
||||
// The owner's chats, may cross orgs. AnyOrganization() authorizes
|
||||
// the caller if they hold read permission on chats owned by
|
||||
// arg.OwnerID in any org they belong to.
|
||||
// TODO(CODAGT-161): the underlying SQL queries filter only by owner_id, not
|
||||
// organization_id.
|
||||
if err := q.authorizeContext(ctx, policy.ActionRead, rbac.ResourceChat.WithOwner(arg.OwnerID.String()).AnyOrganization()); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return q.db.GetChatCostPerChat(ctx, arg)
|
||||
}
|
||||
|
||||
func (q *querier) GetChatCostPerModel(ctx context.Context, arg database.GetChatCostPerModelParams) ([]database.GetChatCostPerModelRow, error) {
|
||||
// See GetChatCostPerChat for the authorization rationale.
|
||||
if err := q.authorizeContext(ctx, policy.ActionRead, rbac.ResourceChat.WithOwner(arg.OwnerID.String()).AnyOrganization()); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return q.db.GetChatCostPerModel(ctx, arg)
|
||||
}
|
||||
|
||||
func (q *querier) GetChatCostPerUser(ctx context.Context, arg database.GetChatCostPerUserParams) ([]database.GetChatCostPerUserRow, error) {
|
||||
if err := q.authorizeContext(ctx, policy.ActionRead, rbac.ResourceChat); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return q.db.GetChatCostPerUser(ctx, arg)
|
||||
}
|
||||
|
||||
func (q *querier) GetChatCostSummary(ctx context.Context, arg database.GetChatCostSummaryParams) (database.GetChatCostSummaryRow, error) {
|
||||
// See GetChatCostPerChat for the authorization rationale.
|
||||
if err := q.authorizeContext(ctx, policy.ActionRead, rbac.ResourceChat.WithOwner(arg.OwnerID.String()).AnyOrganization()); err != nil {
|
||||
return database.GetChatCostSummaryRow{}, err
|
||||
}
|
||||
return q.db.GetChatCostSummary(ctx, arg)
|
||||
}
|
||||
|
||||
func (q *querier) GetChatDebugLoggingAllowUsers(ctx context.Context) (bool, error) {
|
||||
// The allow-users flag is a deployment-wide setting read by any
|
||||
// authenticated chat user. We only require that an explicit actor
|
||||
@@ -3508,13 +3466,6 @@ func (q *querier) GetChatModelConfigsForTelemetry(ctx context.Context) ([]databa
|
||||
return q.db.GetChatModelConfigsForTelemetry(ctx)
|
||||
}
|
||||
|
||||
func (q *querier) GetChatModelUsageCostByChatID(ctx context.Context, chatID uuid.UUID) (database.GetChatModelUsageCostByChatIDRow, error) {
|
||||
if _, err := q.GetChatByID(ctx, chatID); err != nil {
|
||||
return database.GetChatModelUsageCostByChatIDRow{}, err
|
||||
}
|
||||
return q.db.GetChatModelUsageCostByChatID(ctx, chatID)
|
||||
}
|
||||
|
||||
func (q *querier) GetChatPersonalModelOverridesEnabled(ctx context.Context) (bool, error) {
|
||||
// The personal model overrides flag is a deployment-wide setting read by
|
||||
// authenticated chat users. We only require that an explicit actor is
|
||||
|
||||
@@ -872,85 +872,6 @@ func (s *MethodTestSuite) TestChats() {
|
||||
dbm.EXPECT().SoftDeleteContextFileMessages(gomock.Any(), chat.ID).Return(nil).AnyTimes()
|
||||
check.Args(chat.ID).Asserts(chat, policy.ActionUpdate).Returns()
|
||||
}))
|
||||
s.Run("GetChatCostPerChat", s.Mocked(func(dbm *dbmock.MockStore, _ *gofakeit.Faker, check *expects) {
|
||||
arg := database.GetChatCostPerChatParams{
|
||||
OwnerID: uuid.New(),
|
||||
StartDate: time.Date(2025, 1, 1, 0, 0, 0, 0, time.UTC),
|
||||
EndDate: time.Date(2025, 2, 1, 0, 0, 0, 0, time.UTC),
|
||||
}
|
||||
rows := []database.GetChatCostPerChatRow{{
|
||||
RootChatID: uuid.New(),
|
||||
ChatTitle: "chat-cost",
|
||||
TotalCostMicros: 123,
|
||||
MessageCount: 4,
|
||||
TotalInputTokens: 55,
|
||||
TotalOutputTokens: 89,
|
||||
}}
|
||||
dbm.EXPECT().GetChatCostPerChat(gomock.Any(), arg).Return(rows, nil).AnyTimes()
|
||||
check.Args(arg).Asserts(rbac.ResourceChat.WithOwner(arg.OwnerID.String()).AnyOrganization(), policy.ActionRead).Returns(rows)
|
||||
}))
|
||||
s.Run("GetChatCostPerModel", s.Mocked(func(dbm *dbmock.MockStore, _ *gofakeit.Faker, check *expects) {
|
||||
arg := database.GetChatCostPerModelParams{
|
||||
OwnerID: uuid.New(),
|
||||
StartDate: time.Date(2025, 1, 1, 0, 0, 0, 0, time.UTC),
|
||||
EndDate: time.Date(2025, 2, 1, 0, 0, 0, 0, time.UTC),
|
||||
}
|
||||
rows := []database.GetChatCostPerModelRow{{
|
||||
ModelConfigID: uuid.New(),
|
||||
DisplayName: "GPT 4.1",
|
||||
Provider: "openai",
|
||||
Model: "gpt-4.1",
|
||||
TotalCostMicros: 456,
|
||||
MessageCount: 7,
|
||||
TotalInputTokens: 144,
|
||||
TotalOutputTokens: 233,
|
||||
}}
|
||||
dbm.EXPECT().GetChatCostPerModel(gomock.Any(), arg).Return(rows, nil).AnyTimes()
|
||||
check.Args(arg).Asserts(rbac.ResourceChat.WithOwner(arg.OwnerID.String()).AnyOrganization(), policy.ActionRead).Returns(rows)
|
||||
}))
|
||||
s.Run("GetChatCostPerUser", s.Mocked(func(dbm *dbmock.MockStore, _ *gofakeit.Faker, check *expects) {
|
||||
arg := database.GetChatCostPerUserParams{
|
||||
PageOffset: 0,
|
||||
PageLimit: 25,
|
||||
StartDate: time.Date(2025, 1, 1, 0, 0, 0, 0, time.UTC),
|
||||
EndDate: time.Date(2025, 2, 1, 0, 0, 0, 0, time.UTC),
|
||||
Username: "cost-user",
|
||||
}
|
||||
rows := []database.GetChatCostPerUserRow{{
|
||||
UserID: uuid.New(),
|
||||
Username: "cost-user",
|
||||
Name: "Cost User",
|
||||
AvatarURL: "https://example.com/avatar.png",
|
||||
TotalCostMicros: 789,
|
||||
MessageCount: 11,
|
||||
ChatCount: 3,
|
||||
TotalInputTokens: 377,
|
||||
TotalOutputTokens: 610,
|
||||
TotalCount: 1,
|
||||
}}
|
||||
dbm.EXPECT().GetChatCostPerUser(gomock.Any(), arg).Return(rows, nil).AnyTimes()
|
||||
check.Args(arg).Asserts(rbac.ResourceChat, policy.ActionRead).Returns(rows)
|
||||
}))
|
||||
s.Run("GetChatCostSummary", s.Mocked(func(dbm *dbmock.MockStore, _ *gofakeit.Faker, check *expects) {
|
||||
arg := database.GetChatCostSummaryParams{
|
||||
OwnerID: uuid.New(),
|
||||
StartDate: time.Date(2025, 1, 1, 0, 0, 0, 0, time.UTC),
|
||||
EndDate: time.Date(2025, 2, 1, 0, 0, 0, 0, time.UTC),
|
||||
}
|
||||
row := database.GetChatCostSummaryRow{
|
||||
TotalCostMicros: 987,
|
||||
PricedMessageCount: 12,
|
||||
UnpricedMessagesHavingUsageCount: 2,
|
||||
TotalInputTokens: 400,
|
||||
TotalOutputTokens: 800,
|
||||
}
|
||||
dbm.EXPECT().GetChatCostSummary(gomock.Any(), arg).Return(row, nil).AnyTimes()
|
||||
check.Args(arg).Asserts(rbac.ResourceChat.WithOwner(arg.OwnerID.String()).AnyOrganization(), policy.ActionRead).Returns(row)
|
||||
}))
|
||||
s.Run("CountEnabledModelsWithoutPricing", s.Mocked(func(dbm *dbmock.MockStore, _ *gofakeit.Faker, check *expects) {
|
||||
dbm.EXPECT().CountEnabledModelsWithoutPricing(gomock.Any()).Return(int64(3), nil).AnyTimes()
|
||||
check.Args().Asserts(rbac.ResourceDeploymentConfig, policy.ActionRead).Returns(int64(3))
|
||||
}))
|
||||
s.Run("GetChatDiffStatusByChatID", s.Mocked(func(dbm *dbmock.MockStore, faker *gofakeit.Faker, check *expects) {
|
||||
chat := testutil.Fake(s.T(), faker, database.Chat{})
|
||||
diffStatus := testutil.Fake(s.T(), faker, database.ChatDiffStatus{ChatID: chat.ID})
|
||||
@@ -1076,13 +997,6 @@ func (s *MethodTestSuite) TestChats() {
|
||||
dbm.EXPECT().GetAIBridgeChatCost(gomock.Any(), chat.ID).Return(row, nil).AnyTimes()
|
||||
check.Args(chat.ID).Asserts(chat, policy.ActionRead).Returns(row)
|
||||
}))
|
||||
s.Run("GetChatModelUsageCostByChatID", s.Mocked(func(dbm *dbmock.MockStore, faker *gofakeit.Faker, check *expects) {
|
||||
chat := testutil.Fake(s.T(), faker, database.Chat{})
|
||||
row := database.GetChatModelUsageCostByChatIDRow{ChatID: chat.ID, TotalCostMicros: 1000, PricedMessageCount: 2}
|
||||
dbm.EXPECT().GetChatByID(gomock.Any(), chat.ID).Return(chat, nil).AnyTimes()
|
||||
dbm.EXPECT().GetChatModelUsageCostByChatID(gomock.Any(), chat.ID).Return(row, nil).AnyTimes()
|
||||
check.Args(chat.ID).Asserts(chat, policy.ActionRead).Returns(row)
|
||||
}))
|
||||
s.Run("GetChatMessagesByChatIDAscPaginated", s.Mocked(func(dbm *dbmock.MockStore, faker *gofakeit.Faker, check *expects) {
|
||||
chat := testutil.Fake(s.T(), faker, database.Chat{})
|
||||
msgs := []database.ChatMessage{testutil.Fake(s.T(), faker, database.ChatMessage{ChatID: chat.ID})}
|
||||
|
||||
@@ -141,7 +141,6 @@ func ChatMessage(t testing.TB, db database.Store, seed database.ChatMessage) dat
|
||||
CacheReadTokens: []int64{seed.CacheReadTokens.Int64},
|
||||
ContextLimit: []int64{seed.ContextLimit.Int64},
|
||||
Compressed: []bool{seed.Compressed},
|
||||
TotalCostMicros: []int64{seed.TotalCostMicros.Int64},
|
||||
RuntimeMs: []int64{seed.RuntimeMs.Int64},
|
||||
})
|
||||
require.NoError(t, err, "insert chat message")
|
||||
|
||||
@@ -397,7 +397,6 @@ func TestGenerator(t *testing.T) {
|
||||
CacheReadTokens: sql.NullInt64{Int64: 66, Valid: true},
|
||||
ContextLimit: sql.NullInt64{Int64: 77, Valid: true},
|
||||
Compressed: true,
|
||||
TotalCostMicros: sql.NullInt64{Int64: 88, Valid: true},
|
||||
})
|
||||
require.Equal(t, database.ChatMessageRoleAssistant, msg2.Role)
|
||||
require.True(t, msg2.Content.Valid)
|
||||
@@ -410,7 +409,6 @@ func TestGenerator(t *testing.T) {
|
||||
require.Equal(t, sql.NullInt64{Int64: 66, Valid: true}, msg2.CacheReadTokens)
|
||||
require.Equal(t, sql.NullInt64{Int64: 77, Valid: true}, msg2.ContextLimit)
|
||||
require.True(t, msg2.Compressed)
|
||||
require.Equal(t, sql.NullInt64{Int64: 88, Valid: true}, msg2.TotalCostMicros)
|
||||
})
|
||||
|
||||
t.Run("MCPServerConfig", func(t *testing.T) {
|
||||
|
||||
-48
@@ -345,14 +345,6 @@ func (m queryMetricsStore) CountConnectionLogs(ctx context.Context, arg database
|
||||
return r0, r1
|
||||
}
|
||||
|
||||
func (m queryMetricsStore) CountEnabledModelsWithoutPricing(ctx context.Context) (int64, error) {
|
||||
start := time.Now()
|
||||
r0, r1 := m.s.CountEnabledModelsWithoutPricing(ctx)
|
||||
m.queryLatencies.WithLabelValues("CountEnabledModelsWithoutPricing").Observe(time.Since(start).Seconds())
|
||||
m.queryCounts.WithLabelValues(httpmw.ExtractHTTPRoute(ctx), httpmw.ExtractHTTPMethod(ctx), "CountEnabledModelsWithoutPricing").Inc()
|
||||
return r0, r1
|
||||
}
|
||||
|
||||
func (m queryMetricsStore) CountInProgressPrebuilds(ctx context.Context) ([]database.CountInProgressPrebuildsRow, error) {
|
||||
start := time.Now()
|
||||
r0, r1 := m.s.CountInProgressPrebuilds(ctx)
|
||||
@@ -1465,38 +1457,6 @@ func (m queryMetricsStore) GetChatComputerUseProvider(ctx context.Context) (stri
|
||||
return r0, r1
|
||||
}
|
||||
|
||||
func (m queryMetricsStore) GetChatCostPerChat(ctx context.Context, arg database.GetChatCostPerChatParams) ([]database.GetChatCostPerChatRow, error) {
|
||||
start := time.Now()
|
||||
r0, r1 := m.s.GetChatCostPerChat(ctx, arg)
|
||||
m.queryLatencies.WithLabelValues("GetChatCostPerChat").Observe(time.Since(start).Seconds())
|
||||
m.queryCounts.WithLabelValues(httpmw.ExtractHTTPRoute(ctx), httpmw.ExtractHTTPMethod(ctx), "GetChatCostPerChat").Inc()
|
||||
return r0, r1
|
||||
}
|
||||
|
||||
func (m queryMetricsStore) GetChatCostPerModel(ctx context.Context, arg database.GetChatCostPerModelParams) ([]database.GetChatCostPerModelRow, error) {
|
||||
start := time.Now()
|
||||
r0, r1 := m.s.GetChatCostPerModel(ctx, arg)
|
||||
m.queryLatencies.WithLabelValues("GetChatCostPerModel").Observe(time.Since(start).Seconds())
|
||||
m.queryCounts.WithLabelValues(httpmw.ExtractHTTPRoute(ctx), httpmw.ExtractHTTPMethod(ctx), "GetChatCostPerModel").Inc()
|
||||
return r0, r1
|
||||
}
|
||||
|
||||
func (m queryMetricsStore) GetChatCostPerUser(ctx context.Context, arg database.GetChatCostPerUserParams) ([]database.GetChatCostPerUserRow, error) {
|
||||
start := time.Now()
|
||||
r0, r1 := m.s.GetChatCostPerUser(ctx, arg)
|
||||
m.queryLatencies.WithLabelValues("GetChatCostPerUser").Observe(time.Since(start).Seconds())
|
||||
m.queryCounts.WithLabelValues(httpmw.ExtractHTTPRoute(ctx), httpmw.ExtractHTTPMethod(ctx), "GetChatCostPerUser").Inc()
|
||||
return r0, r1
|
||||
}
|
||||
|
||||
func (m queryMetricsStore) GetChatCostSummary(ctx context.Context, arg database.GetChatCostSummaryParams) (database.GetChatCostSummaryRow, error) {
|
||||
start := time.Now()
|
||||
r0, r1 := m.s.GetChatCostSummary(ctx, arg)
|
||||
m.queryLatencies.WithLabelValues("GetChatCostSummary").Observe(time.Since(start).Seconds())
|
||||
m.queryCounts.WithLabelValues(httpmw.ExtractHTTPRoute(ctx), httpmw.ExtractHTTPMethod(ctx), "GetChatCostSummary").Inc()
|
||||
return r0, r1
|
||||
}
|
||||
|
||||
func (m queryMetricsStore) GetChatDebugLoggingAllowUsers(ctx context.Context) (bool, error) {
|
||||
start := time.Now()
|
||||
r0, r1 := m.s.GetChatDebugLoggingAllowUsers(ctx)
|
||||
@@ -1729,14 +1689,6 @@ func (m queryMetricsStore) GetChatModelConfigsForTelemetry(ctx context.Context)
|
||||
return r0, r1
|
||||
}
|
||||
|
||||
func (m queryMetricsStore) GetChatModelUsageCostByChatID(ctx context.Context, chatID uuid.UUID) (database.GetChatModelUsageCostByChatIDRow, error) {
|
||||
start := time.Now()
|
||||
r0, r1 := m.s.GetChatModelUsageCostByChatID(ctx, chatID)
|
||||
m.queryLatencies.WithLabelValues("GetChatModelUsageCostByChatID").Observe(time.Since(start).Seconds())
|
||||
m.queryCounts.WithLabelValues(httpmw.ExtractHTTPRoute(ctx), httpmw.ExtractHTTPMethod(ctx), "GetChatModelUsageCostByChatID").Inc()
|
||||
return r0, r1
|
||||
}
|
||||
|
||||
func (m queryMetricsStore) GetChatPersonalModelOverridesEnabled(ctx context.Context) (bool, error) {
|
||||
start := time.Now()
|
||||
r0, r1 := m.s.GetChatPersonalModelOverridesEnabled(ctx)
|
||||
|
||||
Generated
-90
@@ -528,21 +528,6 @@ func (mr *MockStoreMockRecorder) CountConnectionLogs(ctx, arg any) *gomock.Call
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "CountConnectionLogs", reflect.TypeOf((*MockStore)(nil).CountConnectionLogs), ctx, arg)
|
||||
}
|
||||
|
||||
// CountEnabledModelsWithoutPricing mocks base method.
|
||||
func (m *MockStore) CountEnabledModelsWithoutPricing(ctx context.Context) (int64, error) {
|
||||
m.ctrl.T.Helper()
|
||||
ret := m.ctrl.Call(m, "CountEnabledModelsWithoutPricing", ctx)
|
||||
ret0, _ := ret[0].(int64)
|
||||
ret1, _ := ret[1].(error)
|
||||
return ret0, ret1
|
||||
}
|
||||
|
||||
// CountEnabledModelsWithoutPricing indicates an expected call of CountEnabledModelsWithoutPricing.
|
||||
func (mr *MockStoreMockRecorder) CountEnabledModelsWithoutPricing(ctx any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "CountEnabledModelsWithoutPricing", reflect.TypeOf((*MockStore)(nil).CountEnabledModelsWithoutPricing), ctx)
|
||||
}
|
||||
|
||||
// CountInProgressPrebuilds mocks base method.
|
||||
func (m *MockStore) CountInProgressPrebuilds(ctx context.Context) ([]database.CountInProgressPrebuildsRow, error) {
|
||||
m.ctrl.T.Helper()
|
||||
@@ -2700,66 +2685,6 @@ func (mr *MockStoreMockRecorder) GetChatComputerUseProvider(ctx any) *gomock.Cal
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetChatComputerUseProvider", reflect.TypeOf((*MockStore)(nil).GetChatComputerUseProvider), ctx)
|
||||
}
|
||||
|
||||
// GetChatCostPerChat mocks base method.
|
||||
func (m *MockStore) GetChatCostPerChat(ctx context.Context, arg database.GetChatCostPerChatParams) ([]database.GetChatCostPerChatRow, error) {
|
||||
m.ctrl.T.Helper()
|
||||
ret := m.ctrl.Call(m, "GetChatCostPerChat", ctx, arg)
|
||||
ret0, _ := ret[0].([]database.GetChatCostPerChatRow)
|
||||
ret1, _ := ret[1].(error)
|
||||
return ret0, ret1
|
||||
}
|
||||
|
||||
// GetChatCostPerChat indicates an expected call of GetChatCostPerChat.
|
||||
func (mr *MockStoreMockRecorder) GetChatCostPerChat(ctx, arg any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetChatCostPerChat", reflect.TypeOf((*MockStore)(nil).GetChatCostPerChat), ctx, arg)
|
||||
}
|
||||
|
||||
// GetChatCostPerModel mocks base method.
|
||||
func (m *MockStore) GetChatCostPerModel(ctx context.Context, arg database.GetChatCostPerModelParams) ([]database.GetChatCostPerModelRow, error) {
|
||||
m.ctrl.T.Helper()
|
||||
ret := m.ctrl.Call(m, "GetChatCostPerModel", ctx, arg)
|
||||
ret0, _ := ret[0].([]database.GetChatCostPerModelRow)
|
||||
ret1, _ := ret[1].(error)
|
||||
return ret0, ret1
|
||||
}
|
||||
|
||||
// GetChatCostPerModel indicates an expected call of GetChatCostPerModel.
|
||||
func (mr *MockStoreMockRecorder) GetChatCostPerModel(ctx, arg any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetChatCostPerModel", reflect.TypeOf((*MockStore)(nil).GetChatCostPerModel), ctx, arg)
|
||||
}
|
||||
|
||||
// GetChatCostPerUser mocks base method.
|
||||
func (m *MockStore) GetChatCostPerUser(ctx context.Context, arg database.GetChatCostPerUserParams) ([]database.GetChatCostPerUserRow, error) {
|
||||
m.ctrl.T.Helper()
|
||||
ret := m.ctrl.Call(m, "GetChatCostPerUser", ctx, arg)
|
||||
ret0, _ := ret[0].([]database.GetChatCostPerUserRow)
|
||||
ret1, _ := ret[1].(error)
|
||||
return ret0, ret1
|
||||
}
|
||||
|
||||
// GetChatCostPerUser indicates an expected call of GetChatCostPerUser.
|
||||
func (mr *MockStoreMockRecorder) GetChatCostPerUser(ctx, arg any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetChatCostPerUser", reflect.TypeOf((*MockStore)(nil).GetChatCostPerUser), ctx, arg)
|
||||
}
|
||||
|
||||
// GetChatCostSummary mocks base method.
|
||||
func (m *MockStore) GetChatCostSummary(ctx context.Context, arg database.GetChatCostSummaryParams) (database.GetChatCostSummaryRow, error) {
|
||||
m.ctrl.T.Helper()
|
||||
ret := m.ctrl.Call(m, "GetChatCostSummary", ctx, arg)
|
||||
ret0, _ := ret[0].(database.GetChatCostSummaryRow)
|
||||
ret1, _ := ret[1].(error)
|
||||
return ret0, ret1
|
||||
}
|
||||
|
||||
// GetChatCostSummary indicates an expected call of GetChatCostSummary.
|
||||
func (mr *MockStoreMockRecorder) GetChatCostSummary(ctx, arg any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetChatCostSummary", reflect.TypeOf((*MockStore)(nil).GetChatCostSummary), ctx, arg)
|
||||
}
|
||||
|
||||
// GetChatDebugLoggingAllowUsers mocks base method.
|
||||
func (m *MockStore) GetChatDebugLoggingAllowUsers(ctx context.Context) (bool, error) {
|
||||
m.ctrl.T.Helper()
|
||||
@@ -3195,21 +3120,6 @@ func (mr *MockStoreMockRecorder) GetChatModelConfigsForTelemetry(ctx any) *gomoc
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetChatModelConfigsForTelemetry", reflect.TypeOf((*MockStore)(nil).GetChatModelConfigsForTelemetry), ctx)
|
||||
}
|
||||
|
||||
// GetChatModelUsageCostByChatID mocks base method.
|
||||
func (m *MockStore) GetChatModelUsageCostByChatID(ctx context.Context, chatID uuid.UUID) (database.GetChatModelUsageCostByChatIDRow, error) {
|
||||
m.ctrl.T.Helper()
|
||||
ret := m.ctrl.Call(m, "GetChatModelUsageCostByChatID", ctx, chatID)
|
||||
ret0, _ := ret[0].(database.GetChatModelUsageCostByChatIDRow)
|
||||
ret1, _ := ret[1].(error)
|
||||
return ret0, ret1
|
||||
}
|
||||
|
||||
// GetChatModelUsageCostByChatID indicates an expected call of GetChatModelUsageCostByChatID.
|
||||
func (mr *MockStoreMockRecorder) GetChatModelUsageCostByChatID(ctx, chatID any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetChatModelUsageCostByChatID", reflect.TypeOf((*MockStore)(nil).GetChatModelUsageCostByChatID), ctx, chatID)
|
||||
}
|
||||
|
||||
// GetChatPersonalModelOverridesEnabled mocks base method.
|
||||
func (m *MockStore) GetChatPersonalModelOverridesEnabled(ctx context.Context) (bool, error) {
|
||||
m.ctrl.T.Helper()
|
||||
|
||||
Generated
-21
@@ -100,9 +100,6 @@ type sqlcQuerier interface {
|
||||
// whether the chat is in a "1" sub-state.
|
||||
CountChatQueuedMessages(ctx context.Context, chatID uuid.UUID) (int64, error)
|
||||
CountConnectionLogs(ctx context.Context, arg CountConnectionLogsParams) (int64, error)
|
||||
// Counts enabled, non-deleted model configs that lack both input and
|
||||
// output pricing in their JSONB options.cost configuration.
|
||||
CountEnabledModelsWithoutPricing(ctx context.Context) (int64, error)
|
||||
// CountInProgressPrebuilds returns the number of in-progress prebuilds, grouped by preset ID and transition.
|
||||
// Prebuild considered in-progress if it's in the "pending", "starting", "stopping", or "deleting" state.
|
||||
CountInProgressPrebuilds(ctx context.Context) ([]CountInProgressPrebuildsRow, error)
|
||||
@@ -419,19 +416,6 @@ type sqlcQuerier interface {
|
||||
GetChatByIDForUpdate(ctx context.Context, id uuid.UUID) (Chat, error)
|
||||
GetChatCompactionModelOverride(ctx context.Context) (string, error)
|
||||
GetChatComputerUseProvider(ctx context.Context) (string, error)
|
||||
// Per-root-chat cost breakdown for a single user within a date range.
|
||||
// Groups by root_chat_id so forked chats roll up under their root.
|
||||
// Only counts assistant-role messages.
|
||||
GetChatCostPerChat(ctx context.Context, arg GetChatCostPerChatParams) ([]GetChatCostPerChatRow, error)
|
||||
// Per-model cost breakdown for a single user within a date range.
|
||||
// Only counts assistant-role messages that have a model_config_id.
|
||||
GetChatCostPerModel(ctx context.Context, arg GetChatCostPerModelParams) ([]GetChatCostPerModelRow, error)
|
||||
// Deployment-wide per-user cost rollup within a date range.
|
||||
// Only counts assistant-role messages.
|
||||
GetChatCostPerUser(ctx context.Context, arg GetChatCostPerUserParams) ([]GetChatCostPerUserRow, error)
|
||||
// Aggregate cost summary for a single user within a date range.
|
||||
// Only counts assistant-role messages.
|
||||
GetChatCostSummary(ctx context.Context, arg GetChatCostSummaryParams) (GetChatCostSummaryRow, error)
|
||||
// GetChatDebugLoggingAllowUsers returns the runtime admin setting that
|
||||
// allows users to opt into chat debug logging when the deployment does
|
||||
// not already force debug logging on globally.
|
||||
@@ -497,11 +481,6 @@ type sqlcQuerier interface {
|
||||
// Returns all model configurations for telemetry snapshot collection.
|
||||
// deleted = false guarantees ai_provider_id is non-null, so INNER JOIN is safe.
|
||||
GetChatModelConfigsForTelemetry(ctx context.Context) ([]GetChatModelConfigsForTelemetryRow, error)
|
||||
// Assistant-message cost rolled up over the requested chat's subtree: the
|
||||
// chat itself plus every descendant reachable through parent_chat_id. A
|
||||
// root chat therefore reports its whole tree, while a subagent chat
|
||||
// reports only its own spend plus any nested subagents it spawned.
|
||||
GetChatModelUsageCostByChatID(ctx context.Context, chatID uuid.UUID) (GetChatModelUsageCostByChatIDRow, error)
|
||||
// GetChatPersonalModelOverridesEnabled returns whether users may configure
|
||||
// personal chat model overrides. It defaults to false when unset.
|
||||
GetChatPersonalModelOverridesEnabled(ctx context.Context) (bool, error)
|
||||
|
||||
@@ -12276,7 +12276,6 @@ func TestInsertChatMessages(t *testing.T) {
|
||||
CacheReadTokens: []int64{0},
|
||||
ContextLimit: []int64{0},
|
||||
Compressed: []bool{false},
|
||||
TotalCostMicros: []int64{0},
|
||||
RuntimeMs: []int64{0},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
@@ -12337,7 +12336,6 @@ func TestInsertChatMessages(t *testing.T) {
|
||||
CacheReadTokens: []int64{0, 0, 0},
|
||||
ContextLimit: []int64{0, 0, 0},
|
||||
Compressed: []bool{false, false, false},
|
||||
TotalCostMicros: []int64{0, 100, 0},
|
||||
RuntimeMs: []int64{0, 500, 0},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
@@ -12367,10 +12365,6 @@ func TestInsertChatMessages(t *testing.T) {
|
||||
require.Equal(t, int64(20), msgs[1].OutputTokens.Int64)
|
||||
|
||||
// Verify cost: assistant has cost, others NULL.
|
||||
require.True(t, msgs[1].TotalCostMicros.Valid)
|
||||
require.Equal(t, int64(100), msgs[1].TotalCostMicros.Int64)
|
||||
require.False(t, msgs[0].TotalCostMicros.Valid)
|
||||
require.False(t, msgs[2].TotalCostMicros.Valid)
|
||||
|
||||
// Verify runtime_ms on assistant message.
|
||||
require.True(t, msgs[1].RuntimeMs.Valid)
|
||||
@@ -12414,7 +12408,6 @@ func insertChatMessagesInvertedTimestamps(t *testing.T, db database.Store, sqlDB
|
||||
CacheReadTokens: make([]int64, count),
|
||||
ContextLimit: make([]int64, count),
|
||||
Compressed: make([]bool, count),
|
||||
TotalCostMicros: make([]int64, count),
|
||||
RuntimeMs: make([]int64, count),
|
||||
})
|
||||
require.NoError(t, err)
|
||||
@@ -12586,7 +12579,6 @@ func TestGetChatMessagesForPromptByChatID(t *testing.T) {
|
||||
CacheCreationTokens: []int64{0},
|
||||
CacheReadTokens: []int64{0},
|
||||
ContextLimit: []int64{0},
|
||||
TotalCostMicros: []int64{0},
|
||||
RuntimeMs: []int64{0},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
@@ -15610,7 +15602,6 @@ func TestUpdateChatLastTurnSummary(t *testing.T) {
|
||||
CacheReadTokens: []int64{0},
|
||||
ContextLimit: []int64{0},
|
||||
Compressed: []bool{false},
|
||||
TotalCostMicros: []int64{0},
|
||||
RuntimeMs: []int64{0},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
@@ -15732,7 +15723,6 @@ func TestUpdateChatSummary(t *testing.T) {
|
||||
CacheReadTokens: []int64{0},
|
||||
ContextLimit: []int64{0},
|
||||
Compressed: []bool{false},
|
||||
TotalCostMicros: []int64{0},
|
||||
RuntimeMs: []int64{0},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
@@ -17511,7 +17501,6 @@ func TestGetChatsFilter(t *testing.T) {
|
||||
CacheReadTokens: []int64{0},
|
||||
ContextLimit: []int64{0},
|
||||
Compressed: []bool{false},
|
||||
TotalCostMicros: []int64{0},
|
||||
RuntimeMs: []int64{0},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
@@ -17753,7 +17742,6 @@ func TestGetChatsSearch(t *testing.T) {
|
||||
CacheReadTokens: []int64{0},
|
||||
ContextLimit: []int64{0},
|
||||
Compressed: []bool{false},
|
||||
TotalCostMicros: []int64{0},
|
||||
RuntimeMs: []int64{0},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
@@ -17986,7 +17974,6 @@ func TestChatHasUnread(t *testing.T) {
|
||||
CacheReadTokens: []int64{0},
|
||||
ContextLimit: []int64{0},
|
||||
Compressed: []bool{false},
|
||||
TotalCostMicros: []int64{0},
|
||||
RuntimeMs: []int64{0},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
Generated
+1
-485
@@ -6992,30 +6992,6 @@ func (q *sqlQuerier) CountChatQueuedMessages(ctx context.Context, chatID uuid.UU
|
||||
return count, err
|
||||
}
|
||||
|
||||
const countEnabledModelsWithoutPricing = `-- name: CountEnabledModelsWithoutPricing :one
|
||||
SELECT COUNT(*)::bigint AS count
|
||||
FROM chat_model_configs
|
||||
WHERE enabled = TRUE
|
||||
AND deleted = FALSE
|
||||
AND (
|
||||
options->'cost' IS NULL
|
||||
OR options->'cost' = 'null'::jsonb
|
||||
OR (
|
||||
(options->'cost'->>'input_price_per_million_tokens' IS NULL)
|
||||
AND (options->'cost'->>'output_price_per_million_tokens' IS NULL)
|
||||
)
|
||||
)
|
||||
`
|
||||
|
||||
// Counts enabled, non-deleted model configs that lack both input and
|
||||
// output pricing in their JSONB options.cost configuration.
|
||||
func (q *sqlQuerier) CountEnabledModelsWithoutPricing(ctx context.Context) (int64, error) {
|
||||
row := q.db.QueryRowContext(ctx, countEnabledModelsWithoutPricing)
|
||||
var count int64
|
||||
err := row.Scan(&count)
|
||||
return count, err
|
||||
}
|
||||
|
||||
const deleteAllChatHeartbeats = `-- name: DeleteAllChatHeartbeats :exec
|
||||
DELETE FROM chat_heartbeats WHERE chat_id = $1::uuid
|
||||
`
|
||||
@@ -7705,398 +7681,6 @@ func (q *sqlQuerier) GetChatByIDForUpdate(ctx context.Context, id uuid.UUID) (Ch
|
||||
return i, err
|
||||
}
|
||||
|
||||
const getChatCostPerChat = `-- name: GetChatCostPerChat :many
|
||||
WITH chat_costs AS (
|
||||
SELECT
|
||||
COALESCE(c.root_chat_id, c.id) AS root_chat_id,
|
||||
COALESCE(SUM(cm.total_cost_micros), 0)::bigint AS total_cost_micros,
|
||||
COUNT(*) FILTER (
|
||||
WHERE cm.input_tokens IS NOT NULL
|
||||
OR cm.output_tokens IS NOT NULL
|
||||
OR cm.reasoning_tokens IS NOT NULL
|
||||
OR cm.cache_creation_tokens IS NOT NULL
|
||||
OR cm.cache_read_tokens IS NOT NULL
|
||||
)::bigint AS message_count,
|
||||
COALESCE(SUM(cm.input_tokens), 0)::bigint AS total_input_tokens,
|
||||
COALESCE(SUM(cm.output_tokens), 0)::bigint AS total_output_tokens,
|
||||
COALESCE(SUM(cm.cache_read_tokens), 0)::bigint AS total_cache_read_tokens,
|
||||
COALESCE(SUM(cm.cache_creation_tokens), 0)::bigint AS total_cache_creation_tokens,
|
||||
COALESCE(SUM(cm.runtime_ms), 0)::bigint AS total_runtime_ms
|
||||
FROM chat_messages cm
|
||||
JOIN chats c ON c.id = cm.chat_id
|
||||
WHERE c.owner_id = $1::uuid
|
||||
AND cm.role = 'assistant'
|
||||
AND cm.created_at >= $2::timestamptz
|
||||
AND cm.created_at < $3::timestamptz
|
||||
GROUP BY COALESCE(c.root_chat_id, c.id)
|
||||
)
|
||||
SELECT
|
||||
cc.root_chat_id,
|
||||
COALESCE(rc.title, '') AS chat_title,
|
||||
cc.total_cost_micros,
|
||||
cc.message_count,
|
||||
cc.total_input_tokens,
|
||||
cc.total_output_tokens,
|
||||
cc.total_cache_read_tokens,
|
||||
cc.total_cache_creation_tokens,
|
||||
cc.total_runtime_ms
|
||||
FROM chat_costs cc
|
||||
LEFT JOIN chats rc ON rc.id = cc.root_chat_id
|
||||
ORDER BY cc.total_cost_micros DESC
|
||||
`
|
||||
|
||||
type GetChatCostPerChatParams struct {
|
||||
OwnerID uuid.UUID `db:"owner_id" json:"owner_id"`
|
||||
StartDate time.Time `db:"start_date" json:"start_date"`
|
||||
EndDate time.Time `db:"end_date" json:"end_date"`
|
||||
}
|
||||
|
||||
type GetChatCostPerChatRow struct {
|
||||
RootChatID uuid.UUID `db:"root_chat_id" json:"root_chat_id"`
|
||||
ChatTitle string `db:"chat_title" json:"chat_title"`
|
||||
TotalCostMicros int64 `db:"total_cost_micros" json:"total_cost_micros"`
|
||||
MessageCount int64 `db:"message_count" json:"message_count"`
|
||||
TotalInputTokens int64 `db:"total_input_tokens" json:"total_input_tokens"`
|
||||
TotalOutputTokens int64 `db:"total_output_tokens" json:"total_output_tokens"`
|
||||
TotalCacheReadTokens int64 `db:"total_cache_read_tokens" json:"total_cache_read_tokens"`
|
||||
TotalCacheCreationTokens int64 `db:"total_cache_creation_tokens" json:"total_cache_creation_tokens"`
|
||||
TotalRuntimeMs int64 `db:"total_runtime_ms" json:"total_runtime_ms"`
|
||||
}
|
||||
|
||||
// Per-root-chat cost breakdown for a single user within a date range.
|
||||
// Groups by root_chat_id so forked chats roll up under their root.
|
||||
// Only counts assistant-role messages.
|
||||
func (q *sqlQuerier) GetChatCostPerChat(ctx context.Context, arg GetChatCostPerChatParams) ([]GetChatCostPerChatRow, error) {
|
||||
rows, err := q.db.QueryContext(ctx, getChatCostPerChat, arg.OwnerID, arg.StartDate, arg.EndDate)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer rows.Close()
|
||||
var items []GetChatCostPerChatRow
|
||||
for rows.Next() {
|
||||
var i GetChatCostPerChatRow
|
||||
if err := rows.Scan(
|
||||
&i.RootChatID,
|
||||
&i.ChatTitle,
|
||||
&i.TotalCostMicros,
|
||||
&i.MessageCount,
|
||||
&i.TotalInputTokens,
|
||||
&i.TotalOutputTokens,
|
||||
&i.TotalCacheReadTokens,
|
||||
&i.TotalCacheCreationTokens,
|
||||
&i.TotalRuntimeMs,
|
||||
); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
items = append(items, i)
|
||||
}
|
||||
if err := rows.Close(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if err := rows.Err(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return items, nil
|
||||
}
|
||||
|
||||
const getChatCostPerModel = `-- name: GetChatCostPerModel :many
|
||||
SELECT
|
||||
cmc.id AS model_config_id,
|
||||
cmc.display_name,
|
||||
COALESCE(ap.type::text, '')::text AS provider,
|
||||
cmc.model,
|
||||
COALESCE(SUM(cm.total_cost_micros), 0)::bigint AS total_cost_micros,
|
||||
COUNT(*) FILTER (
|
||||
WHERE cm.input_tokens IS NOT NULL
|
||||
OR cm.output_tokens IS NOT NULL
|
||||
OR cm.reasoning_tokens IS NOT NULL
|
||||
OR cm.cache_creation_tokens IS NOT NULL
|
||||
OR cm.cache_read_tokens IS NOT NULL
|
||||
)::bigint AS message_count,
|
||||
COALESCE(SUM(cm.input_tokens), 0)::bigint AS total_input_tokens,
|
||||
COALESCE(SUM(cm.output_tokens), 0)::bigint AS total_output_tokens,
|
||||
COALESCE(SUM(cm.cache_read_tokens), 0)::bigint AS total_cache_read_tokens,
|
||||
COALESCE(SUM(cm.cache_creation_tokens), 0)::bigint AS total_cache_creation_tokens,
|
||||
COALESCE(SUM(cm.runtime_ms), 0)::bigint AS total_runtime_ms
|
||||
FROM
|
||||
chat_messages cm
|
||||
JOIN
|
||||
chats c ON c.id = cm.chat_id
|
||||
JOIN
|
||||
chat_model_configs cmc ON cmc.id = cm.model_config_id
|
||||
LEFT JOIN
|
||||
ai_providers ap ON ap.id = cmc.ai_provider_id
|
||||
WHERE
|
||||
c.owner_id = $1::uuid
|
||||
AND cm.role = 'assistant'
|
||||
AND cm.created_at >= $2::timestamptz
|
||||
AND cm.created_at < $3::timestamptz
|
||||
GROUP BY
|
||||
cmc.id, cmc.display_name, ap.type, cmc.model
|
||||
ORDER BY
|
||||
total_cost_micros DESC
|
||||
`
|
||||
|
||||
type GetChatCostPerModelParams struct {
|
||||
OwnerID uuid.UUID `db:"owner_id" json:"owner_id"`
|
||||
StartDate time.Time `db:"start_date" json:"start_date"`
|
||||
EndDate time.Time `db:"end_date" json:"end_date"`
|
||||
}
|
||||
|
||||
type GetChatCostPerModelRow struct {
|
||||
ModelConfigID uuid.UUID `db:"model_config_id" json:"model_config_id"`
|
||||
DisplayName string `db:"display_name" json:"display_name"`
|
||||
Provider string `db:"provider" json:"provider"`
|
||||
Model string `db:"model" json:"model"`
|
||||
TotalCostMicros int64 `db:"total_cost_micros" json:"total_cost_micros"`
|
||||
MessageCount int64 `db:"message_count" json:"message_count"`
|
||||
TotalInputTokens int64 `db:"total_input_tokens" json:"total_input_tokens"`
|
||||
TotalOutputTokens int64 `db:"total_output_tokens" json:"total_output_tokens"`
|
||||
TotalCacheReadTokens int64 `db:"total_cache_read_tokens" json:"total_cache_read_tokens"`
|
||||
TotalCacheCreationTokens int64 `db:"total_cache_creation_tokens" json:"total_cache_creation_tokens"`
|
||||
TotalRuntimeMs int64 `db:"total_runtime_ms" json:"total_runtime_ms"`
|
||||
}
|
||||
|
||||
// Per-model cost breakdown for a single user within a date range.
|
||||
// Only counts assistant-role messages that have a model_config_id.
|
||||
func (q *sqlQuerier) GetChatCostPerModel(ctx context.Context, arg GetChatCostPerModelParams) ([]GetChatCostPerModelRow, error) {
|
||||
rows, err := q.db.QueryContext(ctx, getChatCostPerModel, arg.OwnerID, arg.StartDate, arg.EndDate)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer rows.Close()
|
||||
var items []GetChatCostPerModelRow
|
||||
for rows.Next() {
|
||||
var i GetChatCostPerModelRow
|
||||
if err := rows.Scan(
|
||||
&i.ModelConfigID,
|
||||
&i.DisplayName,
|
||||
&i.Provider,
|
||||
&i.Model,
|
||||
&i.TotalCostMicros,
|
||||
&i.MessageCount,
|
||||
&i.TotalInputTokens,
|
||||
&i.TotalOutputTokens,
|
||||
&i.TotalCacheReadTokens,
|
||||
&i.TotalCacheCreationTokens,
|
||||
&i.TotalRuntimeMs,
|
||||
); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
items = append(items, i)
|
||||
}
|
||||
if err := rows.Close(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if err := rows.Err(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return items, nil
|
||||
}
|
||||
|
||||
const getChatCostPerUser = `-- name: GetChatCostPerUser :many
|
||||
WITH chat_cost_users AS (
|
||||
SELECT
|
||||
c.owner_id AS user_id,
|
||||
u.username,
|
||||
u.name,
|
||||
u.avatar_url,
|
||||
COALESCE(SUM(cm.total_cost_micros), 0)::bigint AS total_cost_micros,
|
||||
COUNT(*) FILTER (
|
||||
WHERE cm.input_tokens IS NOT NULL
|
||||
OR cm.output_tokens IS NOT NULL
|
||||
OR cm.reasoning_tokens IS NOT NULL
|
||||
OR cm.cache_creation_tokens IS NOT NULL
|
||||
OR cm.cache_read_tokens IS NOT NULL
|
||||
)::bigint AS message_count,
|
||||
COUNT(DISTINCT COALESCE(c.root_chat_id, c.id))::bigint AS chat_count,
|
||||
COALESCE(SUM(cm.input_tokens), 0)::bigint AS total_input_tokens,
|
||||
COALESCE(SUM(cm.output_tokens), 0)::bigint AS total_output_tokens,
|
||||
COALESCE(SUM(cm.cache_read_tokens), 0)::bigint AS total_cache_read_tokens,
|
||||
COALESCE(SUM(cm.cache_creation_tokens), 0)::bigint AS total_cache_creation_tokens,
|
||||
COALESCE(SUM(cm.runtime_ms), 0)::bigint AS total_runtime_ms
|
||||
FROM
|
||||
chat_messages cm
|
||||
JOIN
|
||||
chats c ON c.id = cm.chat_id
|
||||
JOIN
|
||||
users u ON u.id = c.owner_id
|
||||
WHERE
|
||||
cm.role = 'assistant'
|
||||
AND cm.created_at >= $3::timestamptz
|
||||
AND cm.created_at < $4::timestamptz
|
||||
AND (
|
||||
$5::text = ''
|
||||
OR u.username ILIKE '%' || $5::text || '%'
|
||||
OR u.name ILIKE '%' || $5::text || '%'
|
||||
)
|
||||
GROUP BY
|
||||
c.owner_id,
|
||||
u.username,
|
||||
u.name,
|
||||
u.avatar_url
|
||||
)
|
||||
SELECT
|
||||
user_id,
|
||||
username,
|
||||
name,
|
||||
avatar_url,
|
||||
total_cost_micros,
|
||||
message_count,
|
||||
chat_count,
|
||||
total_input_tokens,
|
||||
total_output_tokens,
|
||||
total_cache_read_tokens,
|
||||
total_cache_creation_tokens,
|
||||
total_runtime_ms,
|
||||
COUNT(*) OVER()::bigint AS total_count
|
||||
FROM
|
||||
chat_cost_users
|
||||
ORDER BY
|
||||
total_cost_micros DESC,
|
||||
username ASC
|
||||
LIMIT
|
||||
$2::int
|
||||
OFFSET
|
||||
$1::int
|
||||
`
|
||||
|
||||
type GetChatCostPerUserParams struct {
|
||||
PageOffset int32 `db:"page_offset" json:"page_offset"`
|
||||
PageLimit int32 `db:"page_limit" json:"page_limit"`
|
||||
StartDate time.Time `db:"start_date" json:"start_date"`
|
||||
EndDate time.Time `db:"end_date" json:"end_date"`
|
||||
Username string `db:"username" json:"username"`
|
||||
}
|
||||
|
||||
type GetChatCostPerUserRow struct {
|
||||
UserID uuid.UUID `db:"user_id" json:"user_id"`
|
||||
Username string `db:"username" json:"username"`
|
||||
Name string `db:"name" json:"name"`
|
||||
AvatarURL string `db:"avatar_url" json:"avatar_url"`
|
||||
TotalCostMicros int64 `db:"total_cost_micros" json:"total_cost_micros"`
|
||||
MessageCount int64 `db:"message_count" json:"message_count"`
|
||||
ChatCount int64 `db:"chat_count" json:"chat_count"`
|
||||
TotalInputTokens int64 `db:"total_input_tokens" json:"total_input_tokens"`
|
||||
TotalOutputTokens int64 `db:"total_output_tokens" json:"total_output_tokens"`
|
||||
TotalCacheReadTokens int64 `db:"total_cache_read_tokens" json:"total_cache_read_tokens"`
|
||||
TotalCacheCreationTokens int64 `db:"total_cache_creation_tokens" json:"total_cache_creation_tokens"`
|
||||
TotalRuntimeMs int64 `db:"total_runtime_ms" json:"total_runtime_ms"`
|
||||
TotalCount int64 `db:"total_count" json:"total_count"`
|
||||
}
|
||||
|
||||
// Deployment-wide per-user cost rollup within a date range.
|
||||
// Only counts assistant-role messages.
|
||||
func (q *sqlQuerier) GetChatCostPerUser(ctx context.Context, arg GetChatCostPerUserParams) ([]GetChatCostPerUserRow, error) {
|
||||
rows, err := q.db.QueryContext(ctx, getChatCostPerUser,
|
||||
arg.PageOffset,
|
||||
arg.PageLimit,
|
||||
arg.StartDate,
|
||||
arg.EndDate,
|
||||
arg.Username,
|
||||
)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer rows.Close()
|
||||
var items []GetChatCostPerUserRow
|
||||
for rows.Next() {
|
||||
var i GetChatCostPerUserRow
|
||||
if err := rows.Scan(
|
||||
&i.UserID,
|
||||
&i.Username,
|
||||
&i.Name,
|
||||
&i.AvatarURL,
|
||||
&i.TotalCostMicros,
|
||||
&i.MessageCount,
|
||||
&i.ChatCount,
|
||||
&i.TotalInputTokens,
|
||||
&i.TotalOutputTokens,
|
||||
&i.TotalCacheReadTokens,
|
||||
&i.TotalCacheCreationTokens,
|
||||
&i.TotalRuntimeMs,
|
||||
&i.TotalCount,
|
||||
); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
items = append(items, i)
|
||||
}
|
||||
if err := rows.Close(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if err := rows.Err(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return items, nil
|
||||
}
|
||||
|
||||
const getChatCostSummary = `-- name: GetChatCostSummary :one
|
||||
SELECT
|
||||
COALESCE(SUM(cm.total_cost_micros), 0)::bigint AS total_cost_micros,
|
||||
COUNT(*) FILTER (
|
||||
WHERE cm.total_cost_micros IS NOT NULL
|
||||
)::bigint AS priced_message_count,
|
||||
COUNT(*) FILTER (
|
||||
WHERE cm.total_cost_micros IS NULL
|
||||
AND (
|
||||
cm.input_tokens IS NOT NULL
|
||||
OR cm.output_tokens IS NOT NULL
|
||||
OR cm.reasoning_tokens IS NOT NULL
|
||||
OR cm.cache_creation_tokens IS NOT NULL
|
||||
OR cm.cache_read_tokens IS NOT NULL
|
||||
)
|
||||
)::bigint AS unpriced_messages_having_usage_count,
|
||||
COALESCE(SUM(cm.input_tokens), 0)::bigint AS total_input_tokens,
|
||||
COALESCE(SUM(cm.output_tokens), 0)::bigint AS total_output_tokens,
|
||||
COALESCE(SUM(cm.cache_read_tokens), 0)::bigint AS total_cache_read_tokens,
|
||||
COALESCE(SUM(cm.cache_creation_tokens), 0)::bigint AS total_cache_creation_tokens,
|
||||
COALESCE(SUM(cm.runtime_ms), 0)::bigint AS total_runtime_ms
|
||||
FROM
|
||||
chat_messages cm
|
||||
JOIN
|
||||
chats c ON c.id = cm.chat_id
|
||||
WHERE
|
||||
c.owner_id = $1::uuid
|
||||
AND cm.role = 'assistant'
|
||||
AND cm.created_at >= $2::timestamptz
|
||||
AND cm.created_at < $3::timestamptz
|
||||
`
|
||||
|
||||
type GetChatCostSummaryParams struct {
|
||||
OwnerID uuid.UUID `db:"owner_id" json:"owner_id"`
|
||||
StartDate time.Time `db:"start_date" json:"start_date"`
|
||||
EndDate time.Time `db:"end_date" json:"end_date"`
|
||||
}
|
||||
|
||||
type GetChatCostSummaryRow struct {
|
||||
TotalCostMicros int64 `db:"total_cost_micros" json:"total_cost_micros"`
|
||||
PricedMessageCount int64 `db:"priced_message_count" json:"priced_message_count"`
|
||||
UnpricedMessagesHavingUsageCount int64 `db:"unpriced_messages_having_usage_count" json:"unpriced_messages_having_usage_count"`
|
||||
TotalInputTokens int64 `db:"total_input_tokens" json:"total_input_tokens"`
|
||||
TotalOutputTokens int64 `db:"total_output_tokens" json:"total_output_tokens"`
|
||||
TotalCacheReadTokens int64 `db:"total_cache_read_tokens" json:"total_cache_read_tokens"`
|
||||
TotalCacheCreationTokens int64 `db:"total_cache_creation_tokens" json:"total_cache_creation_tokens"`
|
||||
TotalRuntimeMs int64 `db:"total_runtime_ms" json:"total_runtime_ms"`
|
||||
}
|
||||
|
||||
// Aggregate cost summary for a single user within a date range.
|
||||
// Only counts assistant-role messages.
|
||||
func (q *sqlQuerier) GetChatCostSummary(ctx context.Context, arg GetChatCostSummaryParams) (GetChatCostSummaryRow, error) {
|
||||
row := q.db.QueryRowContext(ctx, getChatCostSummary, arg.OwnerID, arg.StartDate, arg.EndDate)
|
||||
var i GetChatCostSummaryRow
|
||||
err := row.Scan(
|
||||
&i.TotalCostMicros,
|
||||
&i.PricedMessageCount,
|
||||
&i.UnpricedMessagesHavingUsageCount,
|
||||
&i.TotalInputTokens,
|
||||
&i.TotalOutputTokens,
|
||||
&i.TotalCacheReadTokens,
|
||||
&i.TotalCacheCreationTokens,
|
||||
&i.TotalRuntimeMs,
|
||||
)
|
||||
return i, err
|
||||
}
|
||||
|
||||
const getChatDiffStatusByChatID = `-- name: GetChatDiffStatusByChatID :one
|
||||
SELECT
|
||||
chat_id, url, pull_request_state, changes_requested, additions, deletions, changed_files, refreshed_at, stale_at, created_at, updated_at, git_branch, git_remote_origin, pull_request_title, pull_request_draft, author_login, author_avatar_url, base_branch, pr_number, commits, approved, reviewer_count, head_branch
|
||||
@@ -8340,7 +7924,6 @@ SELECT
|
||||
COALESCE(SUM(cm.reasoning_tokens), 0)::bigint AS total_reasoning_tokens,
|
||||
COALESCE(SUM(cm.cache_creation_tokens), 0)::bigint AS total_cache_creation_tokens,
|
||||
COALESCE(SUM(cm.cache_read_tokens), 0)::bigint AS total_cache_read_tokens,
|
||||
COALESCE(SUM(cm.total_cost_micros), 0)::bigint AS total_cost_micros,
|
||||
COALESCE(SUM(cm.runtime_ms), 0)::bigint AS total_runtime_ms,
|
||||
COUNT(DISTINCT cm.model_config_id)::bigint AS distinct_model_count,
|
||||
COUNT(*) FILTER (WHERE cm.compressed)::bigint AS compressed_message_count
|
||||
@@ -8362,7 +7945,6 @@ type GetChatMessageSummariesPerChatRow struct {
|
||||
TotalReasoningTokens int64 `db:"total_reasoning_tokens" json:"total_reasoning_tokens"`
|
||||
TotalCacheCreationTokens int64 `db:"total_cache_creation_tokens" json:"total_cache_creation_tokens"`
|
||||
TotalCacheReadTokens int64 `db:"total_cache_read_tokens" json:"total_cache_read_tokens"`
|
||||
TotalCostMicros int64 `db:"total_cost_micros" json:"total_cost_micros"`
|
||||
TotalRuntimeMs int64 `db:"total_runtime_ms" json:"total_runtime_ms"`
|
||||
DistinctModelCount int64 `db:"distinct_model_count" json:"distinct_model_count"`
|
||||
CompressedMessageCount int64 `db:"compressed_message_count" json:"compressed_message_count"`
|
||||
@@ -8392,7 +7974,6 @@ func (q *sqlQuerier) GetChatMessageSummariesPerChat(ctx context.Context, created
|
||||
&i.TotalReasoningTokens,
|
||||
&i.TotalCacheCreationTokens,
|
||||
&i.TotalCacheReadTokens,
|
||||
&i.TotalCostMicros,
|
||||
&i.TotalRuntimeMs,
|
||||
&i.DistinctModelCount,
|
||||
&i.CompressedMessageCount,
|
||||
@@ -8855,67 +8436,6 @@ func (q *sqlQuerier) GetChatModelConfigsForTelemetry(ctx context.Context) ([]Get
|
||||
return items, nil
|
||||
}
|
||||
|
||||
const getChatModelUsageCostByChatID = `-- name: GetChatModelUsageCostByChatID :one
|
||||
WITH RECURSIVE target AS (
|
||||
SELECT $1::uuid AS chat_id
|
||||
), subtree AS (
|
||||
SELECT chat_id AS id FROM target
|
||||
UNION ALL
|
||||
SELECT c.id
|
||||
FROM chats c
|
||||
JOIN subtree s ON c.parent_chat_id = s.id
|
||||
), costs AS (
|
||||
SELECT
|
||||
COALESCE(SUM(cm.total_cost_micros), 0)::bigint AS total_cost_micros,
|
||||
COUNT(*) FILTER (
|
||||
WHERE cm.total_cost_micros IS NOT NULL
|
||||
)::bigint AS priced_message_count,
|
||||
COUNT(*) FILTER (
|
||||
WHERE cm.total_cost_micros IS NULL
|
||||
AND (
|
||||
cm.input_tokens IS NOT NULL
|
||||
OR cm.output_tokens IS NOT NULL
|
||||
OR cm.reasoning_tokens IS NOT NULL
|
||||
OR cm.cache_creation_tokens IS NOT NULL
|
||||
OR cm.cache_read_tokens IS NOT NULL
|
||||
)
|
||||
)::bigint AS unpriced_messages_having_usage_count
|
||||
FROM chat_messages cm
|
||||
JOIN subtree s ON s.id = cm.chat_id
|
||||
WHERE cm.role = 'assistant'
|
||||
)
|
||||
SELECT
|
||||
t.chat_id,
|
||||
costs.total_cost_micros,
|
||||
costs.priced_message_count,
|
||||
costs.unpriced_messages_having_usage_count
|
||||
FROM target t
|
||||
CROSS JOIN costs
|
||||
`
|
||||
|
||||
type GetChatModelUsageCostByChatIDRow struct {
|
||||
ChatID uuid.UUID `db:"chat_id" json:"chat_id"`
|
||||
TotalCostMicros int64 `db:"total_cost_micros" json:"total_cost_micros"`
|
||||
PricedMessageCount int64 `db:"priced_message_count" json:"priced_message_count"`
|
||||
UnpricedMessagesHavingUsageCount int64 `db:"unpriced_messages_having_usage_count" json:"unpriced_messages_having_usage_count"`
|
||||
}
|
||||
|
||||
// Assistant-message cost rolled up over the requested chat's subtree: the
|
||||
// chat itself plus every descendant reachable through parent_chat_id. A
|
||||
// root chat therefore reports its whole tree, while a subagent chat
|
||||
// reports only its own spend plus any nested subagents it spawned.
|
||||
func (q *sqlQuerier) GetChatModelUsageCostByChatID(ctx context.Context, chatID uuid.UUID) (GetChatModelUsageCostByChatIDRow, error) {
|
||||
row := q.db.QueryRowContext(ctx, getChatModelUsageCostByChatID, chatID)
|
||||
var i GetChatModelUsageCostByChatIDRow
|
||||
err := row.Scan(
|
||||
&i.ChatID,
|
||||
&i.TotalCostMicros,
|
||||
&i.PricedMessageCount,
|
||||
&i.UnpricedMessagesHavingUsageCount,
|
||||
)
|
||||
return i, err
|
||||
}
|
||||
|
||||
const getChatQueuedMessageByID = `-- name: GetChatQueuedMessageByID :one
|
||||
SELECT id, chat_id, content, created_at, model_config_id, position, created_by, reasoning_effort FROM chat_queued_messages
|
||||
WHERE id = $1::bigint AND chat_id = $2::uuid
|
||||
@@ -10645,7 +10165,6 @@ inserted AS (
|
||||
cache_read_tokens,
|
||||
context_limit,
|
||||
compressed,
|
||||
total_cost_micros,
|
||||
runtime_ms
|
||||
)
|
||||
SELECT
|
||||
@@ -10666,8 +10185,7 @@ inserted AS (
|
||||
NULLIF(($14::bigint[])[allocated.ord], 0),
|
||||
NULLIF(($15::bigint[])[allocated.ord], 0),
|
||||
($16::boolean[])[allocated.ord],
|
||||
NULLIF(($17::bigint[])[allocated.ord], 0),
|
||||
NULLIF(($18::bigint[])[allocated.ord], 0)
|
||||
NULLIF(($17::bigint[])[allocated.ord], 0)
|
||||
FROM allocated
|
||||
RETURNING id, chat_id, model_config_id, created_at, role, content, visibility, input_tokens, output_tokens, total_tokens, reasoning_tokens, cache_creation_tokens, cache_read_tokens, context_limit, compressed, created_by, content_version, total_cost_micros, runtime_ms, deleted, provider_response_id, revision, reasoning_effort, search_tsv
|
||||
)
|
||||
@@ -10693,7 +10211,6 @@ type InsertChatMessagesParams struct {
|
||||
CacheReadTokens []int64 `db:"cache_read_tokens" json:"cache_read_tokens"`
|
||||
ContextLimit []int64 `db:"context_limit" json:"context_limit"`
|
||||
Compressed []bool `db:"compressed" json:"compressed"`
|
||||
TotalCostMicros []int64 `db:"total_cost_micros" json:"total_cost_micros"`
|
||||
RuntimeMs []int64 `db:"runtime_ms" json:"runtime_ms"`
|
||||
}
|
||||
|
||||
@@ -10745,7 +10262,6 @@ func (q *sqlQuerier) InsertChatMessages(ctx context.Context, arg InsertChatMessa
|
||||
pq.Array(arg.CacheReadTokens),
|
||||
pq.Array(arg.ContextLimit),
|
||||
pq.Array(arg.Compressed),
|
||||
pq.Array(arg.TotalCostMicros),
|
||||
pq.Array(arg.RuntimeMs),
|
||||
)
|
||||
if err != nil {
|
||||
|
||||
@@ -941,7 +941,6 @@ inserted AS (
|
||||
cache_read_tokens,
|
||||
context_limit,
|
||||
compressed,
|
||||
total_cost_micros,
|
||||
runtime_ms
|
||||
)
|
||||
SELECT
|
||||
@@ -962,7 +961,6 @@ inserted AS (
|
||||
NULLIF((@cache_read_tokens::bigint[])[allocated.ord], 0),
|
||||
NULLIF((@context_limit::bigint[])[allocated.ord], 0),
|
||||
(@compressed::boolean[])[allocated.ord],
|
||||
NULLIF((@total_cost_micros::bigint[])[allocated.ord], 0),
|
||||
NULLIF((@runtime_ms::bigint[])[allocated.ord], 0)
|
||||
FROM allocated
|
||||
RETURNING *
|
||||
@@ -2211,229 +2209,6 @@ SELECT
|
||||
COUNT(*) FILTER (WHERE pull_request_state = 'closed')::bigint AS closed
|
||||
FROM deduped;
|
||||
|
||||
-- name: GetChatCostSummary :one
|
||||
-- Aggregate cost summary for a single user within a date range.
|
||||
-- Only counts assistant-role messages.
|
||||
SELECT
|
||||
COALESCE(SUM(cm.total_cost_micros), 0)::bigint AS total_cost_micros,
|
||||
COUNT(*) FILTER (
|
||||
WHERE cm.total_cost_micros IS NOT NULL
|
||||
)::bigint AS priced_message_count,
|
||||
COUNT(*) FILTER (
|
||||
WHERE cm.total_cost_micros IS NULL
|
||||
AND (
|
||||
cm.input_tokens IS NOT NULL
|
||||
OR cm.output_tokens IS NOT NULL
|
||||
OR cm.reasoning_tokens IS NOT NULL
|
||||
OR cm.cache_creation_tokens IS NOT NULL
|
||||
OR cm.cache_read_tokens IS NOT NULL
|
||||
)
|
||||
)::bigint AS unpriced_messages_having_usage_count,
|
||||
COALESCE(SUM(cm.input_tokens), 0)::bigint AS total_input_tokens,
|
||||
COALESCE(SUM(cm.output_tokens), 0)::bigint AS total_output_tokens,
|
||||
COALESCE(SUM(cm.cache_read_tokens), 0)::bigint AS total_cache_read_tokens,
|
||||
COALESCE(SUM(cm.cache_creation_tokens), 0)::bigint AS total_cache_creation_tokens,
|
||||
COALESCE(SUM(cm.runtime_ms), 0)::bigint AS total_runtime_ms
|
||||
FROM
|
||||
chat_messages cm
|
||||
JOIN
|
||||
chats c ON c.id = cm.chat_id
|
||||
WHERE
|
||||
c.owner_id = @owner_id::uuid
|
||||
AND cm.role = 'assistant'
|
||||
AND cm.created_at >= @start_date::timestamptz
|
||||
AND cm.created_at < @end_date::timestamptz;
|
||||
|
||||
-- name: GetChatCostPerModel :many
|
||||
-- Per-model cost breakdown for a single user within a date range.
|
||||
-- Only counts assistant-role messages that have a model_config_id.
|
||||
SELECT
|
||||
cmc.id AS model_config_id,
|
||||
cmc.display_name,
|
||||
COALESCE(ap.type::text, '')::text AS provider,
|
||||
cmc.model,
|
||||
COALESCE(SUM(cm.total_cost_micros), 0)::bigint AS total_cost_micros,
|
||||
COUNT(*) FILTER (
|
||||
WHERE cm.input_tokens IS NOT NULL
|
||||
OR cm.output_tokens IS NOT NULL
|
||||
OR cm.reasoning_tokens IS NOT NULL
|
||||
OR cm.cache_creation_tokens IS NOT NULL
|
||||
OR cm.cache_read_tokens IS NOT NULL
|
||||
)::bigint AS message_count,
|
||||
COALESCE(SUM(cm.input_tokens), 0)::bigint AS total_input_tokens,
|
||||
COALESCE(SUM(cm.output_tokens), 0)::bigint AS total_output_tokens,
|
||||
COALESCE(SUM(cm.cache_read_tokens), 0)::bigint AS total_cache_read_tokens,
|
||||
COALESCE(SUM(cm.cache_creation_tokens), 0)::bigint AS total_cache_creation_tokens,
|
||||
COALESCE(SUM(cm.runtime_ms), 0)::bigint AS total_runtime_ms
|
||||
FROM
|
||||
chat_messages cm
|
||||
JOIN
|
||||
chats c ON c.id = cm.chat_id
|
||||
JOIN
|
||||
chat_model_configs cmc ON cmc.id = cm.model_config_id
|
||||
LEFT JOIN
|
||||
ai_providers ap ON ap.id = cmc.ai_provider_id
|
||||
WHERE
|
||||
c.owner_id = @owner_id::uuid
|
||||
AND cm.role = 'assistant'
|
||||
AND cm.created_at >= @start_date::timestamptz
|
||||
AND cm.created_at < @end_date::timestamptz
|
||||
GROUP BY
|
||||
cmc.id, cmc.display_name, ap.type, cmc.model
|
||||
ORDER BY
|
||||
total_cost_micros DESC;
|
||||
|
||||
-- name: GetChatCostPerChat :many
|
||||
-- Per-root-chat cost breakdown for a single user within a date range.
|
||||
-- Groups by root_chat_id so forked chats roll up under their root.
|
||||
-- Only counts assistant-role messages.
|
||||
WITH chat_costs AS (
|
||||
SELECT
|
||||
COALESCE(c.root_chat_id, c.id) AS root_chat_id,
|
||||
COALESCE(SUM(cm.total_cost_micros), 0)::bigint AS total_cost_micros,
|
||||
COUNT(*) FILTER (
|
||||
WHERE cm.input_tokens IS NOT NULL
|
||||
OR cm.output_tokens IS NOT NULL
|
||||
OR cm.reasoning_tokens IS NOT NULL
|
||||
OR cm.cache_creation_tokens IS NOT NULL
|
||||
OR cm.cache_read_tokens IS NOT NULL
|
||||
)::bigint AS message_count,
|
||||
COALESCE(SUM(cm.input_tokens), 0)::bigint AS total_input_tokens,
|
||||
COALESCE(SUM(cm.output_tokens), 0)::bigint AS total_output_tokens,
|
||||
COALESCE(SUM(cm.cache_read_tokens), 0)::bigint AS total_cache_read_tokens,
|
||||
COALESCE(SUM(cm.cache_creation_tokens), 0)::bigint AS total_cache_creation_tokens,
|
||||
COALESCE(SUM(cm.runtime_ms), 0)::bigint AS total_runtime_ms
|
||||
FROM chat_messages cm
|
||||
JOIN chats c ON c.id = cm.chat_id
|
||||
WHERE c.owner_id = @owner_id::uuid
|
||||
AND cm.role = 'assistant'
|
||||
AND cm.created_at >= @start_date::timestamptz
|
||||
AND cm.created_at < @end_date::timestamptz
|
||||
GROUP BY COALESCE(c.root_chat_id, c.id)
|
||||
)
|
||||
SELECT
|
||||
cc.root_chat_id,
|
||||
COALESCE(rc.title, '') AS chat_title,
|
||||
cc.total_cost_micros,
|
||||
cc.message_count,
|
||||
cc.total_input_tokens,
|
||||
cc.total_output_tokens,
|
||||
cc.total_cache_read_tokens,
|
||||
cc.total_cache_creation_tokens,
|
||||
cc.total_runtime_ms
|
||||
FROM chat_costs cc
|
||||
LEFT JOIN chats rc ON rc.id = cc.root_chat_id
|
||||
ORDER BY cc.total_cost_micros DESC;
|
||||
|
||||
-- name: GetChatModelUsageCostByChatID :one
|
||||
-- Assistant-message cost rolled up over the requested chat's subtree: the
|
||||
-- chat itself plus every descendant reachable through parent_chat_id. A
|
||||
-- root chat therefore reports its whole tree, while a subagent chat
|
||||
-- reports only its own spend plus any nested subagents it spawned.
|
||||
WITH RECURSIVE target AS (
|
||||
SELECT @chat_id::uuid AS chat_id
|
||||
), subtree AS (
|
||||
SELECT chat_id AS id FROM target
|
||||
UNION ALL
|
||||
SELECT c.id
|
||||
FROM chats c
|
||||
JOIN subtree s ON c.parent_chat_id = s.id
|
||||
), costs AS (
|
||||
SELECT
|
||||
COALESCE(SUM(cm.total_cost_micros), 0)::bigint AS total_cost_micros,
|
||||
COUNT(*) FILTER (
|
||||
WHERE cm.total_cost_micros IS NOT NULL
|
||||
)::bigint AS priced_message_count,
|
||||
COUNT(*) FILTER (
|
||||
WHERE cm.total_cost_micros IS NULL
|
||||
AND (
|
||||
cm.input_tokens IS NOT NULL
|
||||
OR cm.output_tokens IS NOT NULL
|
||||
OR cm.reasoning_tokens IS NOT NULL
|
||||
OR cm.cache_creation_tokens IS NOT NULL
|
||||
OR cm.cache_read_tokens IS NOT NULL
|
||||
)
|
||||
)::bigint AS unpriced_messages_having_usage_count
|
||||
FROM chat_messages cm
|
||||
JOIN subtree s ON s.id = cm.chat_id
|
||||
WHERE cm.role = 'assistant'
|
||||
)
|
||||
SELECT
|
||||
t.chat_id,
|
||||
costs.total_cost_micros,
|
||||
costs.priced_message_count,
|
||||
costs.unpriced_messages_having_usage_count
|
||||
FROM target t
|
||||
CROSS JOIN costs;
|
||||
|
||||
-- name: GetChatCostPerUser :many
|
||||
-- Deployment-wide per-user cost rollup within a date range.
|
||||
-- Only counts assistant-role messages.
|
||||
WITH chat_cost_users AS (
|
||||
SELECT
|
||||
c.owner_id AS user_id,
|
||||
u.username,
|
||||
u.name,
|
||||
u.avatar_url,
|
||||
COALESCE(SUM(cm.total_cost_micros), 0)::bigint AS total_cost_micros,
|
||||
COUNT(*) FILTER (
|
||||
WHERE cm.input_tokens IS NOT NULL
|
||||
OR cm.output_tokens IS NOT NULL
|
||||
OR cm.reasoning_tokens IS NOT NULL
|
||||
OR cm.cache_creation_tokens IS NOT NULL
|
||||
OR cm.cache_read_tokens IS NOT NULL
|
||||
)::bigint AS message_count,
|
||||
COUNT(DISTINCT COALESCE(c.root_chat_id, c.id))::bigint AS chat_count,
|
||||
COALESCE(SUM(cm.input_tokens), 0)::bigint AS total_input_tokens,
|
||||
COALESCE(SUM(cm.output_tokens), 0)::bigint AS total_output_tokens,
|
||||
COALESCE(SUM(cm.cache_read_tokens), 0)::bigint AS total_cache_read_tokens,
|
||||
COALESCE(SUM(cm.cache_creation_tokens), 0)::bigint AS total_cache_creation_tokens,
|
||||
COALESCE(SUM(cm.runtime_ms), 0)::bigint AS total_runtime_ms
|
||||
FROM
|
||||
chat_messages cm
|
||||
JOIN
|
||||
chats c ON c.id = cm.chat_id
|
||||
JOIN
|
||||
users u ON u.id = c.owner_id
|
||||
WHERE
|
||||
cm.role = 'assistant'
|
||||
AND cm.created_at >= @start_date::timestamptz
|
||||
AND cm.created_at < @end_date::timestamptz
|
||||
AND (
|
||||
@username::text = ''
|
||||
OR u.username ILIKE '%' || @username::text || '%'
|
||||
OR u.name ILIKE '%' || @username::text || '%'
|
||||
)
|
||||
GROUP BY
|
||||
c.owner_id,
|
||||
u.username,
|
||||
u.name,
|
||||
u.avatar_url
|
||||
)
|
||||
SELECT
|
||||
user_id,
|
||||
username,
|
||||
name,
|
||||
avatar_url,
|
||||
total_cost_micros,
|
||||
message_count,
|
||||
chat_count,
|
||||
total_input_tokens,
|
||||
total_output_tokens,
|
||||
total_cache_read_tokens,
|
||||
total_cache_creation_tokens,
|
||||
total_runtime_ms,
|
||||
COUNT(*) OVER()::bigint AS total_count
|
||||
FROM
|
||||
chat_cost_users
|
||||
ORDER BY
|
||||
total_cost_micros DESC,
|
||||
username ASC
|
||||
LIMIT
|
||||
sqlc.arg('page_limit')::int
|
||||
OFFSET
|
||||
sqlc.arg('page_offset')::int;
|
||||
|
||||
-- name: GetTotalChatMessageRuntimeMsInRange :one
|
||||
-- Computes hb_agent_runtime_v1 usage event payloads. Deliberately includes
|
||||
-- soft-deleted messages and messages from all chats.
|
||||
@@ -2443,22 +2218,6 @@ WHERE cm.created_at >= @start_time::timestamptz
|
||||
AND cm.created_at < @end_time::timestamptz
|
||||
AND cm.runtime_ms IS NOT NULL;
|
||||
|
||||
-- name: CountEnabledModelsWithoutPricing :one
|
||||
-- Counts enabled, non-deleted model configs that lack both input and
|
||||
-- output pricing in their JSONB options.cost configuration.
|
||||
SELECT COUNT(*)::bigint AS count
|
||||
FROM chat_model_configs
|
||||
WHERE enabled = TRUE
|
||||
AND deleted = FALSE
|
||||
AND (
|
||||
options->'cost' IS NULL
|
||||
OR options->'cost' = 'null'::jsonb
|
||||
OR (
|
||||
(options->'cost'->>'input_price_per_million_tokens' IS NULL)
|
||||
AND (options->'cost'->>'output_price_per_million_tokens' IS NULL)
|
||||
)
|
||||
);
|
||||
|
||||
-- name: GetChatsByWorkspaceIDs :many
|
||||
SELECT *
|
||||
FROM chats_expanded
|
||||
@@ -2521,7 +2280,6 @@ SELECT
|
||||
COALESCE(SUM(cm.reasoning_tokens), 0)::bigint AS total_reasoning_tokens,
|
||||
COALESCE(SUM(cm.cache_creation_tokens), 0)::bigint AS total_cache_creation_tokens,
|
||||
COALESCE(SUM(cm.cache_read_tokens), 0)::bigint AS total_cache_read_tokens,
|
||||
COALESCE(SUM(cm.total_cost_micros), 0)::bigint AS total_cost_micros,
|
||||
COALESCE(SUM(cm.runtime_ms), 0)::bigint AS total_runtime_ms,
|
||||
COUNT(DISTINCT cm.model_config_id)::bigint AS distinct_model_count,
|
||||
COUNT(*) FILTER (WHERE cm.compressed)::bigint AS compressed_message_count
|
||||
|
||||
@@ -8,7 +8,6 @@ import (
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"math"
|
||||
"mime"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
@@ -21,7 +20,6 @@ import (
|
||||
|
||||
"github.com/go-chi/chi/v5"
|
||||
"github.com/google/uuid"
|
||||
"github.com/shopspring/decimal"
|
||||
"github.com/sqlc-dev/pqtype"
|
||||
"golang.org/x/xerrors"
|
||||
|
||||
@@ -1598,203 +1596,6 @@ func (api *API) listChatModels(rw http.ResponseWriter, r *http.Request) {
|
||||
httpapi.Write(ctx, rw, http.StatusOK, response)
|
||||
}
|
||||
|
||||
func (api *API) chatCostSummary(rw http.ResponseWriter, r *http.Request) {
|
||||
ctx := r.Context()
|
||||
apiKey := httpmw.APIKey(r)
|
||||
|
||||
// Default date range: last 30 days.
|
||||
now := time.Now()
|
||||
defaultStart := now.AddDate(0, 0, -30)
|
||||
|
||||
qp := r.URL.Query()
|
||||
p := httpapi.NewQueryParamParser()
|
||||
startDate := p.Time(qp, defaultStart, "start_date", time.RFC3339)
|
||||
endDate := p.Time(qp, now, "end_date", time.RFC3339)
|
||||
p.ErrorExcessParams(qp)
|
||||
if len(p.Errors) > 0 {
|
||||
httpapi.Write(ctx, rw, http.StatusBadRequest, codersdk.Response{
|
||||
Message: "Invalid query parameters.",
|
||||
Validations: p.Errors,
|
||||
})
|
||||
return
|
||||
}
|
||||
|
||||
targetUser := httpmw.UserParam(r)
|
||||
if targetUser.ID != apiKey.UserID && !api.Authorize(r, policy.ActionRead, rbac.ResourceChat.WithOwner(targetUser.ID.String())) {
|
||||
httpapi.Forbidden(rw)
|
||||
return
|
||||
}
|
||||
|
||||
summary, err := api.Database.GetChatCostSummary(ctx, database.GetChatCostSummaryParams{
|
||||
OwnerID: targetUser.ID,
|
||||
StartDate: startDate,
|
||||
EndDate: endDate,
|
||||
})
|
||||
if err != nil {
|
||||
if dbauthz.IsNotAuthorizedError(err) {
|
||||
httpapi.Forbidden(rw)
|
||||
return
|
||||
}
|
||||
httpapi.InternalServerError(rw, err)
|
||||
return
|
||||
}
|
||||
|
||||
byModel, err := api.Database.GetChatCostPerModel(ctx, database.GetChatCostPerModelParams{
|
||||
OwnerID: targetUser.ID,
|
||||
StartDate: startDate,
|
||||
EndDate: endDate,
|
||||
})
|
||||
if err != nil {
|
||||
if dbauthz.IsNotAuthorizedError(err) {
|
||||
httpapi.Forbidden(rw)
|
||||
return
|
||||
}
|
||||
httpapi.InternalServerError(rw, err)
|
||||
return
|
||||
}
|
||||
|
||||
byChat, err := api.Database.GetChatCostPerChat(ctx, database.GetChatCostPerChatParams{
|
||||
OwnerID: targetUser.ID,
|
||||
StartDate: startDate,
|
||||
EndDate: endDate,
|
||||
})
|
||||
if err != nil {
|
||||
if dbauthz.IsNotAuthorizedError(err) {
|
||||
httpapi.Forbidden(rw)
|
||||
return
|
||||
}
|
||||
httpapi.InternalServerError(rw, err)
|
||||
return
|
||||
}
|
||||
|
||||
modelBreakdowns := make([]codersdk.ChatCostModelBreakdown, 0, len(byModel))
|
||||
for _, model := range byModel {
|
||||
modelBreakdowns = append(modelBreakdowns, convertChatCostModelBreakdown(model))
|
||||
}
|
||||
|
||||
chatBreakdowns := make([]codersdk.ChatCostChatBreakdown, 0, len(byChat))
|
||||
for _, chat := range byChat {
|
||||
chatBreakdowns = append(chatBreakdowns, convertChatCostChatBreakdown(chat))
|
||||
}
|
||||
|
||||
response := codersdk.ChatCostSummary{
|
||||
StartDate: startDate,
|
||||
EndDate: endDate,
|
||||
TotalCostMicros: summary.TotalCostMicros,
|
||||
PricedMessageCount: summary.PricedMessageCount,
|
||||
UnpricedMessagesHavingUsageCount: summary.UnpricedMessagesHavingUsageCount,
|
||||
TotalInputTokens: summary.TotalInputTokens,
|
||||
TotalOutputTokens: summary.TotalOutputTokens,
|
||||
TotalCacheReadTokens: summary.TotalCacheReadTokens,
|
||||
TotalCacheCreationTokens: summary.TotalCacheCreationTokens,
|
||||
TotalRuntimeMs: summary.TotalRuntimeMs,
|
||||
ByModel: modelBreakdowns,
|
||||
ByChat: chatBreakdowns,
|
||||
}
|
||||
httpapi.Write(ctx, rw, http.StatusOK, response)
|
||||
}
|
||||
|
||||
func (api *API) chatCostUsers(rw http.ResponseWriter, r *http.Request) {
|
||||
ctx := r.Context()
|
||||
if !api.Authorize(r, policy.ActionRead, rbac.ResourceChat) {
|
||||
httpapi.Forbidden(rw)
|
||||
return
|
||||
}
|
||||
|
||||
now := time.Now()
|
||||
defaultStart := now.AddDate(0, 0, -30)
|
||||
|
||||
qp := r.URL.Query()
|
||||
p := httpapi.NewQueryParamParser()
|
||||
startDate := p.Time(qp, defaultStart, "start_date", time.RFC3339)
|
||||
endDate := p.Time(qp, now, "end_date", time.RFC3339)
|
||||
username := strings.TrimSpace(p.String(qp, "", "username"))
|
||||
limit := p.Int(qp, 10, "limit")
|
||||
offset := p.Int(qp, 0, "offset")
|
||||
p.ErrorExcessParams(qp)
|
||||
if len(p.Errors) > 0 {
|
||||
httpapi.Write(ctx, rw, http.StatusBadRequest, codersdk.Response{
|
||||
Message: "Invalid query parameters.",
|
||||
Validations: p.Errors,
|
||||
})
|
||||
return
|
||||
}
|
||||
if limit <= 0 {
|
||||
limit = 10
|
||||
}
|
||||
if offset < 0 || offset > math.MaxInt32 || limit > math.MaxInt32 {
|
||||
validations := make([]codersdk.ValidationError, 0, 2)
|
||||
if offset < 0 {
|
||||
validations = append(validations, codersdk.ValidationError{
|
||||
Field: "offset",
|
||||
Detail: "Must be greater than or equal to 0.",
|
||||
})
|
||||
}
|
||||
if offset > math.MaxInt32 {
|
||||
validations = append(validations, codersdk.ValidationError{
|
||||
Field: "offset",
|
||||
Detail: fmt.Sprintf("Must be less than or equal to %d.", math.MaxInt32),
|
||||
})
|
||||
}
|
||||
if limit > math.MaxInt32 {
|
||||
validations = append(validations, codersdk.ValidationError{
|
||||
Field: "limit",
|
||||
Detail: fmt.Sprintf("Must be less than or equal to %d.", math.MaxInt32),
|
||||
})
|
||||
}
|
||||
httpapi.Write(ctx, rw, http.StatusBadRequest, codersdk.Response{
|
||||
Message: "Invalid query parameters.",
|
||||
Validations: validations,
|
||||
})
|
||||
return
|
||||
}
|
||||
|
||||
users, err := api.Database.GetChatCostPerUser(ctx, database.GetChatCostPerUserParams{
|
||||
StartDate: startDate,
|
||||
EndDate: endDate,
|
||||
Username: username,
|
||||
// #nosec G115 - Pagination limits are validated to fit in int32 above.
|
||||
PageLimit: int32(limit),
|
||||
// #nosec G115 - Pagination offsets are validated to fit in int32 above.
|
||||
PageOffset: int32(offset),
|
||||
})
|
||||
if err != nil {
|
||||
httpapi.InternalServerError(rw, err)
|
||||
return
|
||||
}
|
||||
|
||||
rollups := make([]codersdk.ChatCostUserRollup, 0, len(users))
|
||||
count := int64(0)
|
||||
for _, user := range users {
|
||||
count = user.TotalCount
|
||||
rollups = append(rollups, convertChatCostUserRollup(user))
|
||||
}
|
||||
|
||||
if len(users) == 0 && offset > 0 {
|
||||
countUsers, countErr := api.Database.GetChatCostPerUser(ctx, database.GetChatCostPerUserParams{
|
||||
StartDate: startDate,
|
||||
EndDate: endDate,
|
||||
Username: username,
|
||||
PageLimit: 1,
|
||||
PageOffset: 0,
|
||||
})
|
||||
if countErr != nil {
|
||||
httpapi.InternalServerError(rw, countErr)
|
||||
return
|
||||
}
|
||||
if len(countUsers) > 0 {
|
||||
count = countUsers[0].TotalCount
|
||||
}
|
||||
}
|
||||
|
||||
httpapi.Write(ctx, rw, http.StatusOK, codersdk.ChatCostUsersResponse{
|
||||
StartDate: startDate,
|
||||
EndDate: endDate,
|
||||
Count: count,
|
||||
Users: rollups,
|
||||
})
|
||||
}
|
||||
|
||||
// EXPERIMENTAL: this endpoint is experimental and is subject to change.
|
||||
//
|
||||
// @Summary Get chat by ID
|
||||
@@ -6571,57 +6372,6 @@ func (api *API) fetchChatFileMetadata(ctx context.Context, chatID uuid.UUID) []d
|
||||
return rows
|
||||
}
|
||||
|
||||
func convertChatCostModelBreakdown(model database.GetChatCostPerModelRow) codersdk.ChatCostModelBreakdown {
|
||||
displayName := strings.TrimSpace(model.DisplayName)
|
||||
if displayName == "" {
|
||||
displayName = model.Model
|
||||
}
|
||||
return codersdk.ChatCostModelBreakdown{
|
||||
ModelConfigID: model.ModelConfigID,
|
||||
DisplayName: displayName,
|
||||
Provider: model.Provider,
|
||||
Model: model.Model,
|
||||
TotalCostMicros: model.TotalCostMicros,
|
||||
MessageCount: model.MessageCount,
|
||||
TotalInputTokens: model.TotalInputTokens,
|
||||
TotalOutputTokens: model.TotalOutputTokens,
|
||||
TotalCacheReadTokens: model.TotalCacheReadTokens,
|
||||
TotalCacheCreationTokens: model.TotalCacheCreationTokens,
|
||||
TotalRuntimeMs: model.TotalRuntimeMs,
|
||||
}
|
||||
}
|
||||
|
||||
func convertChatCostChatBreakdown(chat database.GetChatCostPerChatRow) codersdk.ChatCostChatBreakdown {
|
||||
return codersdk.ChatCostChatBreakdown{
|
||||
RootChatID: chat.RootChatID,
|
||||
ChatTitle: chat.ChatTitle,
|
||||
TotalCostMicros: chat.TotalCostMicros,
|
||||
MessageCount: chat.MessageCount,
|
||||
TotalInputTokens: chat.TotalInputTokens,
|
||||
TotalOutputTokens: chat.TotalOutputTokens,
|
||||
TotalCacheReadTokens: chat.TotalCacheReadTokens,
|
||||
TotalCacheCreationTokens: chat.TotalCacheCreationTokens,
|
||||
TotalRuntimeMs: chat.TotalRuntimeMs,
|
||||
}
|
||||
}
|
||||
|
||||
func convertChatCostUserRollup(user database.GetChatCostPerUserRow) codersdk.ChatCostUserRollup {
|
||||
return codersdk.ChatCostUserRollup{
|
||||
UserID: user.UserID,
|
||||
Username: user.Username,
|
||||
Name: user.Name,
|
||||
AvatarURL: user.AvatarURL,
|
||||
TotalCostMicros: user.TotalCostMicros,
|
||||
MessageCount: user.MessageCount,
|
||||
ChatCount: user.ChatCount,
|
||||
TotalInputTokens: user.TotalInputTokens,
|
||||
TotalOutputTokens: user.TotalOutputTokens,
|
||||
TotalCacheReadTokens: user.TotalCacheReadTokens,
|
||||
TotalCacheCreationTokens: user.TotalCacheCreationTokens,
|
||||
TotalRuntimeMs: user.TotalRuntimeMs,
|
||||
}
|
||||
}
|
||||
|
||||
func convertChatQueuedMessage(m database.ChatQueuedMessage) codersdk.ChatQueuedMessage {
|
||||
return db2sdk.ChatQueuedMessage(m)
|
||||
}
|
||||
@@ -7579,26 +7329,6 @@ func validateChatModelCallConfig(modelConfig *codersdk.ChatModelCallConfig) erro
|
||||
return nil
|
||||
}
|
||||
|
||||
costConfig := codersdk.ModelCostConfig{}
|
||||
if modelConfig.Cost != nil {
|
||||
costConfig = *modelConfig.Cost
|
||||
}
|
||||
|
||||
pricingFields := []struct {
|
||||
name string
|
||||
value *decimal.Decimal
|
||||
}{
|
||||
{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 := validateNonNegativeDecimalField(field.name, field.value); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
|
||||
if err := validateChatModelReasoningEffortConfig(modelConfig); err != nil {
|
||||
return err
|
||||
}
|
||||
@@ -7641,16 +7371,6 @@ func validateChatModelProviderOptions(options *codersdk.ChatModelProviderOptions
|
||||
return xerrors.Errorf("provider_options.anthropic.thinking_display must be one of summarized, omitted")
|
||||
}
|
||||
|
||||
func validateNonNegativeDecimalField(name string, value *decimal.Decimal) error {
|
||||
if value == nil {
|
||||
return nil
|
||||
}
|
||||
if value.IsNegative() {
|
||||
return xerrors.Errorf("%s must be greater than or equal to zero", name)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func unmarshalChatModelCallConfig(
|
||||
raw json.RawMessage,
|
||||
) *codersdk.ChatModelCallConfig {
|
||||
@@ -7681,7 +7401,6 @@ func isZeroChatModelCallConfig(config *codersdk.ChatModelCallConfig) bool {
|
||||
config.FrequencyPenalty == nil &&
|
||||
config.ReasoningEffort == nil &&
|
||||
isZeroChatModelOpenAIConfig(config.OpenAIConfig) &&
|
||||
isZeroModelCostConfig(config.Cost) &&
|
||||
isZeroChatModelProviderOptions(config.ProviderOptions)
|
||||
}
|
||||
|
||||
@@ -7689,17 +7408,6 @@ func isZeroChatModelOpenAIConfig(config *codersdk.ChatModelOpenAIConfig) bool {
|
||||
return config == nil || config.UseResponsesAPI == nil
|
||||
}
|
||||
|
||||
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
|
||||
|
||||
@@ -419,7 +419,7 @@ func TestSharedReaderStreamChat(t *testing.T) {
|
||||
LastModelConfigID: modelConfig.ID,
|
||||
Title: "shared stream chat",
|
||||
})
|
||||
insertAssistantCostMessage(t, db, chat.ID, modelConfig.ID, 0)
|
||||
insertAssistantMessage(t, db, chat.ID, modelConfig.ID)
|
||||
|
||||
err := client.UpdateChatACL(ctx, chat.ID, codersdk.UpdateChatACL{
|
||||
UserRoles: map[string]codersdk.ChatRole{
|
||||
|
||||
@@ -10,7 +10,6 @@ import (
|
||||
|
||||
"github.com/go-chi/chi/v5"
|
||||
"github.com/google/uuid"
|
||||
"github.com/shopspring/decimal"
|
||||
"github.com/stretchr/testify/require"
|
||||
"go.uber.org/mock/gomock"
|
||||
"golang.org/x/xerrors"
|
||||
@@ -417,7 +416,6 @@ func TestRewriteChatStartWorkspaceManualUpdateResponse(t *testing.T) {
|
||||
func TestIsZeroChatModelCallConfigCoversEveryField(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
costSample := decimal.NewFromInt(3)
|
||||
sampled := codersdk.ChatModelCallConfig{
|
||||
MaxOutputTokens: ptr.Ref(int64(4096)),
|
||||
Temperature: ptr.Ref(0.7),
|
||||
@@ -425,9 +423,6 @@ func TestIsZeroChatModelCallConfigCoversEveryField(t *testing.T) {
|
||||
TopK: ptr.Ref(int64(40)),
|
||||
PresencePenalty: ptr.Ref(0.1),
|
||||
FrequencyPenalty: ptr.Ref(0.2),
|
||||
Cost: &codersdk.ModelCostConfig{
|
||||
InputPricePerMillionTokens: &costSample,
|
||||
},
|
||||
ReasoningEffort: &codersdk.ChatModelReasoningEffortConfig{
|
||||
Default: ptr.Ref("medium"),
|
||||
},
|
||||
|
||||
+5
-536
@@ -20,7 +20,6 @@ import (
|
||||
|
||||
"github.com/google/uuid"
|
||||
"github.com/mark3labs/mcp-go/mcp"
|
||||
"github.com/shopspring/decimal"
|
||||
"github.com/sqlc-dev/pqtype"
|
||||
"github.com/stretchr/testify/require"
|
||||
"golang.org/x/sync/errgroup"
|
||||
@@ -264,12 +263,11 @@ func (s *failNextUpdateChatModelConfigStore) UpdateChatModelConfig(
|
||||
return s.Store.UpdateChatModelConfig(ctx, arg)
|
||||
}
|
||||
|
||||
func insertAssistantCostMessage(
|
||||
func insertAssistantMessage(
|
||||
t *testing.T,
|
||||
db database.Store,
|
||||
chatID uuid.UUID,
|
||||
modelConfigID uuid.UUID,
|
||||
totalCostMicros int64,
|
||||
) {
|
||||
t.Helper()
|
||||
|
||||
@@ -279,11 +277,10 @@ func insertAssistantCostMessage(
|
||||
require.NoError(t, err)
|
||||
|
||||
_ = dbgen.ChatMessage(t, db, database.ChatMessage{
|
||||
ChatID: chatID,
|
||||
ModelConfigID: uuid.NullUUID{UUID: modelConfigID, Valid: true},
|
||||
Role: database.ChatMessageRoleAssistant,
|
||||
Content: assistantContent,
|
||||
TotalCostMicros: sql.NullInt64{Int64: totalCostMicros, Valid: true},
|
||||
ChatID: chatID,
|
||||
ModelConfigID: uuid.NullUUID{UUID: modelConfigID, Valid: true},
|
||||
Role: database.ChatMessageRoleAssistant,
|
||||
Content: assistantContent,
|
||||
})
|
||||
}
|
||||
|
||||
@@ -3878,41 +3875,6 @@ func TestListChatModelConfigs(t *testing.T) {
|
||||
require.Equal(t, enabledConfig.ID, memberConfigs[0].ID)
|
||||
})
|
||||
|
||||
t.Run("DeserializesLegacyPricingJSON", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
ctx := testutil.Context(t, testutil.WaitLong)
|
||||
client, db := newChatClientWithDatabase(t)
|
||||
firstUser := coderdtest.CreateFirstUser(t, client.Client)
|
||||
|
||||
aiProvider := createAIProviderForTest(t, client, "openai", "test-api-key")
|
||||
|
||||
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 := dbgen.ChatModelConfig(t, db, database.ChatModelConfig{
|
||||
AIProviderID: uuid.NullUUID{UUID: aiProvider.ID, Valid: true},
|
||||
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},
|
||||
ContextLimit: 4096,
|
||||
CompressionThreshold: 80,
|
||||
Options: legacyOptions,
|
||||
})
|
||||
|
||||
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: decRef("0.15"),
|
||||
OutputPricePerMillionTokens: decRef("0.6"),
|
||||
CacheReadPricePerMillionTokens: decRef("0.03"),
|
||||
CacheWritePricePerMillionTokens: decRef("0.3"),
|
||||
},
|
||||
})
|
||||
})
|
||||
|
||||
t.Run("SuccessForOrganizationMember", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
@@ -3954,20 +3916,11 @@ func TestCreateChatModelConfig(t *testing.T) {
|
||||
|
||||
contextLimit := int64(4096)
|
||||
isDefault := true
|
||||
pricing := &codersdk.ChatModelCallConfig{
|
||||
Cost: &codersdk.ModelCostConfig{
|
||||
InputPricePerMillionTokens: decRef("0.15"),
|
||||
OutputPricePerMillionTokens: decRef("0.6"),
|
||||
CacheReadPricePerMillionTokens: decRef("0.03"),
|
||||
CacheWritePricePerMillionTokens: decRef("0.3"),
|
||||
},
|
||||
}
|
||||
modelConfig, err := client.CreateChatModelConfig(ctx, codersdk.CreateChatModelConfigRequest{
|
||||
AIProviderID: &aiProvider.ID,
|
||||
Model: "gpt-4o-mini",
|
||||
ContextLimit: &contextLimit,
|
||||
IsDefault: &isDefault,
|
||||
ModelConfig: pricing,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.NotEqual(t, uuid.Nil, modelConfig.ID)
|
||||
@@ -3975,12 +3928,10 @@ 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("ConcurrentCreatesElectSingleDefault", func(t *testing.T) {
|
||||
@@ -4040,35 +3991,6 @@ func TestCreateChatModelConfig(t *testing.T) {
|
||||
require.Equal(t, []uuid.UUID{claimed.ID}, defaults)
|
||||
})
|
||||
|
||||
t.Run("RejectsNegativePricing", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
ctx := testutil.Context(t, testutil.WaitLong)
|
||||
client := newChatClient(t)
|
||||
_ = coderdtest.CreateFirstUser(t, client.Client)
|
||||
|
||||
aiProvider := createAIProviderForTest(t, client, "openai", "test-api-key")
|
||||
|
||||
contextLimit := int64(4096)
|
||||
_, err := client.CreateChatModelConfig(ctx, codersdk.CreateChatModelConfigRequest{
|
||||
AIProviderID: &aiProvider.ID,
|
||||
Model: "gpt-4o-mini",
|
||||
ContextLimit: &contextLimit,
|
||||
ModelConfig: &codersdk.ChatModelCallConfig{
|
||||
Cost: &codersdk.ModelCostConfig{
|
||||
InputPricePerMillionTokens: decRef("-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("ReasoningEffortStored", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
@@ -4358,29 +4280,18 @@ func TestUpdateChatModelConfig(t *testing.T) {
|
||||
modelConfig := createChatModelConfig(t, client)
|
||||
|
||||
contextLimit := int64(8192)
|
||||
pricing := &codersdk.ChatModelCallConfig{
|
||||
Cost: &codersdk.ModelCostConfig{
|
||||
InputPricePerMillionTokens: decRef("0.2"),
|
||||
OutputPricePerMillionTokens: decRef("0.8"),
|
||||
CacheReadPricePerMillionTokens: decRef("0.04"),
|
||||
CacheWritePricePerMillionTokens: decRef("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("UnchangedProviderWithoutAIProviderID", func(t *testing.T) {
|
||||
@@ -4594,30 +4505,6 @@ func TestUpdateChatModelConfig(t *testing.T) {
|
||||
require.True(t, foundForMember)
|
||||
})
|
||||
|
||||
t.Run("RejectsNegativePricing", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
ctx := testutil.Context(t, testutil.WaitLong)
|
||||
client := newChatClient(t)
|
||||
_ = coderdtest.CreateFirstUser(t, client.Client)
|
||||
modelConfig := createChatModelConfig(t, client)
|
||||
|
||||
_, err := client.UpdateChatModelConfig(ctx, modelConfig.ID, codersdk.UpdateChatModelConfigRequest{
|
||||
ModelConfig: &codersdk.ChatModelCallConfig{
|
||||
Cost: &codersdk.ModelCostConfig{
|
||||
OutputPricePerMillionTokens: decRef("-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("UpdateAIProviderID", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
@@ -5314,7 +5201,6 @@ func TestGetChatUserPrompts(t *testing.T) {
|
||||
CacheReadTokens: []int64{0},
|
||||
ContextLimit: []int64{0},
|
||||
Compressed: []bool{false},
|
||||
TotalCostMicros: []int64{0},
|
||||
RuntimeMs: []int64{0},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
@@ -5396,7 +5282,6 @@ func TestGetChatUserPrompts(t *testing.T) {
|
||||
CacheReadTokens: []int64{0},
|
||||
ContextLimit: []int64{0},
|
||||
Compressed: []bool{false},
|
||||
TotalCostMicros: []int64{0},
|
||||
RuntimeMs: []int64{0},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
@@ -5423,7 +5308,6 @@ func TestGetChatUserPrompts(t *testing.T) {
|
||||
CacheReadTokens: []int64{0},
|
||||
ContextLimit: []int64{0},
|
||||
Compressed: []bool{false},
|
||||
TotalCostMicros: []int64{0},
|
||||
RuntimeMs: []int64{0},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
@@ -5618,7 +5502,6 @@ func TestGetChatUserPrompts(t *testing.T) {
|
||||
CacheReadTokens: []int64{0},
|
||||
ContextLimit: []int64{0},
|
||||
Compressed: []bool{false},
|
||||
TotalCostMicros: []int64{0},
|
||||
RuntimeMs: []int64{0},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
@@ -10660,7 +10543,6 @@ func TestPromoteChatQueuedMessage(t *testing.T) {
|
||||
CacheReadTokens: []int64{0},
|
||||
ContextLimit: []int64{0},
|
||||
Compressed: []bool{false},
|
||||
TotalCostMicros: []int64{0},
|
||||
RuntimeMs: []int64{0},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
@@ -11264,194 +11146,6 @@ func TestGetChatFile(t *testing.T) {
|
||||
})
|
||||
}
|
||||
|
||||
type chatCostTestFixture struct {
|
||||
Client *codersdk.ExperimentalClient
|
||||
DB database.Store
|
||||
ModelConfigID uuid.UUID
|
||||
ChatID uuid.UUID
|
||||
EarliestCreatedAt time.Time
|
||||
LatestCreatedAt time.Time
|
||||
}
|
||||
|
||||
// safeOptions returns an explicit time window around the fixture messages to
|
||||
// avoid app-time/database-time boundary flakes in summary tests.
|
||||
func (f chatCostTestFixture) safeOptions() codersdk.ChatCostSummaryOptions {
|
||||
return codersdk.ChatCostSummaryOptions{
|
||||
StartDate: f.EarliestCreatedAt.Add(-time.Minute),
|
||||
EndDate: f.LatestCreatedAt.Add(time.Minute),
|
||||
}
|
||||
}
|
||||
|
||||
func seedChatCostFixture(t *testing.T) chatCostTestFixture {
|
||||
t.Helper()
|
||||
|
||||
client, db := newChatClientWithDatabase(t)
|
||||
firstUser := coderdtest.CreateFirstUser(t, client.Client)
|
||||
modelConfig := createChatModelConfig(t, client)
|
||||
|
||||
chat := dbgen.Chat(t, db, database.Chat{
|
||||
OrganizationID: firstUser.OrganizationID,
|
||||
OwnerID: firstUser.UserID,
|
||||
LastModelConfigID: modelConfig.ID,
|
||||
Title: "test chat",
|
||||
})
|
||||
|
||||
msg1 := dbgen.ChatMessage(t, db, database.ChatMessage{
|
||||
ChatID: chat.ID,
|
||||
ModelConfigID: uuid.NullUUID{UUID: modelConfig.ID, Valid: true},
|
||||
Role: database.ChatMessageRoleAssistant,
|
||||
InputTokens: sql.NullInt64{Int64: 100, Valid: true},
|
||||
OutputTokens: sql.NullInt64{Int64: 50, Valid: true},
|
||||
TotalCostMicros: sql.NullInt64{Int64: 500, Valid: true},
|
||||
RuntimeMs: sql.NullInt64{Int64: 1500, Valid: true},
|
||||
})
|
||||
msg2 := dbgen.ChatMessage(t, db, database.ChatMessage{
|
||||
ChatID: chat.ID,
|
||||
ModelConfigID: uuid.NullUUID{UUID: modelConfig.ID, Valid: true},
|
||||
Role: database.ChatMessageRoleAssistant,
|
||||
InputTokens: sql.NullInt64{Int64: 100, Valid: true},
|
||||
OutputTokens: sql.NullInt64{Int64: 50, Valid: true},
|
||||
TotalCostMicros: sql.NullInt64{Int64: 500, Valid: true},
|
||||
RuntimeMs: sql.NullInt64{Int64: 2500, Valid: true},
|
||||
})
|
||||
results := []database.ChatMessage{msg1, msg2}
|
||||
require.Len(t, results, 2)
|
||||
|
||||
earliestCreatedAt := results[0].CreatedAt
|
||||
latestCreatedAt := results[0].CreatedAt
|
||||
for _, msg := range results {
|
||||
if msg.CreatedAt.Before(earliestCreatedAt) {
|
||||
earliestCreatedAt = msg.CreatedAt
|
||||
}
|
||||
if msg.CreatedAt.After(latestCreatedAt) {
|
||||
latestCreatedAt = msg.CreatedAt
|
||||
}
|
||||
}
|
||||
|
||||
return chatCostTestFixture{
|
||||
Client: client,
|
||||
DB: db,
|
||||
ModelConfigID: modelConfig.ID,
|
||||
ChatID: chat.ID,
|
||||
EarliestCreatedAt: earliestCreatedAt,
|
||||
LatestCreatedAt: latestCreatedAt,
|
||||
}
|
||||
}
|
||||
|
||||
func assertChatCostSummary(t *testing.T, summary codersdk.ChatCostSummary, modelConfigID, chatID uuid.UUID) {
|
||||
t.Helper()
|
||||
|
||||
require.Equal(t, int64(1000), summary.TotalCostMicros)
|
||||
require.Equal(t, int64(2), summary.PricedMessageCount)
|
||||
require.Equal(t, int64(0), summary.UnpricedMessagesHavingUsageCount)
|
||||
require.Equal(t, int64(200), summary.TotalInputTokens)
|
||||
require.Equal(t, int64(100), summary.TotalOutputTokens)
|
||||
require.Equal(t, int64(4000), summary.TotalRuntimeMs)
|
||||
|
||||
require.Len(t, summary.ByModel, 1)
|
||||
require.Equal(t, modelConfigID, summary.ByModel[0].ModelConfigID)
|
||||
require.Equal(t, int64(1000), summary.ByModel[0].TotalCostMicros)
|
||||
require.Equal(t, int64(2), summary.ByModel[0].MessageCount)
|
||||
require.Equal(t, int64(4000), summary.ByModel[0].TotalRuntimeMs)
|
||||
|
||||
require.Len(t, summary.ByChat, 1)
|
||||
require.Equal(t, chatID, summary.ByChat[0].RootChatID)
|
||||
require.Equal(t, int64(1000), summary.ByChat[0].TotalCostMicros)
|
||||
require.Equal(t, int64(2), summary.ByChat[0].MessageCount)
|
||||
require.Equal(t, int64(4000), summary.ByChat[0].TotalRuntimeMs)
|
||||
}
|
||||
|
||||
func TestChatCostSummary(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
t.Run("BasicSummary", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
f := seedChatCostFixture(t)
|
||||
ctx := testutil.Context(t, testutil.WaitLong)
|
||||
|
||||
// Use a window derived from DB timestamps to avoid time boundary flakes.
|
||||
summary, err := f.Client.GetChatCostSummary(ctx, "me", f.safeOptions())
|
||||
require.NoError(t, err)
|
||||
assertChatCostSummary(t, summary, f.ModelConfigID, f.ChatID)
|
||||
})
|
||||
}
|
||||
|
||||
func TestChatCostSummary_AfterModelDeletion(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
f := seedChatCostFixture(t)
|
||||
ctx := testutil.Context(t, testutil.WaitLong)
|
||||
options := f.safeOptions()
|
||||
|
||||
// Baseline: use DB-derived timestamps to avoid time boundary flakes.
|
||||
summary, err := f.Client.GetChatCostSummary(ctx, "me", options)
|
||||
require.NoError(t, err)
|
||||
assertChatCostSummary(t, summary, f.ModelConfigID, f.ChatID)
|
||||
|
||||
// Soft-delete the model config.
|
||||
err = f.Client.DeleteChatModelConfig(ctx, f.ModelConfigID)
|
||||
require.NoError(t, err)
|
||||
|
||||
// Costs must survive the deletion unchanged within the same safe window.
|
||||
summary, err = f.Client.GetChatCostSummary(ctx, "me", options)
|
||||
require.NoError(t, err)
|
||||
assertChatCostSummary(t, summary, f.ModelConfigID, f.ChatID)
|
||||
}
|
||||
|
||||
func TestChatCostSummary_AdminDrilldown(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
client, db := newChatClientWithDatabase(t)
|
||||
firstUser := coderdtest.CreateFirstUser(t, client.Client)
|
||||
memberClientRaw, member := coderdtest.CreateAnotherUser(t, client.Client, firstUser.OrganizationID)
|
||||
memberClient := codersdk.NewExperimentalClient(memberClientRaw)
|
||||
modelConfig := createChatModelConfig(t, client)
|
||||
|
||||
chat := dbgen.Chat(t, db, database.Chat{
|
||||
OrganizationID: firstUser.OrganizationID,
|
||||
OwnerID: member.ID,
|
||||
LastModelConfigID: modelConfig.ID,
|
||||
Title: "member chat",
|
||||
})
|
||||
|
||||
message := dbgen.ChatMessage(t, db, database.ChatMessage{
|
||||
ChatID: chat.ID,
|
||||
ModelConfigID: uuid.NullUUID{UUID: modelConfig.ID, Valid: true},
|
||||
Role: database.ChatMessageRoleAssistant,
|
||||
InputTokens: sql.NullInt64{Int64: 200, Valid: true},
|
||||
OutputTokens: sql.NullInt64{Int64: 100, Valid: true},
|
||||
TotalCostMicros: sql.NullInt64{Int64: 750, Valid: true},
|
||||
})
|
||||
|
||||
options := codersdk.ChatCostSummaryOptions{
|
||||
// Pad the DB-assigned timestamp so the query window cannot race it.
|
||||
StartDate: message.CreatedAt.Add(-time.Minute),
|
||||
EndDate: message.CreatedAt.Add(time.Minute),
|
||||
}
|
||||
|
||||
t.Run("AdminCanDrilldown", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
ctx := testutil.Context(t, testutil.WaitLong)
|
||||
summary, err := client.GetChatCostSummary(ctx, member.ID.String(), options)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, int64(750), summary.TotalCostMicros)
|
||||
require.Equal(t, int64(1), summary.PricedMessageCount)
|
||||
})
|
||||
|
||||
t.Run("MemberCannotDrilldownOtherUser", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
ctx := testutil.Context(t, testutil.WaitLong)
|
||||
_, err := memberClient.GetChatCostSummary(ctx, firstUser.UserID.String(), options)
|
||||
require.Error(t, err)
|
||||
var sdkErr *codersdk.Error
|
||||
require.ErrorAs(t, err, &sdkErr)
|
||||
require.Equal(t, http.StatusNotFound, sdkErr.StatusCode())
|
||||
})
|
||||
}
|
||||
|
||||
// seedChatGatewayRequest records one finished Coder Agents gateway request
|
||||
// under sessionChatID, mirroring aibridged: the session ID is the spawning
|
||||
// chat, and each usage is one provider response within that one request.
|
||||
@@ -11888,231 +11582,6 @@ func TestGetChatCost(t *testing.T) {
|
||||
})
|
||||
}
|
||||
|
||||
func TestChatCostUsers(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
seedCtx := testutil.Context(t, testutil.WaitLong)
|
||||
client, db := newChatClientWithDatabase(t)
|
||||
firstUser := coderdtest.CreateFirstUser(t, client.Client)
|
||||
memberClientRaw, member := coderdtest.CreateAnotherUser(t, client.Client, firstUser.OrganizationID)
|
||||
memberClient := codersdk.NewExperimentalClient(memberClientRaw)
|
||||
firstUserRecord, err := db.GetUserByID(dbauthz.AsSystemRestricted(seedCtx), firstUser.UserID)
|
||||
require.NoError(t, err)
|
||||
modelConfig := createChatModelConfig(t, client)
|
||||
|
||||
adminChat := dbgen.Chat(t, db, database.Chat{
|
||||
OrganizationID: firstUser.OrganizationID,
|
||||
OwnerID: firstUser.UserID,
|
||||
LastModelConfigID: modelConfig.ID,
|
||||
Title: "admin chat",
|
||||
})
|
||||
_ = dbgen.ChatMessage(t, db, database.ChatMessage{
|
||||
ChatID: adminChat.ID,
|
||||
ModelConfigID: uuid.NullUUID{UUID: modelConfig.ID, Valid: true},
|
||||
Role: database.ChatMessageRoleAssistant,
|
||||
InputTokens: sql.NullInt64{Int64: 100, Valid: true},
|
||||
OutputTokens: sql.NullInt64{Int64: 50, Valid: true},
|
||||
TotalCostMicros: sql.NullInt64{Int64: 300, Valid: true},
|
||||
})
|
||||
|
||||
memberChat := dbgen.Chat(t, db, database.Chat{
|
||||
OrganizationID: firstUser.OrganizationID,
|
||||
OwnerID: member.ID,
|
||||
LastModelConfigID: modelConfig.ID,
|
||||
Title: "member chat",
|
||||
})
|
||||
_ = dbgen.ChatMessage(t, db, database.ChatMessage{
|
||||
ChatID: memberChat.ID,
|
||||
ModelConfigID: uuid.NullUUID{UUID: modelConfig.ID, Valid: true},
|
||||
Role: database.ChatMessageRoleAssistant,
|
||||
InputTokens: sql.NullInt64{Int64: 200, Valid: true},
|
||||
OutputTokens: sql.NullInt64{Int64: 100, Valid: true},
|
||||
TotalCostMicros: sql.NullInt64{Int64: 800, Valid: true},
|
||||
})
|
||||
|
||||
t.Run("AdminCanListUsers", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
ctx := testutil.Context(t, testutil.WaitLong)
|
||||
resp, err := client.GetChatCostUsers(ctx, codersdk.ChatCostUsersOptions{})
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, int64(2), resp.Count)
|
||||
require.Len(t, resp.Users, 2)
|
||||
require.Equal(t, member.ID, resp.Users[0].UserID)
|
||||
require.Equal(t, member.Username, resp.Users[0].Username)
|
||||
require.Equal(t, int64(800), resp.Users[0].TotalCostMicros)
|
||||
require.Equal(t, int64(1), resp.Users[0].MessageCount)
|
||||
require.Equal(t, int64(1), resp.Users[0].ChatCount)
|
||||
require.Equal(t, firstUser.UserID, resp.Users[1].UserID)
|
||||
require.Equal(t, firstUserRecord.Username, resp.Users[1].Username)
|
||||
require.Equal(t, int64(300), resp.Users[1].TotalCostMicros)
|
||||
})
|
||||
|
||||
t.Run("AdminCanFilterAndPaginateUsers", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
ctx := testutil.Context(t, testutil.WaitLong)
|
||||
resp, err := client.GetChatCostUsers(ctx, codersdk.ChatCostUsersOptions{
|
||||
Username: member.Username,
|
||||
Pagination: codersdk.Pagination{
|
||||
Limit: 1,
|
||||
Offset: 0,
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, int64(1), resp.Count)
|
||||
require.Len(t, resp.Users, 1)
|
||||
require.Equal(t, member.ID, resp.Users[0].UserID)
|
||||
require.Equal(t, member.Username, resp.Users[0].Username)
|
||||
})
|
||||
|
||||
t.Run("MemberCannotListUsers", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
ctx := testutil.Context(t, testutil.WaitLong)
|
||||
_, err := memberClient.GetChatCostUsers(ctx, codersdk.ChatCostUsersOptions{})
|
||||
require.Error(t, err)
|
||||
var sdkErr *codersdk.Error
|
||||
require.ErrorAs(t, err, &sdkErr)
|
||||
require.Equal(t, http.StatusForbidden, sdkErr.StatusCode())
|
||||
})
|
||||
}
|
||||
|
||||
func TestChatCostSummary_DateRange(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
client, db := newChatClientWithDatabase(t)
|
||||
firstUser := coderdtest.CreateFirstUser(t, client.Client)
|
||||
modelConfig := createChatModelConfig(t, client)
|
||||
|
||||
chat := dbgen.Chat(t, db, database.Chat{
|
||||
OrganizationID: firstUser.OrganizationID,
|
||||
OwnerID: firstUser.UserID,
|
||||
LastModelConfigID: modelConfig.ID,
|
||||
Title: "date range test",
|
||||
})
|
||||
|
||||
_ = dbgen.ChatMessage(t, db, database.ChatMessage{
|
||||
ChatID: chat.ID,
|
||||
ModelConfigID: uuid.NullUUID{UUID: modelConfig.ID, Valid: true},
|
||||
Role: database.ChatMessageRoleAssistant,
|
||||
InputTokens: sql.NullInt64{Int64: 100, Valid: true},
|
||||
OutputTokens: sql.NullInt64{Int64: 50, Valid: true},
|
||||
TotalCostMicros: sql.NullInt64{Int64: 500, Valid: true},
|
||||
})
|
||||
|
||||
now := time.Now()
|
||||
|
||||
t.Run("MessageInRange", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
ctx := testutil.Context(t, testutil.WaitLong)
|
||||
summary, err := client.GetChatCostSummary(ctx, "me", codersdk.ChatCostSummaryOptions{
|
||||
StartDate: now.Add(-time.Hour),
|
||||
EndDate: now.Add(time.Hour),
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, int64(500), summary.TotalCostMicros)
|
||||
require.Equal(t, int64(1), summary.PricedMessageCount)
|
||||
})
|
||||
|
||||
t.Run("MessageOutOfRange", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
ctx := testutil.Context(t, testutil.WaitLong)
|
||||
summary, err := client.GetChatCostSummary(ctx, "me", codersdk.ChatCostSummaryOptions{
|
||||
StartDate: now.Add(time.Hour),
|
||||
EndDate: now.Add(2 * time.Hour),
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, int64(0), summary.TotalCostMicros)
|
||||
require.Equal(t, int64(0), summary.PricedMessageCount)
|
||||
})
|
||||
}
|
||||
|
||||
func TestChatCostSummary_UnpricedMessages(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
ctx := testutil.Context(t, testutil.WaitLong)
|
||||
client, db := newChatClientWithDatabase(t)
|
||||
firstUser := coderdtest.CreateFirstUser(t, client.Client)
|
||||
modelConfig := createChatModelConfig(t, client)
|
||||
|
||||
chat := dbgen.Chat(t, db, database.Chat{
|
||||
OrganizationID: firstUser.OrganizationID,
|
||||
OwnerID: firstUser.UserID,
|
||||
LastModelConfigID: modelConfig.ID,
|
||||
Title: "unpriced test",
|
||||
})
|
||||
|
||||
pricedMessage := dbgen.ChatMessage(t, db, database.ChatMessage{
|
||||
ChatID: chat.ID,
|
||||
ModelConfigID: uuid.NullUUID{UUID: modelConfig.ID, Valid: true},
|
||||
Role: database.ChatMessageRoleAssistant,
|
||||
InputTokens: sql.NullInt64{Int64: 100, Valid: true},
|
||||
OutputTokens: sql.NullInt64{Int64: 50, Valid: true},
|
||||
TotalCostMicros: sql.NullInt64{Int64: 500, Valid: true},
|
||||
})
|
||||
|
||||
unpricedMessage := dbgen.ChatMessage(t, db, database.ChatMessage{
|
||||
ChatID: chat.ID,
|
||||
ModelConfigID: uuid.NullUUID{UUID: modelConfig.ID, Valid: true},
|
||||
Role: database.ChatMessageRoleAssistant,
|
||||
InputTokens: sql.NullInt64{Int64: 200, Valid: true},
|
||||
OutputTokens: sql.NullInt64{Int64: 75, Valid: true},
|
||||
})
|
||||
|
||||
earliestCreatedAt := pricedMessage.CreatedAt
|
||||
latestCreatedAt := pricedMessage.CreatedAt
|
||||
if unpricedMessage.CreatedAt.Before(earliestCreatedAt) {
|
||||
earliestCreatedAt = unpricedMessage.CreatedAt
|
||||
}
|
||||
if unpricedMessage.CreatedAt.After(latestCreatedAt) {
|
||||
latestCreatedAt = unpricedMessage.CreatedAt
|
||||
}
|
||||
options := codersdk.ChatCostSummaryOptions{
|
||||
// Pad the DB-assigned timestamps to avoid time boundary flakes.
|
||||
StartDate: earliestCreatedAt.Add(-time.Minute),
|
||||
EndDate: latestCreatedAt.Add(time.Minute),
|
||||
}
|
||||
|
||||
summary, err := client.GetChatCostSummary(ctx, "me", options)
|
||||
require.NoError(t, err)
|
||||
|
||||
require.Equal(t, int64(500), summary.TotalCostMicros)
|
||||
require.Equal(t, int64(1), summary.PricedMessageCount)
|
||||
require.Equal(t, int64(1), summary.UnpricedMessagesHavingUsageCount)
|
||||
require.Equal(t, int64(300), summary.TotalInputTokens)
|
||||
require.Equal(t, int64(125), summary.TotalOutputTokens)
|
||||
}
|
||||
|
||||
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.True(t, expected.Cost.InputPricePerMillionTokens.Equal(*actual.Cost.InputPricePerMillionTokens))
|
||||
require.True(t, expected.Cost.OutputPricePerMillionTokens.Equal(*actual.Cost.OutputPricePerMillionTokens))
|
||||
require.True(t, expected.Cost.CacheReadPricePerMillionTokens.Equal(*actual.Cost.CacheReadPricePerMillionTokens))
|
||||
require.True(t, expected.Cost.CacheWritePricePerMillionTokens.Equal(*actual.Cost.CacheWritePricePerMillionTokens))
|
||||
}
|
||||
|
||||
func decRef(value string) *decimal.Decimal {
|
||||
d := decimal.RequireFromString(value)
|
||||
return &d
|
||||
}
|
||||
|
||||
func TestWatchChatDesktop(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
|
||||
@@ -2306,7 +2306,6 @@ func ConvertChatMessageSummary(dbRow database.GetChatMessageSummariesPerChatRow)
|
||||
TotalReasoningTokens: dbRow.TotalReasoningTokens,
|
||||
TotalCacheCreationTokens: dbRow.TotalCacheCreationTokens,
|
||||
TotalCacheReadTokens: dbRow.TotalCacheReadTokens,
|
||||
TotalCostMicros: dbRow.TotalCostMicros,
|
||||
TotalRuntimeMs: dbRow.TotalRuntimeMs,
|
||||
DistinctModelCount: dbRow.DistinctModelCount,
|
||||
CompressedMessageCount: dbRow.CompressedMessageCount,
|
||||
@@ -2602,7 +2601,6 @@ type ChatMessageSummary struct {
|
||||
TotalReasoningTokens int64 `json:"total_reasoning_tokens"`
|
||||
TotalCacheCreationTokens int64 `json:"total_cache_creation_tokens"`
|
||||
TotalCacheReadTokens int64 `json:"total_cache_read_tokens"`
|
||||
TotalCostMicros int64 `json:"total_cost_micros"`
|
||||
TotalRuntimeMs int64 `json:"total_runtime_ms"`
|
||||
DistinctModelCount int64 `json:"distinct_model_count"`
|
||||
CompressedMessageCount int64 `json:"compressed_message_count"`
|
||||
|
||||
@@ -1719,7 +1719,6 @@ func TestChatsTelemetry(t *testing.T) {
|
||||
TotalTokens: sql.NullInt64{Int64: 100, Valid: true},
|
||||
CacheCreationTokens: sql.NullInt64{Int64: 50, Valid: true},
|
||||
ContextLimit: sql.NullInt64{Int64: 200000, Valid: true},
|
||||
TotalCostMicros: sql.NullInt64{Int64: 1000, Valid: true},
|
||||
})
|
||||
_ = dbgen.ChatMessage(t, db, database.ChatMessage{
|
||||
ChatID: rootChat.ID,
|
||||
@@ -1732,7 +1731,6 @@ func TestChatsTelemetry(t *testing.T) {
|
||||
ReasoningTokens: sql.NullInt64{Int64: 10, Valid: true},
|
||||
CacheReadTokens: sql.NullInt64{Int64: 25, Valid: true},
|
||||
ContextLimit: sql.NullInt64{Int64: 200000, Valid: true},
|
||||
TotalCostMicros: sql.NullInt64{Int64: 2000, Valid: true},
|
||||
RuntimeMs: sql.NullInt64{Int64: 500, Valid: true},
|
||||
ProviderResponseID: sql.NullString{String: "resp-1", Valid: true},
|
||||
})
|
||||
@@ -1746,7 +1744,6 @@ func TestChatsTelemetry(t *testing.T) {
|
||||
TotalTokens: sql.NullInt64{Int64: 150, Valid: true},
|
||||
CacheCreationTokens: sql.NullInt64{Int64: 30, Valid: true},
|
||||
ContextLimit: sql.NullInt64{Int64: 200000, Valid: true},
|
||||
TotalCostMicros: sql.NullInt64{Int64: 1500, Valid: true},
|
||||
})
|
||||
_ = dbgen.ChatMessage(t, db, database.ChatMessage{
|
||||
ChatID: rootChat.ID,
|
||||
@@ -1759,7 +1756,6 @@ func TestChatsTelemetry(t *testing.T) {
|
||||
ReasoningTokens: sql.NullInt64{Int64: 20, Valid: true},
|
||||
CacheReadTokens: sql.NullInt64{Int64: 40, Valid: true},
|
||||
ContextLimit: sql.NullInt64{Int64: 200000, Valid: true},
|
||||
TotalCostMicros: sql.NullInt64{Int64: 3000, Valid: true},
|
||||
RuntimeMs: sql.NullInt64{Int64: 800, Valid: true},
|
||||
ProviderResponseID: sql.NullString{String: "resp-2", Valid: true},
|
||||
})
|
||||
@@ -1783,7 +1779,6 @@ func TestChatsTelemetry(t *testing.T) {
|
||||
TotalTokens: sql.NullInt64{Int64: 500, Valid: true},
|
||||
CacheCreationTokens: sql.NullInt64{Int64: 100, Valid: true},
|
||||
ContextLimit: sql.NullInt64{Int64: 128000, Valid: true},
|
||||
TotalCostMicros: sql.NullInt64{Int64: 5000, Valid: true},
|
||||
})
|
||||
_ = dbgen.ChatMessage(t, db, database.ChatMessage{
|
||||
ChatID: childChat.ID,
|
||||
@@ -1797,7 +1792,6 @@ func TestChatsTelemetry(t *testing.T) {
|
||||
CacheReadTokens: sql.NullInt64{Int64: 75, Valid: true},
|
||||
ContextLimit: sql.NullInt64{Int64: 128000, Valid: true},
|
||||
Compressed: true,
|
||||
TotalCostMicros: sql.NullInt64{Int64: 8000, Valid: true},
|
||||
RuntimeMs: sql.NullInt64{Int64: 1200, Valid: true},
|
||||
ProviderResponseID: sql.NullString{String: "resp-3", Valid: true},
|
||||
})
|
||||
@@ -1817,7 +1811,6 @@ func TestChatsTelemetry(t *testing.T) {
|
||||
CacheCreationTokens: sql.NullInt64{Int64: 999999, Valid: true},
|
||||
CacheReadTokens: sql.NullInt64{Int64: 999999, Valid: true},
|
||||
ContextLimit: sql.NullInt64{Int64: 200000, Valid: true},
|
||||
TotalCostMicros: sql.NullInt64{Int64: 999999, Valid: true},
|
||||
RuntimeMs: sql.NullInt64{Int64: 999999, Valid: true},
|
||||
})
|
||||
err = db.SoftDeleteChatMessageByID(ctx, poisonMsg.ID)
|
||||
@@ -1893,7 +1886,6 @@ func TestChatsTelemetry(t *testing.T) {
|
||||
assert.Equal(t, int64(30), rootSummary.TotalReasoningTokens) // 0+10+0+20+0
|
||||
assert.Equal(t, int64(80), rootSummary.TotalCacheCreationTokens) // 50+0+30+0+0
|
||||
assert.Equal(t, int64(65), rootSummary.TotalCacheReadTokens) // 0+25+0+40+0
|
||||
assert.Equal(t, int64(7500), rootSummary.TotalCostMicros) // 1000+2000+1500+3000+0
|
||||
assert.Equal(t, int64(1400), rootSummary.TotalRuntimeMs) // 0+500+0+800+100
|
||||
assert.Equal(t, int64(1), rootSummary.DistinctModelCount)
|
||||
assert.Equal(t, int64(0), rootSummary.CompressedMessageCount)
|
||||
@@ -1911,7 +1903,6 @@ func TestChatsTelemetry(t *testing.T) {
|
||||
assert.Equal(t, int64(0), childSummary.SystemMessageCount)
|
||||
assert.Equal(t, int64(100), childSummary.TotalCacheCreationTokens) // 100+0
|
||||
assert.Equal(t, int64(75), childSummary.TotalCacheReadTokens) // 0+75
|
||||
assert.Equal(t, int64(13000), childSummary.TotalCostMicros) // 5000+8000
|
||||
assert.Equal(t, int64(1200), childSummary.TotalRuntimeMs) // 0+1200
|
||||
assert.Equal(t, int64(1), childSummary.DistinctModelCount)
|
||||
assert.Equal(t, int64(1), childSummary.CompressedMessageCount)
|
||||
|
||||
@@ -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