diff --git a/coderd/aibridge/budget/budget.go b/coderd/aibridge/budget/budget.go index b7aa3accba7..aa0e9508bc0 100644 --- a/coderd/aibridge/budget/budget.go +++ b/coderd/aibridge/budget/budget.go @@ -26,6 +26,13 @@ const ( SourceGroup LimitSource = "group" ) +// Store is the subset of database.Store needed to resolve a user's effective +// AI budget. +type Store interface { + GetUserAIBudgetOverride(ctx context.Context, userID uuid.UUID) (database.UserAIBudgetOverride, error) + GetHighestGroupAIBudgetByUser(ctx context.Context, userID uuid.UUID) (database.GetHighestGroupAIBudgetByUserRow, error) +} + // EffectiveBudget is the AI budget that applies to a user after override and // policy resolution. type EffectiveBudget struct { @@ -41,7 +48,7 @@ type EffectiveBudget struct { // return value is false when no budget is configured for the user. A per-user // override wins unconditionally; otherwise the budget is selected from the // user's groups according to policy. -func ResolveUserAIBudget(ctx context.Context, db database.Store, userID uuid.UUID, policy codersdk.AIBudgetPolicy) (EffectiveBudget, bool, error) { +func ResolveUserAIBudget(ctx context.Context, db Store, userID uuid.UUID, policy codersdk.AIBudgetPolicy) (EffectiveBudget, bool, error) { // A per-user override always wins. override, err := db.GetUserAIBudgetOverride(ctx, userID) if err == nil { diff --git a/coderd/aibridgedserver/aibridgedserver.go b/coderd/aibridgedserver/aibridgedserver.go index c9fb665f71b..1690f8bac7e 100644 --- a/coderd/aibridgedserver/aibridgedserver.go +++ b/coderd/aibridgedserver/aibridgedserver.go @@ -67,6 +67,13 @@ type store interface { UpdateAIBridgeInterceptionEnded(ctx context.Context, intcID database.UpdateAIBridgeInterceptionEndedParams) (database.AIBridgeInterception, error) GetAIBridgeInterceptionLineageByToolCallID(ctx context.Context, toolCallID string) (database.GetAIBridgeInterceptionLineageByToolCallIDRow, error) + // Cost-attribution queries, used to snapshot price and effective group on + // each token usage record. + GetAIBridgeInterceptionByID(ctx context.Context, id uuid.UUID) (database.AIBridgeInterception, error) + GetAIModelPriceByProviderModel(ctx context.Context, arg database.GetAIModelPriceByProviderModelParams) (database.AIModelPrice, error) + GetUserAIBudgetOverride(ctx context.Context, userID uuid.UUID) (database.UserAIBudgetOverride, error) + GetHighestGroupAIBudgetByUser(ctx context.Context, userID uuid.UUID) (database.GetHighestGroupAIBudgetByUserRow, error) + // MCPConfigurator-related queries. GetExternalAuthLinksByUserID(ctx context.Context, userID uuid.UUID) ([]database.ExternalAuthLink, error) @@ -87,6 +94,9 @@ type Server struct { coderMCPConfig *proto.MCPServerConfig // may be nil if not available structuredLogging bool aiSeatTracker aiseats.SeatTracker + // budgetPolicy selects the effective group when a user belongs to multiple + // budgeted groups, used for cost attribution on token usage records. + budgetPolicy codersdk.AIBudgetPolicy } func NewServer(lifecycleCtx context.Context, store store, logger slog.Logger, accessURL string, @@ -110,6 +120,7 @@ func NewServer(lifecycleCtx context.Context, store store, logger slog.Logger, ac externalAuthConfigs: eac, structuredLogging: bridgeCfg.StructuredLogging.Value(), aiSeatTracker: aiSeatTracker, + budgetPolicy: codersdk.NewAIBudgetPolicyFromString(bridgeCfg.BudgetPolicy), } if bridgeCfg.InjectCoderMCPTools { @@ -272,6 +283,21 @@ func (s *Server) RecordTokenUsage(ctx context.Context, in *proto.RecordTokenUsag s.logger.Warn(ctx, "failed to marshal aibridge metadata from proto to JSON", slog.F("metadata", in), slog.Error(err)) } + // The interception is always recorded before any of its token usages, + // so it must exist. It carries the provider, model, and initiator needed + // for cost attribution. + intc, err := s.store.GetAIBridgeInterceptionByID(ctx, intcID) + if err != nil { + return nil, xerrors.Errorf("get interception %q: %w", intcID, err) + } + + // Snapshot the effective group, per-token prices and compute cost. A + // missing price row or unbudgeted user yields NULL columns. + cost, err := s.resolveTokenUsageCost(ctx, intc, in) + if err != nil { + return nil, xerrors.Errorf("resolve token usage cost: %w", err) + } + _, err = s.store.InsertAIBridgeTokenUsage(ctx, database.InsertAIBridgeTokenUsageParams{ ID: uuid.New(), InterceptionID: intcID, @@ -282,6 +308,12 @@ func (s *Server) RecordTokenUsage(ctx context.Context, in *proto.RecordTokenUsag CacheWriteInputTokens: in.GetCacheWriteInputTokens(), Metadata: out, CreatedAt: in.GetCreatedAt().AsTime(), + EffectiveGroupID: cost.effectiveGroupID, + InputPriceMicros: cost.inputPriceMicros, + OutputPriceMicros: cost.outputPriceMicros, + CacheReadPriceMicros: cost.cacheReadPriceMicros, + CacheWritePriceMicros: cost.cacheWritePriceMicros, + CostMicros: cost.costMicros, }) if err != nil { return nil, xerrors.Errorf("insert token usage: %w", err) diff --git a/coderd/aibridgedserver/aibridgedserver_test.go b/coderd/aibridgedserver/aibridgedserver_test.go index 2d5ce57e97e..9f06592c8ba 100644 --- a/coderd/aibridgedserver/aibridgedserver_test.go +++ b/coderd/aibridgedserver/aibridgedserver_test.go @@ -13,6 +13,7 @@ import ( "time" "github.com/google/uuid" + "github.com/prometheus/client_golang/prometheus" "github.com/sqlc-dev/pqtype" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" @@ -29,13 +30,16 @@ import ( "github.com/coder/coder/v2/coderd/aibridgedserver" agplaiseats "github.com/coder/coder/v2/coderd/aiseats" "github.com/coder/coder/v2/coderd/apikey" + "github.com/coder/coder/v2/coderd/coderdtest" "github.com/coder/coder/v2/coderd/database" + "github.com/coder/coder/v2/coderd/database/dbauthz" "github.com/coder/coder/v2/coderd/database/dbgen" "github.com/coder/coder/v2/coderd/database/dbmock" "github.com/coder/coder/v2/coderd/database/dbtestutil" "github.com/coder/coder/v2/coderd/database/dbtime" "github.com/coder/coder/v2/coderd/externalauth" codermcp "github.com/coder/coder/v2/coderd/mcp" + "github.com/coder/coder/v2/coderd/rbac" "github.com/coder/coder/v2/coderd/util/ptr" "github.com/coder/coder/v2/codersdk" "github.com/coder/coder/v2/cryptorand" @@ -1172,7 +1176,7 @@ func TestRecordTokenUsage(t *testing.T) { }, []testRecordMethodCase[*proto.RecordTokenUsageRequest]{ { - name: "valid token usage", + name: "valid token usage with null cost", request: &proto.RecordTokenUsageRequest{ InterceptionId: uuid.NewString(), MsgId: "msg_123", @@ -1187,6 +1191,11 @@ func TestRecordTokenUsage(t *testing.T) { interceptionID, err := uuid.Parse(req.GetInterceptionId()) assert.NoError(t, err, "parse interception UUID") + // No budget configured and no price row: tokens recorded + // with NULL cost, prices, and group attribution. + intc := newTestInterception(interceptionID) + expectTokenUsageCostLookups(db, intc, nil, nil, nil) + db.EXPECT().InsertAIBridgeTokenUsage(gomock.Any(), gomock.Cond(func(p database.InsertAIBridgeTokenUsageParams) bool { if !assert.NotEqual(t, uuid.Nil, p.ID, "ID") || !assert.Equal(t, interceptionID, p.InterceptionID, "interception ID") || @@ -1196,7 +1205,13 @@ func TestRecordTokenUsage(t *testing.T) { !assert.Equal(t, req.GetCacheReadInputTokens(), p.CacheReadInputTokens, "cache read input tokens") || !assert.Equal(t, req.GetCacheWriteInputTokens(), p.CacheWriteInputTokens, "cache write input tokens") || !assert.JSONEq(t, metadataJSON, string(p.Metadata), "metadata") || - !assert.WithinDuration(t, req.GetCreatedAt().AsTime(), p.CreatedAt, time.Second, "created at") { + !assert.WithinDuration(t, req.GetCreatedAt().AsTime(), p.CreatedAt, time.Second, "created at") || + !assert.False(t, p.EffectiveGroupID.Valid, "effective group ID null") || + !assert.False(t, p.InputPriceMicros.Valid, "input price null") || + !assert.False(t, p.OutputPriceMicros.Valid, "output price null") || + !assert.False(t, p.CacheReadPriceMicros.Valid, "cache read price null") || + !assert.False(t, p.CacheWritePriceMicros.Valid, "cache write price null") || + !assert.False(t, p.CostMicros.Valid, "cost null") { return false } return true @@ -1216,6 +1231,256 @@ func TestRecordTokenUsage(t *testing.T) { }, nil) }, }, + { + name: "valid token usage with cost", + request: &proto.RecordTokenUsageRequest{ + InterceptionId: uuid.NewString(), + MsgId: "msg_123", + InputTokens: 100, + OutputTokens: 200, + CacheReadInputTokens: 50, + CacheWriteInputTokens: 10, + CreatedAt: timestamppb.Now(), + }, + setupMocks: func(t *testing.T, db *dbmock.MockStore, req *proto.RecordTokenUsageRequest) { + interceptionID, err := uuid.Parse(req.GetInterceptionId()) + assert.NoError(t, err, "parse interception UUID") + + intc := newTestInterception(interceptionID) + groupID := uuid.New() + group := &database.GetHighestGroupAIBudgetByUserRow{GroupID: groupID, SpendLimitMicros: 1_000_000_000} + price := &database.AIModelPrice{ + Provider: intc.Provider, + Model: intc.Model, + InputPrice: sql.NullInt64{Int64: 3_000_000, Valid: true}, + OutputPrice: sql.NullInt64{Int64: 6_000_000, Valid: true}, + CacheReadPrice: sql.NullInt64{Int64: 300_000, Valid: true}, + CacheWritePrice: sql.NullInt64{Int64: 4_000_000, Valid: true}, + } + // No override + expectTokenUsageCostLookups(db, intc, nil, group, price) + + // input 300 + output 1200 + cache read 15 + cache write 40. + const wantCost int64 = 1555 + + db.EXPECT().InsertAIBridgeTokenUsage(gomock.Any(), gomock.Cond(func(p database.InsertAIBridgeTokenUsageParams) bool { + if !assert.Equal(t, uuid.NullUUID{UUID: groupID, Valid: true}, p.EffectiveGroupID, "effective group ID") || + !assert.Equal(t, price.InputPrice, p.InputPriceMicros, "input price") || + !assert.Equal(t, price.OutputPrice, p.OutputPriceMicros, "output price") || + !assert.Equal(t, price.CacheReadPrice, p.CacheReadPriceMicros, "cache read price") || + !assert.Equal(t, price.CacheWritePrice, p.CacheWritePriceMicros, "cache write price") || + !assert.Equal(t, sql.NullInt64{Int64: wantCost, Valid: true}, p.CostMicros, "cost") { + return false + } + return true + })).Return(database.AIBridgeTokenUsage{ID: uuid.New(), InterceptionID: interceptionID}, nil) + }, + }, + { + name: "valid token usage with user override", + request: &proto.RecordTokenUsageRequest{ + InterceptionId: uuid.NewString(), + MsgId: "msg_123", + InputTokens: 100, + CreatedAt: timestamppb.Now(), + }, + setupMocks: func(t *testing.T, db *dbmock.MockStore, req *proto.RecordTokenUsageRequest) { + interceptionID, err := uuid.Parse(req.GetInterceptionId()) + assert.NoError(t, err, "parse interception UUID") + + intc := newTestInterception(interceptionID) + overrideGroupID := uuid.New() + override := &database.UserAIBudgetOverride{ + UserID: intc.InitiatorID, + GroupID: overrideGroupID, + SpendLimitMicros: 1_500_000_000, + } + price := &database.AIModelPrice{ + Provider: intc.Provider, + Model: intc.Model, + InputPrice: sql.NullInt64{Int64: 3_000_000, Valid: true}, + } + // No group + expectTokenUsageCostLookups(db, intc, override, nil, price) + + db.EXPECT().InsertAIBridgeTokenUsage(gomock.Any(), gomock.Cond(func(p database.InsertAIBridgeTokenUsageParams) bool { + // Override group wins. + if !assert.Equal(t, uuid.NullUUID{UUID: overrideGroupID, Valid: true}, p.EffectiveGroupID, "effective group ID") || + !assert.Equal(t, sql.NullInt64{Int64: 300, Valid: true}, p.CostMicros, "cost") { + return false + } + return true + })).Return(database.AIBridgeTokenUsage{ID: uuid.New(), InterceptionID: interceptionID}, nil) + }, + }, + { + name: "valid token usage with budget but no price", + request: &proto.RecordTokenUsageRequest{ + InterceptionId: uuid.NewString(), + MsgId: "msg_123", + InputTokens: 100, + OutputTokens: 200, + CacheReadInputTokens: 50, + CacheWriteInputTokens: 10, + CreatedAt: timestamppb.Now(), + }, + setupMocks: func(t *testing.T, db *dbmock.MockStore, req *proto.RecordTokenUsageRequest) { + interceptionID, err := uuid.Parse(req.GetInterceptionId()) + assert.NoError(t, err, "parse interception UUID") + + intc := newTestInterception(interceptionID) + groupID := uuid.New() + group := &database.GetHighestGroupAIBudgetByUserRow{GroupID: groupID, SpendLimitMicros: 1_000_000_000} + // Budget resolves to a group, but the model has no price row. + // The resolved group must survive the price lookup's early + // return on sql.ErrNoRows, while prices and cost stay NULL. + expectTokenUsageCostLookups(db, intc, nil, group, nil) + + db.EXPECT().InsertAIBridgeTokenUsage(gomock.Any(), gomock.Cond(func(p database.InsertAIBridgeTokenUsageParams) bool { + if !assert.Equal(t, uuid.NullUUID{UUID: groupID, Valid: true}, p.EffectiveGroupID, "effective group ID") || + !assert.False(t, p.InputPriceMicros.Valid, "input price null") || + !assert.False(t, p.OutputPriceMicros.Valid, "output price null") || + !assert.False(t, p.CacheReadPriceMicros.Valid, "cache read price null") || + !assert.False(t, p.CacheWritePriceMicros.Valid, "cache write price null") || + !assert.False(t, p.CostMicros.Valid, "cost null") { + return false + } + return true + })).Return(database.AIBridgeTokenUsage{ID: uuid.New(), InterceptionID: interceptionID}, nil) + }, + }, + { + name: "valid token usage with price but no budget", + request: &proto.RecordTokenUsageRequest{ + InterceptionId: uuid.NewString(), + MsgId: "msg_123", + InputTokens: 100, + OutputTokens: 200, + CacheReadInputTokens: 50, + CacheWriteInputTokens: 10, + CreatedAt: timestamppb.Now(), + }, + setupMocks: func(t *testing.T, db *dbmock.MockStore, req *proto.RecordTokenUsageRequest) { + interceptionID, err := uuid.Parse(req.GetInterceptionId()) + assert.NoError(t, err, "parse interception UUID") + + intc := newTestInterception(interceptionID) + price := &database.AIModelPrice{ + Provider: intc.Provider, + Model: intc.Model, + InputPrice: sql.NullInt64{Int64: 3_000_000, Valid: true}, + OutputPrice: sql.NullInt64{Int64: 6_000_000, Valid: true}, + CacheReadPrice: sql.NullInt64{Int64: 300_000, Valid: true}, + CacheWritePrice: sql.NullInt64{Int64: 4_000_000, Valid: true}, + } + // No budget configured, but the model is priced: cost is + // computed independently of budget resolution, and the group + // attribution stays NULL. + expectTokenUsageCostLookups(db, intc, nil, nil, price) + + // input 300 + output 1200 + cache read 15 + cache write 40. + const wantCost int64 = 1555 + + db.EXPECT().InsertAIBridgeTokenUsage(gomock.Any(), gomock.Cond(func(p database.InsertAIBridgeTokenUsageParams) bool { + if !assert.False(t, p.EffectiveGroupID.Valid, "effective group ID null") || + !assert.Equal(t, price.InputPrice, p.InputPriceMicros, "input price") || + !assert.Equal(t, price.OutputPrice, p.OutputPriceMicros, "output price") || + !assert.Equal(t, price.CacheReadPrice, p.CacheReadPriceMicros, "cache read price") || + !assert.Equal(t, price.CacheWritePrice, p.CacheWritePriceMicros, "cache write price") || + !assert.Equal(t, sql.NullInt64{Int64: wantCost, Valid: true}, p.CostMicros, "cost") { + return false + } + return true + })).Return(database.AIBridgeTokenUsage{ID: uuid.New(), InterceptionID: interceptionID}, nil) + }, + }, + { + name: "valid token usage with zero prices", + request: &proto.RecordTokenUsageRequest{ + InterceptionId: uuid.NewString(), + MsgId: "msg_123", + InputTokens: 100, + OutputTokens: 200, + CacheReadInputTokens: 50, + CacheWriteInputTokens: 10, + CreatedAt: timestamppb.Now(), + }, + setupMocks: func(t *testing.T, db *dbmock.MockStore, req *proto.RecordTokenUsageRequest) { + interceptionID, err := uuid.Parse(req.GetInterceptionId()) + assert.NoError(t, err, "parse interception UUID") + + intc := newTestInterception(interceptionID) + // A model priced at zero is distinct from an unpriced model: + // the price columns and cost are recorded as 0, not NULL. + price := &database.AIModelPrice{ + Provider: intc.Provider, + Model: intc.Model, + InputPrice: sql.NullInt64{Int64: 0, Valid: true}, + OutputPrice: sql.NullInt64{Int64: 0, Valid: true}, + CacheReadPrice: sql.NullInt64{Int64: 0, Valid: true}, + CacheWritePrice: sql.NullInt64{Int64: 0, Valid: true}, + } + expectTokenUsageCostLookups(db, intc, nil, nil, price) + + db.EXPECT().InsertAIBridgeTokenUsage(gomock.Any(), gomock.Cond(func(p database.InsertAIBridgeTokenUsageParams) bool { + zero := sql.NullInt64{Int64: 0, Valid: true} + if !assert.Equal(t, zero, p.InputPriceMicros, "input price zero") || + !assert.Equal(t, zero, p.OutputPriceMicros, "output price zero") || + !assert.Equal(t, zero, p.CacheReadPriceMicros, "cache read price zero") || + !assert.Equal(t, zero, p.CacheWritePriceMicros, "cache write price zero") || + // Cost is 0 but recorded (Valid), not NULL. + !assert.Equal(t, zero, p.CostMicros, "cost zero") { + return false + } + return true + })).Return(database.AIBridgeTokenUsage{ID: uuid.New(), InterceptionID: interceptionID}, nil) + }, + }, + { + name: "valid token usage with all null prices", + request: &proto.RecordTokenUsageRequest{ + InterceptionId: uuid.NewString(), + MsgId: "msg_123", + InputTokens: 100, + OutputTokens: 200, + CacheReadInputTokens: 50, + CacheWriteInputTokens: 10, + CreatedAt: timestamppb.Now(), + }, + setupMocks: func(t *testing.T, db *dbmock.MockStore, req *proto.RecordTokenUsageRequest) { + interceptionID, err := uuid.Parse(req.GetInterceptionId()) + assert.NoError(t, err, "parse interception UUID") + + intc := newTestInterception(interceptionID) + // The price row exists but every price column is NULL. Each + // category is treated as zero for cost, so the columns are + // recorded as NULL while cost is recorded as 0 (not NULL): + // cost's NULL-ness tracks price row presence, not the price + // values. + price := &database.AIModelPrice{ + Provider: intc.Provider, + Model: intc.Model, + InputPrice: sql.NullInt64{Valid: false}, + OutputPrice: sql.NullInt64{Valid: false}, + CacheReadPrice: sql.NullInt64{Valid: false}, + CacheWritePrice: sql.NullInt64{Valid: false}, + } + expectTokenUsageCostLookups(db, intc, nil, nil, price) + + db.EXPECT().InsertAIBridgeTokenUsage(gomock.Any(), gomock.Cond(func(p database.InsertAIBridgeTokenUsageParams) bool { + if !assert.False(t, p.InputPriceMicros.Valid, "input price null") || + !assert.False(t, p.OutputPriceMicros.Valid, "output price null") || + !assert.False(t, p.CacheReadPriceMicros.Valid, "cache read price null") || + !assert.False(t, p.CacheWritePriceMicros.Valid, "cache write price null") || + // Cost is recorded as 0 (Valid), not NULL, because the + // price row exists. + !assert.Equal(t, sql.NullInt64{Int64: 0, Valid: true}, p.CostMicros, "cost zero") { + return false + } + return true + })).Return(database.AIBridgeTokenUsage{ID: uuid.New(), InterceptionID: interceptionID}, nil) + }, + }, { name: "invalid interception ID", request: &proto.RecordTokenUsageRequest{ @@ -1228,7 +1493,7 @@ func TestRecordTokenUsage(t *testing.T) { expectedErr: "failed to parse interception_id", }, { - name: "database error", + name: "interception lookup error", request: &proto.RecordTokenUsageRequest{ InterceptionId: uuid.NewString(), MsgId: "msg_123", @@ -1237,6 +1502,56 @@ func TestRecordTokenUsage(t *testing.T) { CreatedAt: timestamppb.Now(), }, setupMocks: func(t *testing.T, db *dbmock.MockStore, req *proto.RecordTokenUsageRequest) { + interceptionID, err := uuid.Parse(req.GetInterceptionId()) + assert.NoError(t, err, "parse interception UUID") + + // An unexpected interception lookup error fails the record; + // no token usage is inserted. + db.EXPECT().GetAIBridgeInterceptionByID(gomock.Any(), interceptionID). + Return(database.AIBridgeInterception{}, sql.ErrConnDone) + }, + expectedErr: "get interception", + }, + { + name: "price lookup error", + request: &proto.RecordTokenUsageRequest{ + InterceptionId: uuid.NewString(), + MsgId: "msg_123", + InputTokens: 100, + OutputTokens: 200, + CreatedAt: timestamppb.Now(), + }, + setupMocks: func(t *testing.T, db *dbmock.MockStore, req *proto.RecordTokenUsageRequest) { + interceptionID, err := uuid.Parse(req.GetInterceptionId()) + assert.NoError(t, err, "parse interception UUID") + + // An unexpected price lookup error (not sql.ErrNoRows) fails + // the record. + intc := newTestInterception(interceptionID) + db.EXPECT().GetAIBridgeInterceptionByID(gomock.Any(), interceptionID).Return(intc, nil) + db.EXPECT().GetUserAIBudgetOverride(gomock.Any(), intc.InitiatorID). + Return(database.UserAIBudgetOverride{}, sql.ErrNoRows) + db.EXPECT().GetHighestGroupAIBudgetByUser(gomock.Any(), intc.InitiatorID). + Return(database.GetHighestGroupAIBudgetByUserRow{}, sql.ErrNoRows) + db.EXPECT().GetAIModelPriceByProviderModel(gomock.Any(), gomock.Any()). + Return(database.AIModelPrice{}, sql.ErrConnDone) + }, + expectedErr: "resolve token usage cost", + }, + { + name: "insert error", + request: &proto.RecordTokenUsageRequest{ + InterceptionId: uuid.NewString(), + MsgId: "msg_123", + InputTokens: 100, + OutputTokens: 200, + CreatedAt: timestamppb.Now(), + }, + setupMocks: func(t *testing.T, db *dbmock.MockStore, req *proto.RecordTokenUsageRequest) { + interceptionID, err := uuid.Parse(req.GetInterceptionId()) + assert.NoError(t, err, "parse interception UUID") + + expectTokenUsageCostLookups(db, newTestInterception(interceptionID), nil, nil, nil) db.EXPECT().InsertAIBridgeTokenUsage(gomock.Any(), gomock.Any()).Return(database.AIBridgeTokenUsage{}, sql.ErrConnDone) }, expectedErr: "insert token usage", @@ -1245,6 +1560,128 @@ func TestRecordTokenUsage(t *testing.T) { ) } +// TestRecordTokenUsageAuthorized exercises RecordTokenUsage end-to-end against a +// real database through the dbauthz layer as subjectAibridged. This catches missing +// RBAC grants on the aibridged subject and verifies the cost columns round-trip to storage. +func TestRecordTokenUsageAuthorized(t *testing.T) { + t.Parallel() + + ctx := testutil.Context(t, testutil.WaitLong) + logger := testutil.Logger(t) + + rawDB, _ := dbtestutil.NewDB(t) + authzDB := dbauthz.New(rawDB, rbac.NewStrictAuthorizer(prometheus.NewRegistry()), logger, coderdtest.AccessControlStorePointer()) + + // Seed prerequisites via the raw (unauthorized) store. The user belongs to a + // group with a budget, so the effective group resolves to that group. + org := dbgen.Organization(t, rawDB, database.Organization{}) + user := dbgen.User(t, rawDB, database.User{}) + dbgen.OrganizationMember(t, rawDB, database.OrganizationMember{OrganizationID: org.ID, UserID: user.ID}) + group := dbgen.Group(t, rawDB, database.Group{OrganizationID: org.ID}) + dbgen.GroupMember(t, rawDB, database.GroupMemberTable{UserID: user.ID, GroupID: group.ID}) + + _, err := rawDB.UpsertGroupAIBudget(ctx, database.UpsertGroupAIBudgetParams{ + GroupID: group.ID, + SpendLimitMicros: 1_000_000_000, + }) + require.NoError(t, err, "upsert group AI budget") + + const provider, model = "anthropic", "claude-sonnet-4-6" + priceSeed, err := json.Marshal([]map[string]any{{ + "provider": provider, + "model": model, + "input_price": 3_000_000, + "output_price": 6_000_000, + "cache_read_price": 300_000, + "cache_write_price": 4_000_000, + }}) + require.NoError(t, err) + require.NoError(t, rawDB.UpsertAIModelPrices(ctx, priceSeed), "seed model prices") + + intc := dbgen.AIBridgeInterception(t, rawDB, database.InsertAIBridgeInterceptionParams{ + InitiatorID: user.ID, + Provider: provider, + Model: model, + }, nil) + + // The server runs every store call as subjectAibridged via the authzDB. + srv, err := aibridgedserver.NewServer(ctx, authzDB, logger, "/", codersdk.AIBridgeConfig{}, nil, requiredExperiments, agplaiseats.Noop{}) + require.NoError(t, err) + + _, err = srv.RecordTokenUsage(ctx, &proto.RecordTokenUsageRequest{ + InterceptionId: intc.ID.String(), + MsgId: "msg_e2e", + InputTokens: 100, + OutputTokens: 200, + CacheReadInputTokens: 50, + CacheWriteInputTokens: 10, + CreatedAt: timestamppb.Now(), + }) + require.NoError(t, err, "record token usage") + + // Read the persisted row back via the raw store and verify the snapshot. + usages, err := rawDB.GetAIBridgeTokenUsagesByInterceptionID(ctx, intc.ID) + require.NoError(t, err) + require.Len(t, usages, 1) + got := usages[0] + + require.Equal(t, uuid.NullUUID{UUID: group.ID, Valid: true}, got.EffectiveGroupID, "effective group") + require.Equal(t, sql.NullInt64{Int64: 3_000_000, Valid: true}, got.InputPriceMicros, "input price") + require.Equal(t, sql.NullInt64{Int64: 6_000_000, Valid: true}, got.OutputPriceMicros, "output price") + require.Equal(t, sql.NullInt64{Int64: 300_000, Valid: true}, got.CacheReadPriceMicros, "cache read price") + require.Equal(t, sql.NullInt64{Int64: 4_000_000, Valid: true}, got.CacheWritePriceMicros, "cache write price") + // input 300 + output 1200 + cache read 15 + cache write 40. + require.Equal(t, sql.NullInt64{Int64: 1555, Valid: true}, got.CostMicros, "cost") +} + +// newTestInterception returns an interception with a fixed initiator, provider, +// and model for cost-attribution test setup. +func newTestInterception(id uuid.UUID) database.AIBridgeInterception { + return database.AIBridgeInterception{ + ID: id, + InitiatorID: uuid.New(), + Provider: "anthropic", + Model: "claude-sonnet-4-6", + } +} + +// expectTokenUsageCostLookups mocks the store lookups made by resolveTokenUsageCost +// (budget resolution and the price lookup). A nil override, group, or price makes that +// lookup return sql.ErrNoRows. Budget resolution mirrors production code: a non-nil override +// wins and skips the group lookup, so group is consulted only when override is nil. +func expectTokenUsageCostLookups( + db *dbmock.MockStore, + intc database.AIBridgeInterception, + override *database.UserAIBudgetOverride, + group *database.GetHighestGroupAIBudgetByUserRow, + price *database.AIModelPrice, +) { + db.EXPECT().GetAIBridgeInterceptionByID(gomock.Any(), intc.ID).Return(intc, nil) + + if override != nil { + db.EXPECT().GetUserAIBudgetOverride(gomock.Any(), intc.InitiatorID).Return(*override, nil) + } else { + db.EXPECT().GetUserAIBudgetOverride(gomock.Any(), intc.InitiatorID). + Return(database.UserAIBudgetOverride{}, sql.ErrNoRows) + if group != nil { + db.EXPECT().GetHighestGroupAIBudgetByUser(gomock.Any(), intc.InitiatorID).Return(*group, nil) + } else { + db.EXPECT().GetHighestGroupAIBudgetByUser(gomock.Any(), intc.InitiatorID). + Return(database.GetHighestGroupAIBudgetByUserRow{}, sql.ErrNoRows) + } + } + + if price != nil { + db.EXPECT().GetAIModelPriceByProviderModel(gomock.Any(), database.GetAIModelPriceByProviderModelParams{ + Provider: intc.Provider, + Model: intc.Model, + }).Return(*price, nil) + } else { + db.EXPECT().GetAIModelPriceByProviderModel(gomock.Any(), gomock.Any()). + Return(database.AIModelPrice{}, sql.ErrNoRows) + } +} + func TestRecordPromptUsage(t *testing.T) { t.Parallel() @@ -1722,6 +2159,7 @@ func TestStructuredLogging(t *testing.T) { name: "RecordTokenUsage_logs_when_enabled", structuredLogging: true, setupMocks: func(db *dbmock.MockStore, intcID uuid.UUID) { + expectTokenUsageCostLookups(db, newTestInterception(intcID), nil, nil, nil) db.EXPECT().InsertAIBridgeTokenUsage(gomock.Any(), gomock.Any()).Return(database.AIBridgeTokenUsage{ ID: uuid.New(), InterceptionID: intcID, diff --git a/coderd/aibridgedserver/cost.go b/coderd/aibridgedserver/cost.go new file mode 100644 index 00000000000..0a6b9ba5b63 --- /dev/null +++ b/coderd/aibridgedserver/cost.go @@ -0,0 +1,97 @@ +package aibridgedserver + +import ( + "context" + "database/sql" + "errors" + + "github.com/google/uuid" + "golang.org/x/xerrors" + + "cdr.dev/slog/v3" + "github.com/coder/coder/v2/coderd/aibridge/budget" + "github.com/coder/coder/v2/coderd/aibridged/proto" + "github.com/coder/coder/v2/coderd/database" +) + +// tokensPerMillion is the divisor for prices, which are quoted per million +// tokens. +const tokensPerMillion = 1_000_000 + +// tokenUsageCost holds the cost-attribution columns snapshotted onto a token +// usage record. A field left unset (Valid == false) is recorded as SQL NULL; a +// price or cost of 0 is recorded as 0, which is distinct from NULL. +type tokenUsageCost struct { + effectiveGroupID uuid.NullUUID + inputPriceMicros sql.NullInt64 + outputPriceMicros sql.NullInt64 + cacheReadPriceMicros sql.NullInt64 + cacheWritePriceMicros sql.NullInt64 + costMicros sql.NullInt64 +} + +// resolveTokenUsageCost resolves the effective group and per-token prices for an +// interception and computes its cost. Two outcomes are expected and yield NULL +// columns rather than an error: a user with no configured budget (yields a NULL +// group) and a model absent from the price table (yields NULL prices and cost). +// Any other error is returned. A NULL cost unambiguously means "model not priced". +func (s *Server) resolveTokenUsageCost(ctx context.Context, intc database.AIBridgeInterception, in *proto.RecordTokenUsageRequest) (tokenUsageCost, error) { + var result tokenUsageCost + + // Resolve the effective group for attribution. This is independent of + // whether the model is priced. ok is false when no budget is configured, + // which leaves the group attribution NULL. + effectiveBudget, ok, err := budget.ResolveUserAIBudget(ctx, s.store, intc.InitiatorID, s.budgetPolicy) + if err != nil { + return tokenUsageCost{}, xerrors.Errorf("resolve effective AI budget for user %q with policy %q: %w", intc.InitiatorID, s.budgetPolicy, err) + } + if ok { + result.effectiveGroupID = uuid.NullUUID{UUID: effectiveBudget.GroupID, Valid: true} + } + + // Snapshot the price for this (provider, model) and compute cost. + price, err := s.store.GetAIModelPriceByProviderModel(ctx, database.GetAIModelPriceByProviderModelParams{ + Provider: intc.Provider, + Model: intc.Model, + }) + switch { + case errors.Is(err, sql.ErrNoRows): + // Model not in the price table: record tokens but leave cost NULL. + s.logger.Debug(ctx, "no price found for model, recording token usage with NULL cost", + slog.F("provider", intc.Provider), slog.F("model", intc.Model)) + return result, nil + case err != nil: + return tokenUsageCost{}, xerrors.Errorf("look up model price for %s/%s: %w", intc.Provider, intc.Model, err) + } + + result.inputPriceMicros = price.InputPrice + result.outputPriceMicros = price.OutputPrice + result.cacheReadPriceMicros = price.CacheReadPrice + result.cacheWritePriceMicros = price.CacheWritePrice + result.costMicros = sql.NullInt64{ + Int64: computeCost(price, + in.GetInputTokens(), in.GetOutputTokens(), + in.GetCacheReadInputTokens(), in.GetCacheWriteInputTokens()), + Valid: true, + } + return result, nil +} + +// computeCost returns the cost of an interception in micro-units, snapshotting +// the per-token prices from the price table. Prices are expressed per million +// tokens; a NULL price column is treated as zero (e.g. providers that do not +// charge for cache writes). +func computeCost(price database.AIModelPrice, inputTokens, outputTokens, cacheReadTokens, cacheWriteTokens int64) int64 { + return tokenCost(inputTokens, price.InputPrice) + + tokenCost(outputTokens, price.OutputPrice) + + tokenCost(cacheReadTokens, price.CacheReadPrice) + + tokenCost(cacheWriteTokens, price.CacheWritePrice) +} + +// tokenCost returns tokens * price / 1,000,000, treating a NULL price as zero. +func tokenCost(tokens int64, pricePerMillion sql.NullInt64) int64 { + if !pricePerMillion.Valid { + return 0 + } + return tokens * pricePerMillion.Int64 / tokensPerMillion +} diff --git a/coderd/aibridgedserver/cost_internal_test.go b/coderd/aibridgedserver/cost_internal_test.go new file mode 100644 index 00000000000..a36d24e16e1 --- /dev/null +++ b/coderd/aibridgedserver/cost_internal_test.go @@ -0,0 +1,124 @@ +package aibridgedserver + +import ( + "database/sql" + "testing" + + "github.com/coder/coder/v2/coderd/database" +) + +func TestComputeCost(t *testing.T) { + t.Parallel() + + nullInt64 := func(v int64) sql.NullInt64 { return sql.NullInt64{Int64: v, Valid: true} } + + tests := []struct { + name string + price database.AIModelPrice + inputTokens, outputTokens, cacheReadTokens, cacheWriteTokens int64 + want int64 + }{ + { + name: "all priced", + price: database.AIModelPrice{ + InputPrice: nullInt64(3_000_000), + OutputPrice: nullInt64(6_000_000), + CacheReadPrice: nullInt64(300_000), + CacheWritePrice: nullInt64(3_750_000), + }, + inputTokens: 100, + outputTokens: 200, + cacheReadTokens: 50, + cacheWriteTokens: 10, + // 300 + 1200 + 15 + 37 (10*3_750_000/1e6 = 37, integer division). + want: 1552, + }, + { + name: "null cache write price treated as zero", + price: database.AIModelPrice{ + InputPrice: nullInt64(3_000_000), + OutputPrice: nullInt64(6_000_000), + CacheReadPrice: nullInt64(300_000), + CacheWritePrice: sql.NullInt64{Valid: false}, + }, + inputTokens: 100, + outputTokens: 200, + cacheReadTokens: 50, + cacheWriteTokens: 10, + // 300 + 1200 + 15 + 0. + want: 1515, + }, + { + name: "all prices null is zero cost", + price: database.AIModelPrice{}, + inputTokens: 100, + outputTokens: 200, + cacheReadTokens: 50, + cacheWriteTokens: 10, + want: 0, + }, + { + name: "zero tokens is zero cost", + price: database.AIModelPrice{ + InputPrice: nullInt64(3_000_000), + OutputPrice: nullInt64(6_000_000), + }, + want: 0, + }, + { + name: "integer division truncates", + price: database.AIModelPrice{ + // 1 token at 1 micro-unit per million tokens rounds down to 0. + InputPrice: nullInt64(1), + }, + inputTokens: 1, + want: 0, + }, + { + name: "price just below one micro-unit per token floors to zero", + price: database.AIModelPrice{ + InputPrice: nullInt64(999_999), + }, + inputTokens: 1, // 1 * 999_999 = 999_999, below 1_000_000 + want: 0, + }, + { + name: "sub-unit price summed across tokens still floors to zero", + price: database.AIModelPrice{ + InputPrice: nullInt64(999), + }, + inputTokens: 1000, // 1000 * 999 = 999_000, below 1_000_000 + want: 0, + }, + { + name: "sub-unit price crosses one micro-unit once the product reaches 1e6", + price: database.AIModelPrice{ + InputPrice: nullInt64(999), + }, + inputTokens: 1002, // 1002 * 999 = 1_000_998 + want: 1, + }, + { + // Stress the per-term numerator near the int64 ceiling. At a $75/M + // model the overflow point is ~123e9 tokens (123e9 * 75e6 = 9.225e18, + // just over int64 max 9.223e18); 122e9 stays just under. + name: "large token count at a high price does not overflow", + price: database.AIModelPrice{ + InputPrice: nullInt64(75_000_000), // $75 per 1M tokens + }, + inputTokens: 122_000_000_000, // 122e9 * 75e6 = 9.15e18 < int64 max + want: 9_150_000_000_000, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + t.Parallel() + + got := computeCost(tt.price, tt.inputTokens, tt.outputTokens, tt.cacheReadTokens, tt.cacheWriteTokens) + if got != tt.want { + t.Fatalf("computeCost = %d, want %d", got, tt.want) + } + }) + } +} diff --git a/coderd/database/check_constraint.go b/coderd/database/check_constraint.go index 67122363f58..935c38c0d17 100644 --- a/coderd/database/check_constraint.go +++ b/coderd/database/check_constraint.go @@ -6,51 +6,56 @@ type CheckConstraint string // CheckConstraint enums. const ( - CheckAIGatewayKeysHashedSecretCheck CheckConstraint = "ai_gateway_keys_hashed_secret_check" // ai_gateway_keys - CheckAIGatewayKeysNameCheck CheckConstraint = "ai_gateway_keys_name_check" // ai_gateway_keys - CheckAIGatewayKeysSecretPrefixCheck CheckConstraint = "ai_gateway_keys_secret_prefix_check" // ai_gateway_keys - CheckAIModelPricesCacheReadPriceCheck CheckConstraint = "ai_model_prices_cache_read_price_check" // ai_model_prices - CheckAIModelPricesCacheWritePriceCheck CheckConstraint = "ai_model_prices_cache_write_price_check" // ai_model_prices - CheckAIModelPricesInputPriceCheck CheckConstraint = "ai_model_prices_input_price_check" // ai_model_prices - CheckAIModelPricesOutputPriceCheck CheckConstraint = "ai_model_prices_output_price_check" // ai_model_prices - CheckAIProvidersNameCheck CheckConstraint = "ai_providers_name_check" // ai_providers - CheckAPIKeysAllowListNotEmpty CheckConstraint = "api_keys_allow_list_not_empty" // api_keys - CheckBoundaryLogsSequenceNumberCheck CheckConstraint = "boundary_logs_sequence_number_check" // boundary_logs - CheckChatModelConfigsAIProviderRequiredWhenActive CheckConstraint = "chat_model_configs_ai_provider_required_when_active" // chat_model_configs - CheckChatModelConfigsCompressionThresholdCheck CheckConstraint = "chat_model_configs_compression_threshold_check" // chat_model_configs - CheckChatModelConfigsContextLimitCheck CheckConstraint = "chat_model_configs_context_limit_check" // chat_model_configs - CheckChatUsageLimitConfigDefaultLimitMicrosCheck CheckConstraint = "chat_usage_limit_config_default_limit_micros_check" // chat_usage_limit_config - CheckChatUsageLimitConfigPeriodCheck CheckConstraint = "chat_usage_limit_config_period_check" // chat_usage_limit_config - CheckChatUsageLimitConfigSingletonCheck CheckConstraint = "chat_usage_limit_config_singleton_check" // chat_usage_limit_config - CheckChatAclOnlyOnRootChats CheckConstraint = "chat_acl_only_on_root_chats" // chats - CheckChatGroupAclNotNullJsonb CheckConstraint = "chat_group_acl_not_null_jsonb" // chats - CheckChatUserAclNotNullJsonb CheckConstraint = "chat_user_acl_not_null_jsonb" // chats - CheckChatsPinOrderArchivedCheck CheckConstraint = "chats_pin_order_archived_check" // chats - CheckChatsPinOrderParentCheck CheckConstraint = "chats_pin_order_parent_check" // chats - CheckOneTimePasscodeSet CheckConstraint = "one_time_passcode_set" // users - CheckUsersChatSpendLimitMicrosCheck CheckConstraint = "users_chat_spend_limit_micros_check" // users - CheckUsersEmailNotEmpty CheckConstraint = "users_email_not_empty" // users - CheckUsersServiceAccountLoginType CheckConstraint = "users_service_account_login_type" // users - CheckUsersUsernameMinLength CheckConstraint = "users_username_min_length" // users - CheckOrganizationIDNotZero CheckConstraint = "organization_id_not_zero" // custom_roles - CheckGroupAIBudgetsSpendLimitMicrosCheck CheckConstraint = "group_ai_budgets_spend_limit_micros_check" // group_ai_budgets - CheckGroupsChatSpendLimitMicrosCheck CheckConstraint = "groups_chat_spend_limit_micros_check" // groups - CheckMcpServerConfigsAuthTypeCheck CheckConstraint = "mcp_server_configs_auth_type_check" // mcp_server_configs - CheckMcpServerConfigsAvailabilityCheck CheckConstraint = "mcp_server_configs_availability_check" // mcp_server_configs - CheckMcpServerConfigsTransportCheck CheckConstraint = "mcp_server_configs_transport_check" // mcp_server_configs - CheckMaxProvisionerLogsLength CheckConstraint = "max_provisioner_logs_length" // provisioner_jobs - CheckMaxLogsLength CheckConstraint = "max_logs_length" // workspace_agents - CheckSubsystemsNotNone CheckConstraint = "subsystems_not_none" // workspace_agents - CheckWorkspaceBuildsDeadlineBelowMaxDeadline CheckConstraint = "workspace_builds_deadline_below_max_deadline" // workspace_builds - CheckGroupAclIsObject CheckConstraint = "group_acl_is_object" // workspaces - CheckUserAclIsObject CheckConstraint = "user_acl_is_object" // workspaces - CheckTelemetryLockEventTypeConstraint CheckConstraint = "telemetry_lock_event_type_constraint" // telemetry_locks - CheckValidationMonotonicOrder CheckConstraint = "validation_monotonic_order" // template_version_parameters - CheckUsageEventTypeCheck CheckConstraint = "usage_event_type_check" // usage_events - CheckUserAIBudgetOverridesSpendLimitMicrosCheck CheckConstraint = "user_ai_budget_overrides_spend_limit_micros_check" // user_ai_budget_overrides - CheckUserAIProviderKeysAPIKeyCheck CheckConstraint = "user_ai_provider_keys_api_key_check" // user_ai_provider_keys - CheckUserSkillsContentSize CheckConstraint = "user_skills_content_size" // user_skills - CheckUserSkillsDescriptionSize CheckConstraint = "user_skills_description_size" // user_skills - CheckUserSkillsNameFormat CheckConstraint = "user_skills_name_format" // user_skills - CheckUserSkillsNameSize CheckConstraint = "user_skills_name_size" // user_skills + CheckAIGatewayKeysHashedSecretCheck CheckConstraint = "ai_gateway_keys_hashed_secret_check" // ai_gateway_keys + CheckAIGatewayKeysNameCheck CheckConstraint = "ai_gateway_keys_name_check" // ai_gateway_keys + CheckAIGatewayKeysSecretPrefixCheck CheckConstraint = "ai_gateway_keys_secret_prefix_check" // ai_gateway_keys + CheckAIModelPricesCacheReadPriceCheck CheckConstraint = "ai_model_prices_cache_read_price_check" // ai_model_prices + CheckAIModelPricesCacheWritePriceCheck CheckConstraint = "ai_model_prices_cache_write_price_check" // ai_model_prices + CheckAIModelPricesInputPriceCheck CheckConstraint = "ai_model_prices_input_price_check" // ai_model_prices + CheckAIModelPricesOutputPriceCheck CheckConstraint = "ai_model_prices_output_price_check" // ai_model_prices + CheckAIProvidersNameCheck CheckConstraint = "ai_providers_name_check" // ai_providers + CheckAibridgeTokenUsagesCacheReadPriceMicrosCheck CheckConstraint = "aibridge_token_usages_cache_read_price_micros_check" // aibridge_token_usages + CheckAibridgeTokenUsagesCacheWritePriceMicrosCheck CheckConstraint = "aibridge_token_usages_cache_write_price_micros_check" // aibridge_token_usages + CheckAibridgeTokenUsagesCostMicrosCheck CheckConstraint = "aibridge_token_usages_cost_micros_check" // aibridge_token_usages + CheckAibridgeTokenUsagesInputPriceMicrosCheck CheckConstraint = "aibridge_token_usages_input_price_micros_check" // aibridge_token_usages + CheckAibridgeTokenUsagesOutputPriceMicrosCheck CheckConstraint = "aibridge_token_usages_output_price_micros_check" // aibridge_token_usages + CheckAPIKeysAllowListNotEmpty CheckConstraint = "api_keys_allow_list_not_empty" // api_keys + CheckBoundaryLogsSequenceNumberCheck CheckConstraint = "boundary_logs_sequence_number_check" // boundary_logs + CheckChatModelConfigsAIProviderRequiredWhenActive CheckConstraint = "chat_model_configs_ai_provider_required_when_active" // chat_model_configs + CheckChatModelConfigsCompressionThresholdCheck CheckConstraint = "chat_model_configs_compression_threshold_check" // chat_model_configs + CheckChatModelConfigsContextLimitCheck CheckConstraint = "chat_model_configs_context_limit_check" // chat_model_configs + CheckChatUsageLimitConfigDefaultLimitMicrosCheck CheckConstraint = "chat_usage_limit_config_default_limit_micros_check" // chat_usage_limit_config + CheckChatUsageLimitConfigPeriodCheck CheckConstraint = "chat_usage_limit_config_period_check" // chat_usage_limit_config + CheckChatUsageLimitConfigSingletonCheck CheckConstraint = "chat_usage_limit_config_singleton_check" // chat_usage_limit_config + CheckChatAclOnlyOnRootChats CheckConstraint = "chat_acl_only_on_root_chats" // chats + CheckChatGroupAclNotNullJsonb CheckConstraint = "chat_group_acl_not_null_jsonb" // chats + CheckChatUserAclNotNullJsonb CheckConstraint = "chat_user_acl_not_null_jsonb" // chats + CheckChatsPinOrderArchivedCheck CheckConstraint = "chats_pin_order_archived_check" // chats + CheckChatsPinOrderParentCheck CheckConstraint = "chats_pin_order_parent_check" // chats + CheckOneTimePasscodeSet CheckConstraint = "one_time_passcode_set" // users + CheckUsersChatSpendLimitMicrosCheck CheckConstraint = "users_chat_spend_limit_micros_check" // users + CheckUsersEmailNotEmpty CheckConstraint = "users_email_not_empty" // users + CheckUsersServiceAccountLoginType CheckConstraint = "users_service_account_login_type" // users + CheckUsersUsernameMinLength CheckConstraint = "users_username_min_length" // users + CheckOrganizationIDNotZero CheckConstraint = "organization_id_not_zero" // custom_roles + CheckGroupAIBudgetsSpendLimitMicrosCheck CheckConstraint = "group_ai_budgets_spend_limit_micros_check" // group_ai_budgets + CheckGroupsChatSpendLimitMicrosCheck CheckConstraint = "groups_chat_spend_limit_micros_check" // groups + CheckMcpServerConfigsAuthTypeCheck CheckConstraint = "mcp_server_configs_auth_type_check" // mcp_server_configs + CheckMcpServerConfigsAvailabilityCheck CheckConstraint = "mcp_server_configs_availability_check" // mcp_server_configs + CheckMcpServerConfigsTransportCheck CheckConstraint = "mcp_server_configs_transport_check" // mcp_server_configs + CheckMaxProvisionerLogsLength CheckConstraint = "max_provisioner_logs_length" // provisioner_jobs + CheckMaxLogsLength CheckConstraint = "max_logs_length" // workspace_agents + CheckSubsystemsNotNone CheckConstraint = "subsystems_not_none" // workspace_agents + CheckWorkspaceBuildsDeadlineBelowMaxDeadline CheckConstraint = "workspace_builds_deadline_below_max_deadline" // workspace_builds + CheckGroupAclIsObject CheckConstraint = "group_acl_is_object" // workspaces + CheckUserAclIsObject CheckConstraint = "user_acl_is_object" // workspaces + CheckTelemetryLockEventTypeConstraint CheckConstraint = "telemetry_lock_event_type_constraint" // telemetry_locks + CheckValidationMonotonicOrder CheckConstraint = "validation_monotonic_order" // template_version_parameters + CheckUsageEventTypeCheck CheckConstraint = "usage_event_type_check" // usage_events + CheckUserAIBudgetOverridesSpendLimitMicrosCheck CheckConstraint = "user_ai_budget_overrides_spend_limit_micros_check" // user_ai_budget_overrides + CheckUserAIProviderKeysAPIKeyCheck CheckConstraint = "user_ai_provider_keys_api_key_check" // user_ai_provider_keys + CheckUserSkillsContentSize CheckConstraint = "user_skills_content_size" // user_skills + CheckUserSkillsDescriptionSize CheckConstraint = "user_skills_description_size" // user_skills + CheckUserSkillsNameFormat CheckConstraint = "user_skills_name_format" // user_skills + CheckUserSkillsNameSize CheckConstraint = "user_skills_name_size" // user_skills ) diff --git a/coderd/database/dbauthz/dbauthz.go b/coderd/database/dbauthz/dbauthz.go index b0765b5c2c9..a91c45b6fc3 100644 --- a/coderd/database/dbauthz/dbauthz.go +++ b/coderd/database/dbauthz/dbauthz.go @@ -650,9 +650,9 @@ var ( }, rbac.ResourceApiKey.Type: {policy.ActionRead}, // Validate API keys. rbac.ResourceAibridgeInterception.Type: {policy.ActionCreate, policy.ActionRead, policy.ActionUpdate, policy.ActionDelete}, - rbac.ResourceAiModelPrice.Type: {policy.ActionUpdate}, // Required for the startup price seeder. - rbac.ResourceAiSeat.Type: {policy.ActionCreate}, // Required for UpsertAISeatState. - rbac.ResourceAIProvider.Type: {policy.ActionRead}, // Required to load the provider snapshot (and per-provider keys) at startup. + rbac.ResourceAiModelPrice.Type: {policy.ActionRead, policy.ActionUpdate}, // Read: per-interception cost lookup. Update: startup price seeder. + rbac.ResourceAiSeat.Type: {policy.ActionCreate}, // Required for UpsertAISeatState. + rbac.ResourceAIProvider.Type: {policy.ActionRead}, // Required to load the provider snapshot (and per-provider keys) at startup. }), User: []rbac.Permission{}, ByOrgID: map[string]rbac.OrgPermissions{}, diff --git a/coderd/database/dbgen/dbgen.go b/coderd/database/dbgen/dbgen.go index f23123d8565..b56ce64ee0d 100644 --- a/coderd/database/dbgen/dbgen.go +++ b/coderd/database/dbgen/dbgen.go @@ -2037,6 +2037,12 @@ func AIBridgeTokenUsage(t testing.TB, db database.Store, seed database.InsertAIB CacheWriteInputTokens: seed.CacheWriteInputTokens, Metadata: takeFirstSlice(seed.Metadata, json.RawMessage("{}")), CreatedAt: takeFirst(seed.CreatedAt, dbtime.Now()), + EffectiveGroupID: seed.EffectiveGroupID, + InputPriceMicros: seed.InputPriceMicros, + OutputPriceMicros: seed.OutputPriceMicros, + CacheReadPriceMicros: seed.CacheReadPriceMicros, + CacheWritePriceMicros: seed.CacheWritePriceMicros, + CostMicros: seed.CostMicros, }) require.NoError(t, err, "insert aibridge token usage") return usage diff --git a/coderd/database/dump.sql b/coderd/database/dump.sql index f08d96e4fab..8018482a1ce 100644 --- a/coderd/database/dump.sql +++ b/coderd/database/dump.sql @@ -1576,7 +1576,18 @@ CREATE TABLE aibridge_token_usages ( metadata jsonb, created_at timestamp with time zone NOT NULL, cache_read_input_tokens bigint DEFAULT 0 NOT NULL, - cache_write_input_tokens bigint DEFAULT 0 NOT NULL + cache_write_input_tokens bigint DEFAULT 0 NOT NULL, + effective_group_id uuid, + input_price_micros bigint, + output_price_micros bigint, + cache_read_price_micros bigint, + cache_write_price_micros bigint, + cost_micros bigint, + CONSTRAINT aibridge_token_usages_cache_read_price_micros_check CHECK ((cache_read_price_micros >= 0)), + CONSTRAINT aibridge_token_usages_cache_write_price_micros_check CHECK ((cache_write_price_micros >= 0)), + CONSTRAINT aibridge_token_usages_cost_micros_check CHECK ((cost_micros >= 0)), + CONSTRAINT aibridge_token_usages_input_price_micros_check CHECK ((input_price_micros >= 0)), + CONSTRAINT aibridge_token_usages_output_price_micros_check CHECK ((output_price_micros >= 0)) ); COMMENT ON TABLE aibridge_token_usages IS 'Audit log of tokens used by intercepted requests in AI Bridge'; diff --git a/coderd/database/migrations/000524_aibridge_token_usage_cost.down.sql b/coderd/database/migrations/000524_aibridge_token_usage_cost.down.sql new file mode 100644 index 00000000000..fc6b205af16 --- /dev/null +++ b/coderd/database/migrations/000524_aibridge_token_usage_cost.down.sql @@ -0,0 +1,7 @@ +ALTER TABLE aibridge_token_usages + DROP COLUMN effective_group_id, + DROP COLUMN input_price_micros, + DROP COLUMN output_price_micros, + DROP COLUMN cache_read_price_micros, + DROP COLUMN cache_write_price_micros, + DROP COLUMN cost_micros; diff --git a/coderd/database/migrations/000524_aibridge_token_usage_cost.up.sql b/coderd/database/migrations/000524_aibridge_token_usage_cost.up.sql new file mode 100644 index 00000000000..2e24f34d226 --- /dev/null +++ b/coderd/database/migrations/000524_aibridge_token_usage_cost.up.sql @@ -0,0 +1,15 @@ +ALTER TABLE aibridge_token_usages + -- Effective group this interception's spend is attributed to. NULL if the + -- user has no effective group (no budget configured). Intentionally not a + -- foreign key: this is an immutable historical attribution that must + -- survive group deletion, so the id is retained even after the group is gone. + ADD COLUMN effective_group_id UUID, + -- Snapshotted prices at interception time, in micro-units per million + -- tokens. NULL if the model is not present in ai_model_prices. + ADD COLUMN input_price_micros BIGINT CHECK (input_price_micros >= 0), + ADD COLUMN output_price_micros BIGINT CHECK (output_price_micros >= 0), + ADD COLUMN cache_read_price_micros BIGINT CHECK (cache_read_price_micros >= 0), + ADD COLUMN cache_write_price_micros BIGINT CHECK (cache_write_price_micros >= 0), + -- Computed cost in micro-units at interception time. NULL if the model is + -- not present in ai_model_prices. + ADD COLUMN cost_micros BIGINT CHECK (cost_micros >= 0); diff --git a/coderd/database/models.go b/coderd/database/models.go index 705c4418162..617da5aa16c 100644 --- a/coderd/database/models.go +++ b/coderd/database/models.go @@ -4571,6 +4571,12 @@ type AIBridgeTokenUsage struct { CreatedAt time.Time `db:"created_at" json:"created_at"` CacheReadInputTokens int64 `db:"cache_read_input_tokens" json:"cache_read_input_tokens"` CacheWriteInputTokens int64 `db:"cache_write_input_tokens" json:"cache_write_input_tokens"` + EffectiveGroupID uuid.NullUUID `db:"effective_group_id" json:"effective_group_id"` + InputPriceMicros sql.NullInt64 `db:"input_price_micros" json:"input_price_micros"` + OutputPriceMicros sql.NullInt64 `db:"output_price_micros" json:"output_price_micros"` + CacheReadPriceMicros sql.NullInt64 `db:"cache_read_price_micros" json:"cache_read_price_micros"` + CacheWritePriceMicros sql.NullInt64 `db:"cache_write_price_micros" json:"cache_write_price_micros"` + CostMicros sql.NullInt64 `db:"cost_micros" json:"cost_micros"` } // Audit log of tool calls in intercepted requests in AI Bridge diff --git a/coderd/database/queries.sql.go b/coderd/database/queries.sql.go index d28ed289a3a..4cb67cbe96c 100644 --- a/coderd/database/queries.sql.go +++ b/coderd/database/queries.sql.go @@ -1216,7 +1216,7 @@ func (q *sqlQuerier) GetAIBridgeInterceptions(ctx context.Context) ([]AIBridgeIn const getAIBridgeTokenUsagesByInterceptionID = `-- name: GetAIBridgeTokenUsagesByInterceptionID :many SELECT - id, interception_id, provider_response_id, input_tokens, output_tokens, metadata, created_at, cache_read_input_tokens, cache_write_input_tokens + id, interception_id, provider_response_id, input_tokens, output_tokens, metadata, created_at, cache_read_input_tokens, cache_write_input_tokens, effective_group_id, input_price_micros, output_price_micros, cache_read_price_micros, cache_write_price_micros, cost_micros FROM aibridge_token_usages WHERE interception_id = $1::uuid ORDER BY @@ -1243,6 +1243,12 @@ func (q *sqlQuerier) GetAIBridgeTokenUsagesByInterceptionID(ctx context.Context, &i.CreatedAt, &i.CacheReadInputTokens, &i.CacheWriteInputTokens, + &i.EffectiveGroupID, + &i.InputPriceMicros, + &i.OutputPriceMicros, + &i.CacheReadPriceMicros, + &i.CacheWritePriceMicros, + &i.CostMicros, ); err != nil { return nil, err } @@ -1452,11 +1458,13 @@ func (q *sqlQuerier) InsertAIBridgeModelThought(ctx context.Context, arg InsertA const insertAIBridgeTokenUsage = `-- name: InsertAIBridgeTokenUsage :one INSERT INTO aibridge_token_usages ( - id, interception_id, provider_response_id, input_tokens, output_tokens, cache_read_input_tokens, cache_write_input_tokens, metadata, created_at + id, interception_id, provider_response_id, input_tokens, output_tokens, cache_read_input_tokens, cache_write_input_tokens, metadata, created_at, + effective_group_id, input_price_micros, output_price_micros, cache_read_price_micros, cache_write_price_micros, cost_micros ) VALUES ( - $1, $2, $3, $4, $5, $6, $7, COALESCE($8::jsonb, '{}'::jsonb), $9 + $1, $2, $3, $4, $5, $6, $7, COALESCE($8::jsonb, '{}'::jsonb), $9, + $10, $11, $12, $13, $14, $15 ) -RETURNING id, interception_id, provider_response_id, input_tokens, output_tokens, metadata, created_at, cache_read_input_tokens, cache_write_input_tokens +RETURNING id, interception_id, provider_response_id, input_tokens, output_tokens, metadata, created_at, cache_read_input_tokens, cache_write_input_tokens, effective_group_id, input_price_micros, output_price_micros, cache_read_price_micros, cache_write_price_micros, cost_micros ` type InsertAIBridgeTokenUsageParams struct { @@ -1469,6 +1477,12 @@ type InsertAIBridgeTokenUsageParams struct { CacheWriteInputTokens int64 `db:"cache_write_input_tokens" json:"cache_write_input_tokens"` Metadata json.RawMessage `db:"metadata" json:"metadata"` CreatedAt time.Time `db:"created_at" json:"created_at"` + EffectiveGroupID uuid.NullUUID `db:"effective_group_id" json:"effective_group_id"` + InputPriceMicros sql.NullInt64 `db:"input_price_micros" json:"input_price_micros"` + OutputPriceMicros sql.NullInt64 `db:"output_price_micros" json:"output_price_micros"` + CacheReadPriceMicros sql.NullInt64 `db:"cache_read_price_micros" json:"cache_read_price_micros"` + CacheWritePriceMicros sql.NullInt64 `db:"cache_write_price_micros" json:"cache_write_price_micros"` + CostMicros sql.NullInt64 `db:"cost_micros" json:"cost_micros"` } func (q *sqlQuerier) InsertAIBridgeTokenUsage(ctx context.Context, arg InsertAIBridgeTokenUsageParams) (AIBridgeTokenUsage, error) { @@ -1482,6 +1496,12 @@ func (q *sqlQuerier) InsertAIBridgeTokenUsage(ctx context.Context, arg InsertAIB arg.CacheWriteInputTokens, arg.Metadata, arg.CreatedAt, + arg.EffectiveGroupID, + arg.InputPriceMicros, + arg.OutputPriceMicros, + arg.CacheReadPriceMicros, + arg.CacheWritePriceMicros, + arg.CostMicros, ) var i AIBridgeTokenUsage err := row.Scan( @@ -1494,6 +1514,12 @@ func (q *sqlQuerier) InsertAIBridgeTokenUsage(ctx context.Context, arg InsertAIB &i.CreatedAt, &i.CacheReadInputTokens, &i.CacheWriteInputTokens, + &i.EffectiveGroupID, + &i.InputPriceMicros, + &i.OutputPriceMicros, + &i.CacheReadPriceMicros, + &i.CacheWritePriceMicros, + &i.CostMicros, ) return i, err } @@ -2162,7 +2188,7 @@ func (q *sqlQuerier) ListAIBridgeSessions(ctx context.Context, arg ListAIBridgeS const listAIBridgeTokenUsagesByInterceptionIDs = `-- name: ListAIBridgeTokenUsagesByInterceptionIDs :many SELECT - id, interception_id, provider_response_id, input_tokens, output_tokens, metadata, created_at, cache_read_input_tokens, cache_write_input_tokens + id, interception_id, provider_response_id, input_tokens, output_tokens, metadata, created_at, cache_read_input_tokens, cache_write_input_tokens, effective_group_id, input_price_micros, output_price_micros, cache_read_price_micros, cache_write_price_micros, cost_micros FROM aibridge_token_usages WHERE @@ -2191,6 +2217,12 @@ func (q *sqlQuerier) ListAIBridgeTokenUsagesByInterceptionIDs(ctx context.Contex &i.CreatedAt, &i.CacheReadInputTokens, &i.CacheWriteInputTokens, + &i.EffectiveGroupID, + &i.InputPriceMicros, + &i.OutputPriceMicros, + &i.CacheReadPriceMicros, + &i.CacheWritePriceMicros, + &i.CostMicros, ); err != nil { return nil, err } diff --git a/coderd/database/queries/aibridge.sql b/coderd/database/queries/aibridge.sql index 73996dfa21a..e4ebf08e64c 100644 --- a/coderd/database/queries/aibridge.sql +++ b/coderd/database/queries/aibridge.sql @@ -38,9 +38,11 @@ WHERE aibridge_interceptions.id = ( -- name: InsertAIBridgeTokenUsage :one INSERT INTO aibridge_token_usages ( - id, interception_id, provider_response_id, input_tokens, output_tokens, cache_read_input_tokens, cache_write_input_tokens, metadata, created_at + id, interception_id, provider_response_id, input_tokens, output_tokens, cache_read_input_tokens, cache_write_input_tokens, metadata, created_at, + effective_group_id, input_price_micros, output_price_micros, cache_read_price_micros, cache_write_price_micros, cost_micros ) VALUES ( - @id, @interception_id, @provider_response_id, @input_tokens, @output_tokens, @cache_read_input_tokens, @cache_write_input_tokens, COALESCE(@metadata::jsonb, '{}'::jsonb), @created_at + @id, @interception_id, @provider_response_id, @input_tokens, @output_tokens, @cache_read_input_tokens, @cache_write_input_tokens, COALESCE(@metadata::jsonb, '{}'::jsonb), @created_at, + @effective_group_id, @input_price_micros, @output_price_micros, @cache_read_price_micros, @cache_write_price_micros, @cost_micros ) RETURNING *; diff --git a/codersdk/deployment.go b/codersdk/deployment.go index 1aec8ff8a12..5916ef7ab1d 100644 --- a/codersdk/deployment.go +++ b/codersdk/deployment.go @@ -589,6 +589,15 @@ var AIBudgetPolicies = []string{ string(AIBudgetPolicyHighest), } +// NewAIBudgetPolicyFromString converts s to an AIBudgetPolicy, falling back to +// AIBudgetPolicyHighest when s is empty or not a recognized policy. +func NewAIBudgetPolicyFromString(s string) AIBudgetPolicy { + if slices.Contains(AIBudgetPolicies, s) { + return AIBudgetPolicy(s) + } + return AIBudgetPolicyHighest +} + // AIBudgetPeriod determines when accumulated AI spend resets to zero, // aligned to UTC calendar boundaries. type AIBudgetPeriod string diff --git a/codersdk/deployment_test.go b/codersdk/deployment_test.go index 1abaa5e0005..a70fef938a6 100644 --- a/codersdk/deployment_test.go +++ b/codersdk/deployment_test.go @@ -1249,6 +1249,27 @@ func TestAIBudgetConfigParsing(t *testing.T) { }) } +func TestNewAIBudgetPolicyFromString(t *testing.T) { + t.Parallel() + + tests := []struct { + name string + in string + want codersdk.AIBudgetPolicy + }{ + {name: "supported", in: "highest", want: codersdk.AIBudgetPolicyHighest}, + {name: "empty falls back to highest", in: "", want: codersdk.AIBudgetPolicyHighest}, + {name: "unknown falls back to highest", in: "unsupported", want: codersdk.AIBudgetPolicyHighest}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + t.Parallel() + assert.Equal(t, tt.want, codersdk.NewAIBudgetPolicyFromString(tt.in)) + }) + } +} + func TestComputeMaxIdleConns(t *testing.T) { t.Parallel()