From 1031da9738a22cc7158f7b3e02e31933e7771243 Mon Sep 17 00:00:00 2001 From: Michael Suchacz <203725896+ibetitsmike@users.noreply.github.com> Date: Tue, 17 Mar 2026 01:24:03 +0100 Subject: [PATCH] feat: add agent chat spend limiting (backend) (#23071) Introduces deployment-scoped spend limiting for Coder Agents, enabling administrators to control LLM costs at global, group, and individual user levels. ## Changes - **Database migration (000437)**: `chat_usage_limit_config` (singleton), `chat_usage_limit_overrides` (per-user), `chat_usage_limit_group_overrides` (per-group) - **Single-query limit resolution**: individual override > min(group) > global default via `ResolveUserChatSpendLimit` - **Fail-open enforcement** in chatd with documented TOCTOU trade-off - **Experimental API** under `/api/experimental/chats/usage-limits` for CRUD on limits - **`AsChatd` RBAC subject** for narrowly-scoped daemon access (replaces `AsSystemRestricted`) - **Generated TypeScript types** for the frontend SDK ## Hierarchy 1. Individual user override (highest) 2. Minimum of group limits 3. Global default 4. Disabled / unlimited Currency stored as micro-dollars (`1,000,000` = $1.00). Frontend PR: #23072 --- coderd/chatd/chatd.go | 191 +++++-- coderd/chatd/chatd_test.go | 512 +++++++++++++++++ coderd/chatd/usagelimit.go | 128 +++++ coderd/chatd/usagelimit_test.go | 132 +++++ coderd/chats.go | 533 +++++++++++++++++- coderd/chats_test.go | 388 +++++++++++++ coderd/coderd.go | 13 + coderd/database/check_constraint.go | 41 +- coderd/database/dbauthz/dbauthz.go | 98 ++++ coderd/database/dbauthz/dbauthz_test.go | 140 +++++ coderd/database/dbmetrics/querymetrics.go | 112 ++++ coderd/database/dbmock/dbmock.go | 208 +++++++ coderd/database/dump.sql | 38 +- .../000441_chat_usage_limits.down.sql | 4 + .../000441_chat_usage_limits.up.sql | 32 ++ .../fixtures/000441_chat_usage_limits.up.sql | 5 + coderd/database/modelqueries.go | 1 + coderd/database/models.go | 16 +- coderd/database/querier.go | 23 + coderd/database/queries.sql.go | 426 +++++++++++++- coderd/database/queries/chats.sql | 125 ++++ coderd/database/unique_constraint.go | 2 + codersdk/chats.go | 347 +++++++++++- codersdk/chats_test.go | 80 +++ docs/admin/security/audit-logs.md | 4 +- enterprise/audit/table.go | 18 +- .../converted_state.state.golden | 2 +- .../devcontainer-multiple-agents.tfstate.json | 2 +- .../converted_state.plan.golden | 3 +- .../duplicate-env-keys.tfplan.json | 298 ++++++++-- site/src/api/typesGenerated.ts | 129 +++++ 31 files changed, 3904 insertions(+), 147 deletions(-) create mode 100644 coderd/chatd/usagelimit.go create mode 100644 coderd/chatd/usagelimit_test.go create mode 100644 coderd/database/migrations/000441_chat_usage_limits.down.sql create mode 100644 coderd/database/migrations/000441_chat_usage_limits.up.sql create mode 100644 coderd/database/migrations/testdata/fixtures/000441_chat_usage_limits.up.sql diff --git a/coderd/chatd/chatd.go b/coderd/chatd/chatd.go index a00d4104af..ef8c9aafdb 100644 --- a/coderd/chatd/chatd.go +++ b/coderd/chatd/chatd.go @@ -14,6 +14,7 @@ import ( "charm.land/fantasy" "charm.land/fantasy/providers/anthropic" "github.com/google/uuid" + "github.com/shopspring/decimal" "github.com/sqlc-dev/pqtype" "golang.org/x/sync/errgroup" "golang.org/x/xerrors" @@ -172,6 +173,27 @@ var ( errChatTakenByOtherWorker = xerrors.New("chat acquired by another worker") ) +// UsageLimitExceededError indicates the user has exceeded their chat spend +// limit. +type UsageLimitExceededError struct { + LimitMicros int64 + ConsumedMicros int64 + PeriodEnd time.Time +} + +func formatMicrosAsDollars(micros int64) string { + return "$" + decimal.NewFromInt(micros).Shift(-6).StringFixed(2) +} + +func (e *UsageLimitExceededError) Error() string { + return fmt.Sprintf( + "usage limit exceeded: spent %s of %s limit, resets at %s", + formatMicrosAsDollars(e.ConsumedMicros), + formatMicrosAsDollars(e.LimitMicros), + e.PeriodEnd.Format(time.RFC3339), + ) +} + // CreateOptions controls chat creation in the shared chat mutation path. type CreateOptions struct { OwnerID uuid.UUID @@ -257,6 +279,10 @@ func (p *Server) CreateChat(ctx context.Context, opts CreateOptions) (database.C var chat database.Chat txErr := p.db.InTx(func(tx database.Store) error { + if limitErr := p.checkUsageLimit(ctx, tx, opts.OwnerID); limitErr != nil { + return limitErr + } + insertedChat, err := tx.InsertChat(ctx, database.InsertChatParams{ OwnerID: opts.OwnerID, WorkspaceID: opts.WorkspaceID, @@ -389,24 +415,33 @@ func (p *Server) SendMessage( if err != nil { return xerrors.Errorf("lock chat: %w", err) } + + // Enforce usage limits before queueing or inserting. + if limitErr := p.checkUsageLimit(ctx, tx, lockedChat.OwnerID); limitErr != nil { + return limitErr + } + modelConfigID := lockedChat.LastModelConfigID if opts.ModelConfigID != nil { modelConfigID = *opts.ModelConfigID } + existingQueued, err := tx.GetChatQueuedMessages(ctx, opts.ChatID) + if err != nil { + return xerrors.Errorf("get queued messages: %w", err) + } + // Both queue and interrupt behaviors queue messages - // when the chat is busy. Interrupt additionally - // signals the running loop to stop so the queued - // message is promoted sooner. Crucially, this - // guarantees the interrupted assistant response is - // persisted (with a lower id/created_at) before the - // user message is promoted into chat_messages, - // preserving correct conversation order. - if shouldQueueUserMessage(lockedChat.Status) { - existingQueued, err := tx.GetChatQueuedMessages(ctx, opts.ChatID) - if err != nil { - return xerrors.Errorf("get queued messages: %w", err) - } + // when the chat is busy. We also keep queueing while a + // backlog exists so waiting chats blocked by spend limits + // preserve FIFO user-message order. Interrupt additionally + // signals the running loop to stop so the queued message + // is promoted sooner. Crucially, this guarantees the + // interrupted assistant response is persisted (with a + // lower id/created_at) before the user message is + // promoted into chat_messages, preserving correct + // conversation order. + if shouldQueueUserMessage(lockedChat.Status) || len(existingQueued) > 0 { if len(existingQueued) >= MaxQueueSize { return ErrMessageQueueFull } @@ -492,6 +527,31 @@ func (p *Server) SendMessage( return result, nil } +func (p *Server) checkUsageLimit(ctx context.Context, store database.Store, ownerID uuid.UUID) error { + status, err := ResolveUsageLimitStatus(ctx, store, ownerID, time.Now()) + if err != nil { + // Fail open: never block chat due to a limit-resolution failure. + p.logger.Warn(ctx, "usage limit check failed, allowing message", + slog.F("owner_id", ownerID), + slog.Error(err), + ) + return nil + } + if status == nil { + return nil + } + // Block when current spend reaches or exceeds limit (>= ensures + // the user cannot start new conversations once the limit is hit). + if status.SpendLimitMicros != nil && status.CurrentSpend >= *status.SpendLimitMicros { + return &UsageLimitExceededError{ + LimitMicros: *status.SpendLimitMicros, + ConsumedMicros: status.CurrentSpend, + PeriodEnd: status.PeriodEnd, + } + } + return nil +} + // EditMessage updates a user message in-place, truncates all following messages, // clears queued messages, and moves the chat into pending status. func (p *Server) EditMessage( @@ -515,11 +575,15 @@ func (p *Server) EditMessage( var result EditMessageResult txErr := p.db.InTx(func(tx database.Store) error { - _, err := tx.GetChatByIDForUpdate(ctx, opts.ChatID) + lockedChat, err := tx.GetChatByIDForUpdate(ctx, opts.ChatID) if err != nil { return xerrors.Errorf("lock chat: %w", err) } + if limitErr := p.checkUsageLimit(ctx, tx, lockedChat.OwnerID); limitErr != nil { + return limitErr + } + existing, err := tx.GetChatMessageByID(ctx, opts.EditedMessageID) if err != nil { if errors.Is(err, sql.ErrNoRows) { @@ -1849,6 +1913,62 @@ func (p *Server) chatFileResolver() chatprompt.FileResolver { } } +// tryAutoPromoteQueuedMessage pops the next queued message and converts it +// into a pending user message inside the caller's transaction. Queued +// messages were already admitted through SendMessage, so this preserves FIFO +// order without re-checking usage limits. +func (p *Server) tryAutoPromoteQueuedMessage( + ctx context.Context, + tx database.Store, + chat database.Chat, +) (*database.ChatMessage, []database.ChatQueuedMessage, bool, error) { + logger := p.logger.With(slog.F("chat_id", chat.ID)) + + nextQueued, err := tx.PopNextQueuedMessage(ctx, chat.ID) + if errors.Is(err, sql.ErrNoRows) { + return nil, nil, false, nil + } + if err != nil { + return nil, nil, false, xerrors.Errorf("pop next queued message: %w", err) + } + + msg, err := insertChatMessageWithStore(ctx, tx, database.InsertChatMessageParams{ + ChatID: chat.ID, + ModelConfigID: uuid.NullUUID{UUID: chat.LastModelConfigID, Valid: true}, + Role: database.ChatMessageRoleUser, + ContentVersion: chatprompt.CurrentContentVersion, + Content: pqtype.NullRawMessage{ + RawMessage: nextQueued.Content, + Valid: len(nextQueued.Content) > 0, + }, + CreatedBy: uuid.NullUUID{UUID: chat.OwnerID, Valid: chat.OwnerID != uuid.Nil}, + Visibility: database.ChatMessageVisibilityBoth, + InputTokens: sql.NullInt64{}, + OutputTokens: sql.NullInt64{}, + TotalTokens: sql.NullInt64{}, + ReasoningTokens: sql.NullInt64{}, + CacheCreationTokens: sql.NullInt64{}, + CacheReadTokens: sql.NullInt64{}, + ContextLimit: sql.NullInt64{}, + TotalCostMicros: sql.NullInt64{}, + Compressed: sql.NullBool{}, + }) + if err != nil { + logger.Error(ctx, "failed to promote queued message", + slog.F("queued_message_id", nextQueued.ID), slog.Error(err)) + return nil, nil, false, nil + } + + remainingQueuedMessages, err := tx.GetChatQueuedMessages(ctx, chat.ID) + if err != nil { + logger.Error(ctx, "failed to load remaining queued messages after auto-promotion", + slog.F("queued_message_id", nextQueued.ID), slog.Error(err)) + return &msg, nil, false, nil + } + + return &msg, remainingQueuedMessages, true, nil +} + func (p *Server) processChat(ctx context.Context, chat database.Chat) { logger := p.logger.With(slog.F("chat_id", chat.ID)) logger.Info(ctx, "processing chat request") @@ -1964,43 +2084,14 @@ func (p *Server) processChat(ctx context.Context, chat database.Chat) { if latestChat.Status == database.ChatStatusPending { status = database.ChatStatusPending } else if status == database.ChatStatusWaiting { - // Try to auto-promote the next queued message. - nextQueued, popErr := tx.PopNextQueuedMessage(cleanupCtx, chat.ID) - if popErr == nil { - msg, insertErr := tx.InsertChatMessage(cleanupCtx, database.InsertChatMessageParams{ - ChatID: chat.ID, - ModelConfigID: uuid.NullUUID{UUID: latestChat.LastModelConfigID, Valid: true}, - Role: database.ChatMessageRoleUser, - ContentVersion: chatprompt.CurrentContentVersion, - Content: pqtype.NullRawMessage{ - RawMessage: nextQueued.Content, - Valid: len(nextQueued.Content) > 0, - }, - CreatedBy: uuid.NullUUID{UUID: chat.OwnerID, Valid: chat.OwnerID != uuid.Nil}, - Visibility: database.ChatMessageVisibilityBoth, - InputTokens: sql.NullInt64{}, - OutputTokens: sql.NullInt64{}, - TotalTokens: sql.NullInt64{}, - ReasoningTokens: sql.NullInt64{}, - CacheCreationTokens: sql.NullInt64{}, - CacheReadTokens: sql.NullInt64{}, - ContextLimit: sql.NullInt64{}, - TotalCostMicros: sql.NullInt64{}, - Compressed: sql.NullBool{}, - }) - if insertErr != nil { - logger.Error(cleanupCtx, "failed to promote queued message", - slog.F("queued_message_id", nextQueued.ID), slog.Error(insertErr)) - } else { - status = database.ChatStatusPending - promotedMessage = &msg - - remaining, qErr := tx.GetChatQueuedMessages(cleanupCtx, chat.ID) - if qErr == nil { - remainingQueuedMessages = remaining - shouldPublishQueueUpdate = true - } - } + // Queued messages were already admitted through SendMessage, + // so auto-promotion only preserves FIFO order here. + var promoteErr error + promotedMessage, remainingQueuedMessages, shouldPublishQueueUpdate, promoteErr = p.tryAutoPromoteQueuedMessage(cleanupCtx, tx, latestChat) + if promoteErr != nil { + logger.Error(cleanupCtx, "failed to auto-promote queued message", slog.Error(promoteErr)) + } else if promotedMessage != nil { + status = database.ChatStatusPending } } diff --git a/coderd/chatd/chatd_test.go b/coderd/chatd/chatd_test.go index fff0754e49..11920a042c 100644 --- a/coderd/chatd/chatd_test.go +++ b/coderd/chatd/chatd_test.go @@ -374,6 +374,72 @@ func TestSendMessageQueueBehaviorQueuesWhenBusy(t *testing.T) { require.Len(t, messages, 1) } +func TestSendMessageQueuesWhenWaitingWithQueuedBacklog(t *testing.T) { + t.Parallel() + + db, ps := dbtestutil.NewDB(t) + replica := newTestServer(t, db, ps, uuid.New()) + + ctx := testutil.Context(t, testutil.WaitLong) + user, model := seedChatDependencies(ctx, t, db) + + chat, err := replica.CreateChat(ctx, chatd.CreateOptions{ + OwnerID: user.ID, + Title: "queue-when-waiting-with-backlog", + ModelConfigID: model.ID, + InitialUserContent: []codersdk.ChatMessagePart{codersdk.ChatMessageText("hello")}, + }) + require.NoError(t, err) + + queuedContent, err := json.Marshal([]codersdk.ChatMessagePart{ + codersdk.ChatMessageText("older queued"), + }) + require.NoError(t, err) + _, err = db.InsertChatQueuedMessage(ctx, database.InsertChatQueuedMessageParams{ + ChatID: chat.ID, + Content: queuedContent, + }) + require.NoError(t, err) + + chat, err = db.UpdateChatStatus(ctx, database.UpdateChatStatusParams{ + ID: chat.ID, + Status: database.ChatStatusWaiting, + WorkerID: uuid.NullUUID{}, + StartedAt: sql.NullTime{}, + HeartbeatAt: sql.NullTime{}, + LastError: sql.NullString{}, + }) + require.NoError(t, err) + + result, err := replica.SendMessage(ctx, chatd.SendMessageOptions{ + ChatID: chat.ID, + Content: []codersdk.ChatMessagePart{codersdk.ChatMessageText("newer queued")}, + }) + require.NoError(t, err) + require.True(t, result.Queued) + require.NotNil(t, result.QueuedMessage) + require.Equal(t, database.ChatStatusWaiting, result.Chat.Status) + + queued, err := db.GetChatQueuedMessages(ctx, chat.ID) + require.NoError(t, err) + require.Len(t, queued, 2) + + olderSDK := db2sdk.ChatQueuedMessage(queued[0]) + require.Len(t, olderSDK.Content, 1) + require.Equal(t, "older queued", olderSDK.Content[0].Text) + + newerSDK := db2sdk.ChatQueuedMessage(queued[1]) + require.Len(t, newerSDK.Content, 1) + require.Equal(t, "newer queued", newerSDK.Content[0].Text) + + messages, err := db.GetChatMessagesByChatID(ctx, database.GetChatMessagesByChatIDParams{ + ChatID: chat.ID, + AfterID: 0, + }) + require.NoError(t, err) + require.Len(t, messages, 1) +} + func TestSendMessageInterruptBehaviorQueuesAndInterruptsWhenBusy(t *testing.T) { t.Parallel() @@ -525,6 +591,452 @@ func TestEditMessageUpdatesAndTruncatesAndClearsQueue(t *testing.T) { require.False(t, chatFromDB.WorkerID.Valid) } +func TestCreateChatRejectsWhenUsageLimitReached(t *testing.T) { + t.Parallel() + + db, ps := dbtestutil.NewDB(t) + replica := newTestServer(t, db, ps, uuid.New()) + + ctx := testutil.Context(t, testutil.WaitLong) + user, model := seedChatDependencies(ctx, t, db) + + _, err := db.UpsertChatUsageLimitConfig(ctx, database.UpsertChatUsageLimitConfigParams{ + Enabled: true, + DefaultLimitMicros: 100, + Period: string(codersdk.ChatUsageLimitPeriodDay), + }) + require.NoError(t, err) + + existingChat, err := db.InsertChat(ctx, database.InsertChatParams{ + OwnerID: user.ID, + Title: "existing-limit-chat", + LastModelConfigID: model.ID, + }) + require.NoError(t, err) + + assistantContent, err := chatprompt.MarshalParts([]codersdk.ChatMessagePart{ + codersdk.ChatMessageText("assistant"), + }) + require.NoError(t, err) + + _, err = db.InsertChatMessage(ctx, database.InsertChatMessageParams{ + ChatID: existingChat.ID, + ModelConfigID: uuid.NullUUID{UUID: model.ID, Valid: true}, + Role: database.ChatMessageRoleAssistant, + ContentVersion: chatprompt.CurrentContentVersion, + Content: assistantContent, + Visibility: database.ChatMessageVisibilityBoth, + InputTokens: sql.NullInt64{}, + OutputTokens: sql.NullInt64{}, + TotalTokens: sql.NullInt64{}, + ReasoningTokens: sql.NullInt64{}, + CacheCreationTokens: sql.NullInt64{}, + CacheReadTokens: sql.NullInt64{}, + ContextLimit: sql.NullInt64{}, + Compressed: sql.NullBool{}, + TotalCostMicros: sql.NullInt64{Int64: 100, Valid: true}, + }) + require.NoError(t, err) + + beforeChats, err := db.GetChatsByOwnerID(ctx, database.GetChatsByOwnerIDParams{ + OwnerID: user.ID, + AfterID: uuid.Nil, + OffsetOpt: 0, + LimitOpt: 100, + }) + require.NoError(t, err) + require.Len(t, beforeChats, 1) + + _, err = replica.CreateChat(ctx, chatd.CreateOptions{ + OwnerID: user.ID, + Title: "over-limit", + ModelConfigID: model.ID, + InitialUserContent: []codersdk.ChatMessagePart{codersdk.ChatMessageText("hello")}, + }) + require.Error(t, err) + + var limitErr *chatd.UsageLimitExceededError + require.ErrorAs(t, err, &limitErr) + require.Equal(t, int64(100), limitErr.LimitMicros) + require.Equal(t, int64(100), limitErr.ConsumedMicros) + + afterChats, err := db.GetChatsByOwnerID(ctx, database.GetChatsByOwnerIDParams{ + OwnerID: user.ID, + AfterID: uuid.Nil, + OffsetOpt: 0, + LimitOpt: 100, + }) + require.NoError(t, err) + require.Len(t, afterChats, len(beforeChats)) +} + +func TestPromoteQueuedAllowsAlreadyQueuedMessageWhenUsageLimitReached(t *testing.T) { + t.Parallel() + + db, ps := dbtestutil.NewDB(t) + replica := newTestServer(t, db, ps, uuid.New()) + + ctx := testutil.Context(t, testutil.WaitLong) + user, model := seedChatDependencies(ctx, t, db) + + _, err := db.UpsertChatUsageLimitConfig(ctx, database.UpsertChatUsageLimitConfigParams{ + Enabled: true, + DefaultLimitMicros: 100, + Period: string(codersdk.ChatUsageLimitPeriodDay), + }) + require.NoError(t, err) + + chat, err := replica.CreateChat(ctx, chatd.CreateOptions{ + OwnerID: user.ID, + Title: "queued-limit-reached", + ModelConfigID: model.ID, + InitialUserContent: []codersdk.ChatMessagePart{codersdk.ChatMessageText("hello")}, + }) + require.NoError(t, err) + + chat, err = db.UpdateChatStatus(ctx, database.UpdateChatStatusParams{ + ID: chat.ID, + Status: database.ChatStatusRunning, + WorkerID: uuid.NullUUID{UUID: uuid.New(), Valid: true}, + StartedAt: sql.NullTime{Time: time.Now(), Valid: true}, + HeartbeatAt: sql.NullTime{Time: time.Now(), Valid: true}, + }) + require.NoError(t, err) + + queuedResult, err := replica.SendMessage(ctx, chatd.SendMessageOptions{ + ChatID: chat.ID, + Content: []codersdk.ChatMessagePart{codersdk.ChatMessageText("queued")}, + BusyBehavior: chatd.SendMessageBusyBehaviorQueue, + }) + require.NoError(t, err) + require.True(t, queuedResult.Queued) + require.NotNil(t, queuedResult.QueuedMessage) + + assistantContent, err := chatprompt.MarshalParts([]codersdk.ChatMessagePart{ + codersdk.ChatMessageText("assistant"), + }) + require.NoError(t, err) + + _, err = db.InsertChatMessage(ctx, database.InsertChatMessageParams{ + ChatID: chat.ID, + ModelConfigID: uuid.NullUUID{UUID: model.ID, Valid: true}, + Role: database.ChatMessageRoleAssistant, + ContentVersion: chatprompt.CurrentContentVersion, + Content: assistantContent, + Visibility: database.ChatMessageVisibilityBoth, + InputTokens: sql.NullInt64{}, + OutputTokens: sql.NullInt64{}, + TotalTokens: sql.NullInt64{}, + ReasoningTokens: sql.NullInt64{}, + CacheCreationTokens: sql.NullInt64{}, + CacheReadTokens: sql.NullInt64{}, + ContextLimit: sql.NullInt64{}, + Compressed: sql.NullBool{}, + TotalCostMicros: sql.NullInt64{Int64: 100, Valid: true}, + }) + require.NoError(t, err) + + chat, err = db.UpdateChatStatus(ctx, database.UpdateChatStatusParams{ + ID: chat.ID, + Status: database.ChatStatusWaiting, + WorkerID: uuid.NullUUID{}, + StartedAt: sql.NullTime{}, + HeartbeatAt: sql.NullTime{}, + LastError: sql.NullString{}, + }) + require.NoError(t, err) + + result, err := replica.PromoteQueued(ctx, chatd.PromoteQueuedOptions{ + ChatID: chat.ID, + QueuedMessageID: queuedResult.QueuedMessage.ID, + CreatedBy: user.ID, + }) + require.NoError(t, err) + require.Equal(t, database.ChatMessageRoleUser, result.PromotedMessage.Role) + + chat, err = db.GetChatByID(ctx, chat.ID) + require.NoError(t, err) + require.Equal(t, database.ChatStatusPending, chat.Status) + + queued, err := db.GetChatQueuedMessages(ctx, chat.ID) + require.NoError(t, err) + require.Empty(t, queued) + + messages, err := db.GetChatMessagesByChatID(ctx, database.GetChatMessagesByChatIDParams{ + ChatID: chat.ID, + AfterID: 0, + }) + require.NoError(t, err) + require.Len(t, messages, 3) + require.Equal(t, database.ChatMessageRoleUser, messages[2].Role) +} + +func TestInterruptAutoPromotionIgnoresLaterUsageLimitIncrease(t *testing.T) { + t.Parallel() + + db, ps := dbtestutil.NewDB(t) + ctx := testutil.Context(t, testutil.WaitLong) + + _, err := db.UpsertChatUsageLimitConfig(ctx, database.UpsertChatUsageLimitConfigParams{ + Enabled: true, + DefaultLimitMicros: 100, + Period: string(codersdk.ChatUsageLimitPeriodDay), + }) + require.NoError(t, err) + + streamStarted := make(chan struct{}) + interrupted := make(chan struct{}) + allowFinish := make(chan struct{}) + var requestCount atomic.Int32 + openAIURL := chattest.NewOpenAI(t, func(req *chattest.OpenAIRequest) chattest.OpenAIResponse { + if !req.Stream { + return chattest.OpenAINonStreamingResponse("title") + } + if requestCount.Add(1) == 1 { + chunks := make(chan chattest.OpenAIChunk, 1) + go func() { + defer close(chunks) + chunks <- chattest.OpenAITextChunks("partial")[0] + select { + case <-streamStarted: + default: + close(streamStarted) + } + <-req.Context().Done() + select { + case <-interrupted: + default: + close(interrupted) + } + <-allowFinish + }() + return chattest.OpenAIResponse{StreamingChunks: chunks} + } + return chattest.OpenAIStreamingResponse( + chattest.OpenAITextChunks("done")..., + ) + }) + + logger := slogtest.Make(t, &slogtest.Options{IgnoreErrors: true}) + server := chatd.New(chatd.Config{ + Logger: logger, + Database: db, + ReplicaID: uuid.New(), + Pubsub: ps, + PendingChatAcquireInterval: 10 * time.Millisecond, + InFlightChatStaleAfter: testutil.WaitSuperLong, + }) + t.Cleanup(func() { + require.NoError(t, server.Close()) + }) + + user, model := seedChatDependencies(ctx, t, db) + setOpenAIProviderBaseURL(ctx, t, db, openAIURL) + + chat, err := server.CreateChat(ctx, chatd.CreateOptions{ + OwnerID: user.ID, + Title: "interrupt-autopromote-limit", + ModelConfigID: model.ID, + InitialUserContent: []codersdk.ChatMessagePart{codersdk.ChatMessageText("hello")}, + }) + require.NoError(t, err) + + require.Eventually(t, func() bool { + fromDB, dbErr := db.GetChatByID(ctx, chat.ID) + if dbErr != nil { + return false + } + return fromDB.Status == database.ChatStatusRunning && fromDB.WorkerID.Valid + }, testutil.WaitMedium, testutil.IntervalFast) + + require.Eventually(t, func() bool { + select { + case <-streamStarted: + return true + default: + return false + } + }, testutil.WaitMedium, testutil.IntervalFast) + + queuedResult, err := server.SendMessage(ctx, chatd.SendMessageOptions{ + ChatID: chat.ID, + Content: []codersdk.ChatMessagePart{codersdk.ChatMessageText("queued")}, + BusyBehavior: chatd.SendMessageBusyBehaviorInterrupt, + }) + require.NoError(t, err) + require.True(t, queuedResult.Queued) + require.NotNil(t, queuedResult.QueuedMessage) + + require.Eventually(t, func() bool { + select { + case <-interrupted: + return true + default: + return false + } + }, testutil.WaitMedium, testutil.IntervalFast) + + laterQueuedResult, err := server.SendMessage(ctx, chatd.SendMessageOptions{ + ChatID: chat.ID, + Content: []codersdk.ChatMessagePart{codersdk.ChatMessageText("later queued")}, + }) + require.NoError(t, err) + require.True(t, laterQueuedResult.Queued) + require.NotNil(t, laterQueuedResult.QueuedMessage) + + spendChat, err := db.InsertChat(ctx, database.InsertChatParams{ + OwnerID: user.ID, + WorkspaceID: uuid.NullUUID{}, + ParentChatID: uuid.NullUUID{}, + RootChatID: uuid.NullUUID{}, + LastModelConfigID: model.ID, + Title: "other-spend", + Mode: database.NullChatMode{}, + }) + require.NoError(t, err) + + assistantContent, err := chatprompt.MarshalParts([]codersdk.ChatMessagePart{ + codersdk.ChatMessageText("spent elsewhere"), + }) + require.NoError(t, err) + + _, err = db.InsertChatMessage(ctx, database.InsertChatMessageParams{ + ChatID: spendChat.ID, + ModelConfigID: uuid.NullUUID{UUID: model.ID, Valid: true}, + Role: database.ChatMessageRoleAssistant, + ContentVersion: chatprompt.CurrentContentVersion, + Content: assistantContent, + Visibility: database.ChatMessageVisibilityBoth, + InputTokens: sql.NullInt64{}, + OutputTokens: sql.NullInt64{}, + TotalTokens: sql.NullInt64{}, + ReasoningTokens: sql.NullInt64{}, + CacheCreationTokens: sql.NullInt64{}, + CacheReadTokens: sql.NullInt64{}, + ContextLimit: sql.NullInt64{}, + Compressed: sql.NullBool{}, + TotalCostMicros: sql.NullInt64{Int64: 100, Valid: true}, + }) + require.NoError(t, err) + + close(allowFinish) + + require.Eventually(t, func() bool { + queued, dbErr := db.GetChatQueuedMessages(ctx, chat.ID) + if dbErr != nil || len(queued) != 0 { + return false + } + + fromDB, dbErr := db.GetChatByID(ctx, chat.ID) + if dbErr != nil || fromDB.Status != database.ChatStatusWaiting { + return false + } + + messages, dbErr := db.GetChatMessagesByChatID(ctx, database.GetChatMessagesByChatIDParams{ + ChatID: chat.ID, + AfterID: 0, + }) + if dbErr != nil { + return false + } + + userTexts := make([]string, 0, 3) + for _, message := range messages { + if message.Role != database.ChatMessageRoleUser { + continue + } + sdkMessage := db2sdk.ChatMessage(message) + if len(sdkMessage.Content) != 1 { + continue + } + userTexts = append(userTexts, sdkMessage.Content[0].Text) + } + if len(userTexts) != 3 { + return false + } + return userTexts[0] == "hello" && userTexts[1] == "queued" && userTexts[2] == "later queued" + }, testutil.WaitLong, testutil.IntervalFast) +} + +func TestEditMessageRejectsWhenUsageLimitReached(t *testing.T) { + t.Parallel() + + db, ps := dbtestutil.NewDB(t) + replica := newTestServer(t, db, ps, uuid.New()) + + ctx := testutil.Context(t, testutil.WaitLong) + user, model := seedChatDependencies(ctx, t, db) + + _, err := db.UpsertChatUsageLimitConfig(ctx, database.UpsertChatUsageLimitConfigParams{ + Enabled: true, + DefaultLimitMicros: 100, + Period: string(codersdk.ChatUsageLimitPeriodDay), + }) + require.NoError(t, err) + + chat, err := replica.CreateChat(ctx, chatd.CreateOptions{ + OwnerID: user.ID, + Title: "edit-limit-reached", + ModelConfigID: model.ID, + InitialUserContent: []codersdk.ChatMessagePart{codersdk.ChatMessageText("original")}, + }) + require.NoError(t, err) + + messages, err := db.GetChatMessagesByChatID(ctx, database.GetChatMessagesByChatIDParams{ + ChatID: chat.ID, + AfterID: 0, + }) + require.NoError(t, err) + require.Len(t, messages, 1) + editedMessageID := messages[0].ID + + assistantContent, err := chatprompt.MarshalParts([]codersdk.ChatMessagePart{ + codersdk.ChatMessageText("assistant"), + }) + require.NoError(t, err) + + _, err = db.InsertChatMessage(ctx, database.InsertChatMessageParams{ + ChatID: chat.ID, + ModelConfigID: uuid.NullUUID{UUID: model.ID, Valid: true}, + Role: database.ChatMessageRoleAssistant, + ContentVersion: chatprompt.CurrentContentVersion, + Content: assistantContent, + Visibility: database.ChatMessageVisibilityBoth, + InputTokens: sql.NullInt64{}, + OutputTokens: sql.NullInt64{}, + TotalTokens: sql.NullInt64{}, + ReasoningTokens: sql.NullInt64{}, + CacheCreationTokens: sql.NullInt64{}, + CacheReadTokens: sql.NullInt64{}, + ContextLimit: sql.NullInt64{}, + Compressed: sql.NullBool{}, + TotalCostMicros: sql.NullInt64{Int64: 100, Valid: true}, + }) + require.NoError(t, err) + + _, err = replica.EditMessage(ctx, chatd.EditMessageOptions{ + ChatID: chat.ID, + EditedMessageID: editedMessageID, + Content: []codersdk.ChatMessagePart{codersdk.ChatMessageText("edited")}, + }) + require.Error(t, err) + + var limitErr *chatd.UsageLimitExceededError + require.ErrorAs(t, err, &limitErr) + require.Equal(t, int64(100), limitErr.LimitMicros) + require.Equal(t, int64(100), limitErr.ConsumedMicros) + + messages, err = db.GetChatMessagesByChatID(ctx, database.GetChatMessagesByChatIDParams{ + ChatID: chat.ID, + AfterID: 0, + }) + require.NoError(t, err) + require.Len(t, messages, 2) + originalMessage := db2sdk.ChatMessage(messages[0]) + require.Len(t, originalMessage.Content, 1) + require.Equal(t, "original", originalMessage.Content[0].Text) +} + func TestEditMessageRejectsMissingMessage(t *testing.T) { t.Parallel() diff --git a/coderd/chatd/usagelimit.go b/coderd/chatd/usagelimit.go new file mode 100644 index 0000000000..12535421d4 --- /dev/null +++ b/coderd/chatd/usagelimit.go @@ -0,0 +1,128 @@ +package chatd + +import ( + "context" + "database/sql" + "errors" + "fmt" + "time" + + "github.com/google/uuid" + "golang.org/x/xerrors" + + "github.com/coder/coder/v2/coderd/database" + "github.com/coder/coder/v2/coderd/database/dbauthz" + "github.com/coder/coder/v2/codersdk" +) + +// ComputeUsagePeriodBounds returns the UTC-aligned start and end bounds for the +// active usage-limit period containing now. +func ComputeUsagePeriodBounds(now time.Time, period codersdk.ChatUsageLimitPeriod) (start, end time.Time) { + utcNow := now.UTC() + + switch period { + case codersdk.ChatUsageLimitPeriodDay: + start = time.Date(utcNow.Year(), utcNow.Month(), utcNow.Day(), 0, 0, 0, 0, time.UTC) + end = start.AddDate(0, 0, 1) + case codersdk.ChatUsageLimitPeriodWeek: + // Walk backward to Monday of the current ISO week. + // ISO 8601 weeks always start on Monday, so this never + // crosses an ISO-week boundary. + start = time.Date(utcNow.Year(), utcNow.Month(), utcNow.Day(), 0, 0, 0, 0, time.UTC) + for start.Weekday() != time.Monday { + start = start.AddDate(0, 0, -1) + } + end = start.AddDate(0, 0, 7) + case codersdk.ChatUsageLimitPeriodMonth: + start = time.Date(utcNow.Year(), utcNow.Month(), 1, 0, 0, 0, 0, time.UTC) + end = start.AddDate(0, 1, 0) + default: + panic(fmt.Sprintf("unknown chat usage limit period: %q", period)) + } + + return start, end +} + +// ResolveUsageLimitStatus resolves the current usage-limit status for userID. +// +// Note: There is a potential race condition where two concurrent messages +// from the same user can both pass the limit check if processed in +// parallel, allowing brief overage. This is acceptable because: +// - Cost is only known after the LLM API returns. +// - Overage is bounded by message cost × concurrency. +// - Fail-open is the deliberate design choice for this feature. +// +// Architecture note: today this path enforces one period globally +// (day/week/month) from config. +// To support simultaneous periods, add nullable +// daily/weekly/monthly_limit_micros columns on override tables, where NULL +// means no limit for that period. +// Then scan spend once over the widest active window with conditional SUMs +// for each period and compare each spend/limit pair Go-side, blocking on +// whichever period is tightest. +func ResolveUsageLimitStatus(ctx context.Context, db database.Store, userID uuid.UUID, now time.Time) (*codersdk.ChatUsageLimitStatus, error) { + //nolint:gocritic // AsChatd provides narrowly-scoped daemon access for + // deployment config reads and cross-user chat spend aggregation. + authCtx := dbauthz.AsChatd(ctx) + + config, err := db.GetChatUsageLimitConfig(authCtx) + if err != nil { + if errors.Is(err, sql.ErrNoRows) { + return nil, nil //nolint:nilnil // Nil status cleanly signals disabled limits. + } + return nil, err + } + if !config.Enabled { + return nil, nil //nolint:nilnil // Nil status cleanly signals disabled limits. + } + + period, ok := mapDBPeriodToSDK(config.Period) + if !ok { + return nil, xerrors.Errorf("invalid chat usage limit period %q", config.Period) + } + + // Resolve effective limit in a single query: + // individual override > group limit > global default. + effectiveLimit, err := db.ResolveUserChatSpendLimit(authCtx, userID) + if err != nil { + return nil, err + } + // -1 means limits are disabled (shouldn't happen since we checked above, + // but handle gracefully). + if effectiveLimit < 0 { + return nil, nil //nolint:nilnil // Nil status cleanly signals disabled limits. + } + + start, end := ComputeUsagePeriodBounds(now, period) + + spendTotal, err := db.GetUserChatSpendInPeriod(authCtx, database.GetUserChatSpendInPeriodParams{ + UserID: userID, + StartTime: start, + EndTime: end, + }) + if err != nil { + return nil, err + } + + return &codersdk.ChatUsageLimitStatus{ + IsLimited: true, + Period: period, + SpendLimitMicros: &effectiveLimit, + CurrentSpend: spendTotal, + PeriodStart: start, + PeriodEnd: end, + }, nil +} + +func mapDBPeriodToSDK(dbPeriod string) (codersdk.ChatUsageLimitPeriod, bool) { + switch dbPeriod { + case string(codersdk.ChatUsageLimitPeriodDay): + return codersdk.ChatUsageLimitPeriodDay, true + case string(codersdk.ChatUsageLimitPeriodWeek): + return codersdk.ChatUsageLimitPeriodWeek, true + case string(codersdk.ChatUsageLimitPeriodMonth): + return codersdk.ChatUsageLimitPeriodMonth, true + default: + return "", false + } +} diff --git a/coderd/chatd/usagelimit_test.go b/coderd/chatd/usagelimit_test.go new file mode 100644 index 0000000000..d618f8e44b --- /dev/null +++ b/coderd/chatd/usagelimit_test.go @@ -0,0 +1,132 @@ +package chatd //nolint:testpackage // Keeps chatd unit tests in the package. + +import ( + "testing" + "time" + + "github.com/coder/coder/v2/codersdk" +) + +func TestComputeUsagePeriodBounds(t *testing.T) { + t.Parallel() + + newYork, err := time.LoadLocation("America/New_York") + if err != nil { + t.Fatalf("load America/New_York: %v", err) + } + + tests := []struct { + name string + now time.Time + period codersdk.ChatUsageLimitPeriod + wantStart time.Time + wantEnd time.Time + }{ + { + name: "day/mid_day", + now: time.Date(2025, time.June, 15, 14, 30, 0, 0, time.UTC), + period: codersdk.ChatUsageLimitPeriodDay, + wantStart: time.Date(2025, time.June, 15, 0, 0, 0, 0, time.UTC), + wantEnd: time.Date(2025, time.June, 16, 0, 0, 0, 0, time.UTC), + }, + { + name: "day/midnight_exactly", + now: time.Date(2025, time.June, 15, 0, 0, 0, 0, time.UTC), + period: codersdk.ChatUsageLimitPeriodDay, + wantStart: time.Date(2025, time.June, 15, 0, 0, 0, 0, time.UTC), + wantEnd: time.Date(2025, time.June, 16, 0, 0, 0, 0, time.UTC), + }, + { + name: "day/end_of_day", + now: time.Date(2025, time.June, 15, 23, 59, 59, 0, time.UTC), + period: codersdk.ChatUsageLimitPeriodDay, + wantStart: time.Date(2025, time.June, 15, 0, 0, 0, 0, time.UTC), + wantEnd: time.Date(2025, time.June, 16, 0, 0, 0, 0, time.UTC), + }, + { + name: "week/wednesday", + now: time.Date(2025, time.June, 11, 10, 0, 0, 0, time.UTC), + period: codersdk.ChatUsageLimitPeriodWeek, + wantStart: time.Date(2025, time.June, 9, 0, 0, 0, 0, time.UTC), + wantEnd: time.Date(2025, time.June, 16, 0, 0, 0, 0, time.UTC), + }, + { + name: "week/monday", + now: time.Date(2025, time.June, 9, 0, 0, 0, 0, time.UTC), + period: codersdk.ChatUsageLimitPeriodWeek, + wantStart: time.Date(2025, time.June, 9, 0, 0, 0, 0, time.UTC), + wantEnd: time.Date(2025, time.June, 16, 0, 0, 0, 0, time.UTC), + }, + { + name: "week/sunday", + now: time.Date(2025, time.June, 15, 23, 0, 0, 0, time.UTC), + period: codersdk.ChatUsageLimitPeriodWeek, + wantStart: time.Date(2025, time.June, 9, 0, 0, 0, 0, time.UTC), + wantEnd: time.Date(2025, time.June, 16, 0, 0, 0, 0, time.UTC), + }, + { + name: "week/year_boundary", + now: time.Date(2024, time.December, 31, 12, 0, 0, 0, time.UTC), + period: codersdk.ChatUsageLimitPeriodWeek, + wantStart: time.Date(2024, time.December, 30, 0, 0, 0, 0, time.UTC), + wantEnd: time.Date(2025, time.January, 6, 0, 0, 0, 0, time.UTC), + }, + { + name: "month/mid_month", + now: time.Date(2025, time.June, 15, 0, 0, 0, 0, time.UTC), + period: codersdk.ChatUsageLimitPeriodMonth, + wantStart: time.Date(2025, time.June, 1, 0, 0, 0, 0, time.UTC), + wantEnd: time.Date(2025, time.July, 1, 0, 0, 0, 0, time.UTC), + }, + { + name: "month/first_day", + now: time.Date(2025, time.June, 1, 0, 0, 0, 0, time.UTC), + period: codersdk.ChatUsageLimitPeriodMonth, + wantStart: time.Date(2025, time.June, 1, 0, 0, 0, 0, time.UTC), + wantEnd: time.Date(2025, time.July, 1, 0, 0, 0, 0, time.UTC), + }, + { + name: "month/last_day", + now: time.Date(2025, time.June, 30, 23, 59, 59, 0, time.UTC), + period: codersdk.ChatUsageLimitPeriodMonth, + wantStart: time.Date(2025, time.June, 1, 0, 0, 0, 0, time.UTC), + wantEnd: time.Date(2025, time.July, 1, 0, 0, 0, 0, time.UTC), + }, + { + name: "month/february", + now: time.Date(2025, time.February, 15, 12, 0, 0, 0, time.UTC), + period: codersdk.ChatUsageLimitPeriodMonth, + wantStart: time.Date(2025, time.February, 1, 0, 0, 0, 0, time.UTC), + wantEnd: time.Date(2025, time.March, 1, 0, 0, 0, 0, time.UTC), + }, + { + name: "month/leap_year_february", + now: time.Date(2024, time.February, 29, 12, 0, 0, 0, time.UTC), + period: codersdk.ChatUsageLimitPeriodMonth, + wantStart: time.Date(2024, time.February, 1, 0, 0, 0, 0, time.UTC), + wantEnd: time.Date(2024, time.March, 1, 0, 0, 0, 0, time.UTC), + }, + { + name: "day/non_utc_timezone", + now: time.Date(2025, time.June, 15, 22, 0, 0, 0, newYork), + period: codersdk.ChatUsageLimitPeriodDay, + wantStart: time.Date(2025, time.June, 16, 0, 0, 0, 0, time.UTC), + wantEnd: time.Date(2025, time.June, 17, 0, 0, 0, 0, time.UTC), + }, + } + + for _, tc := range tests { + tc := tc + t.Run(tc.name, func(t *testing.T) { + t.Parallel() + + start, end := ComputeUsagePeriodBounds(tc.now, tc.period) + if !start.Equal(tc.wantStart) { + t.Errorf("start: got %v, want %v", start, tc.wantStart) + } + if !end.Equal(tc.wantEnd) { + t.Errorf("end: got %v, want %v", end, tc.wantEnd) + } + }) + } +} diff --git a/coderd/chats.go b/coderd/chats.go index 055edf9474..ab05c5c0b8 100644 --- a/coderd/chats.go +++ b/coderd/chats.go @@ -82,6 +82,30 @@ type chatDiffReference struct { RepositoryRef *chatRepositoryRef } +func writeChatUsageLimitExceeded( + ctx context.Context, + rw http.ResponseWriter, + limitErr *chatd.UsageLimitExceededError, +) { + httpapi.Write(ctx, rw, http.StatusConflict, codersdk.ChatUsageLimitExceededResponse{ + Response: codersdk.Response{ + Message: "Chat usage limit exceeded.", + }, + SpentMicros: limitErr.ConsumedMicros, + LimitMicros: limitErr.LimitMicros, + ResetsAt: limitErr.PeriodEnd, + }) +} + +func maybeWriteLimitErr(ctx context.Context, rw http.ResponseWriter, err error) bool { + var limitErr *chatd.UsageLimitExceededError + if errors.As(err, &limitErr) { + writeChatUsageLimitExceeded(ctx, rw, limitErr) + return true + } + return false +} + // EXPERIMENTAL: this endpoint is experimental and is subject to change. func (api *API) watchChats(rw http.ResponseWriter, r *http.Request) { ctx := r.Context() @@ -268,6 +292,9 @@ func (api *API) postChats(rw http.ResponseWriter, r *http.Request) { InitialUserContent: contentBlocks, }) if err != nil { + if maybeWriteLimitErr(ctx, rw, err) { + return + } if database.IsForeignKeyViolation( err, database.ForeignKeyChatsLastModelConfigID, @@ -431,7 +458,12 @@ func (api *API) chatCostSummary(rw http.ResponseWriter, r *http.Request) { chatBreakdowns = append(chatBreakdowns, convertChatCostChatBreakdown(chat)) } - httpapi.Write(ctx, rw, http.StatusOK, codersdk.ChatCostSummary{ + usageStatus, err := chatd.ResolveUsageLimitStatus(ctx, api.Database, targetUser.ID, time.Now()) + if err != nil { + api.Logger.Warn(ctx, "failed to resolve usage limit status", slog.Error(err)) + } + + response := codersdk.ChatCostSummary{ StartDate: startDate, EndDate: endDate, TotalCostMicros: summary.TotalCostMicros, @@ -443,7 +475,12 @@ func (api *API) chatCostSummary(rw http.ResponseWriter, r *http.Request) { TotalCacheCreationTokens: summary.TotalCacheCreationTokens, ByModel: modelBreakdowns, ByChat: chatBreakdowns, - }) + } + if usageStatus != nil { + response.UsageLimit = usageStatus + } + + httpapi.Write(ctx, rw, http.StatusOK, response) } func (api *API) chatCostUsers(rw http.ResponseWriter, r *http.Request) { @@ -547,6 +584,445 @@ func (api *API) chatCostUsers(rw http.ResponseWriter, r *http.Request) { }) } +// @Summary Get chat usage limit config +// @x-apidocgen {"skip": true} +// EXPERIMENTAL: this endpoint is experimental and is subject to change. +// +//nolint:revive // HTTP handler writes to ResponseWriter. +func (api *API) getChatUsageLimitConfig(rw http.ResponseWriter, r *http.Request) { + ctx := r.Context() + + if !api.Authorize(r, policy.ActionRead, rbac.ResourceDeploymentConfig) { + httpapi.Forbidden(rw) + return + } + + config, configErr := api.Database.GetChatUsageLimitConfig(ctx) + if configErr != nil && !errors.Is(configErr, sql.ErrNoRows) { + httpapi.Write(ctx, rw, http.StatusInternalServerError, codersdk.Response{ + Message: "Failed to get chat usage limit config.", + Detail: configErr.Error(), + }) + return + } + + overrideRows, err := api.Database.ListChatUsageLimitOverrides(ctx) + if err != nil { + httpapi.Write(ctx, rw, http.StatusInternalServerError, codersdk.Response{ + Message: "Failed to list chat usage limit overrides.", + Detail: err.Error(), + }) + return + } + + groupOverrides, err := api.Database.ListChatUsageLimitGroupOverrides(ctx) + if err != nil { + httpapi.Write(ctx, rw, http.StatusInternalServerError, codersdk.Response{ + Message: "Failed to list group usage limit overrides.", + Detail: err.Error(), + }) + return + } + + unpricedModelCount, err := api.Database.CountEnabledModelsWithoutPricing(ctx) + if err != nil { + httpapi.Write(ctx, rw, http.StatusInternalServerError, codersdk.Response{ + Message: "Failed to count unpriced chat models.", + Detail: err.Error(), + }) + return + } + + response := codersdk.ChatUsageLimitConfigResponse{ + ChatUsageLimitConfig: codersdk.ChatUsageLimitConfig{}, + UnpricedModelCount: unpricedModelCount, + Overrides: make([]codersdk.ChatUsageLimitOverride, 0, len(overrideRows)), + GroupOverrides: make([]codersdk.ChatUsageLimitGroupOverride, 0, len(groupOverrides)), + } + if configErr == nil { + response.Period = codersdk.ChatUsageLimitPeriod(config.Period) + response.UpdatedAt = config.UpdatedAt + if config.Enabled { + response.SpendLimitMicros = ptr.Ref(config.DefaultLimitMicros) + } + } + + for _, row := range overrideRows { + response.Overrides = append(response.Overrides, codersdk.ChatUsageLimitOverride{ + UserID: row.UserID, + Username: row.Username, + Name: row.Name, + AvatarURL: row.AvatarURL, + SpendLimitMicros: nullInt64Ptr(row.SpendLimitMicros), + }) + } + + for _, glo := range groupOverrides { + response.GroupOverrides = append(response.GroupOverrides, codersdk.ChatUsageLimitGroupOverride{ + GroupID: glo.GroupID, + GroupName: glo.GroupName, + GroupDisplayName: glo.GroupDisplayName, + GroupAvatarURL: glo.GroupAvatarUrl, + MemberCount: glo.MemberCount, + SpendLimitMicros: nullInt64Ptr(glo.SpendLimitMicros), + }) + } + httpapi.Write(ctx, rw, http.StatusOK, response) +} + +// @Summary Update chat usage limit config +// @x-apidocgen {"skip": true} +// EXPERIMENTAL: this endpoint is experimental and is subject to change. +func (api *API) updateChatUsageLimitConfig(rw http.ResponseWriter, r *http.Request) { + ctx := r.Context() + if !api.Authorize(r, policy.ActionUpdate, rbac.ResourceDeploymentConfig) { + httpapi.Forbidden(rw) + return + } + + var req codersdk.ChatUsageLimitConfig + if !httpapi.Read(ctx, rw, r, &req) { + return + } + + params := database.UpsertChatUsageLimitConfigParams{ + Enabled: false, + DefaultLimitMicros: 0, + Period: "", + } + if req.SpendLimitMicros == nil { + if req.Period != "" && !req.Period.Valid() { + httpapi.Write(ctx, rw, http.StatusBadRequest, codersdk.Response{ + Message: "Invalid chat usage limit period.", + Detail: "Period must be one of: day, week, month.", + }) + return + } + + params.Enabled = false + params.DefaultLimitMicros = 0 + params.Period = string(req.Period) + if params.Period == "" { + params.Period = string(codersdk.ChatUsageLimitPeriodMonth) + } + } else { + if *req.SpendLimitMicros <= 0 { + httpapi.Write(ctx, rw, http.StatusBadRequest, codersdk.Response{ + Message: "Invalid chat usage limit spend limit.", + Detail: "Spend limit must be greater than 0.", + }) + return + } + if !req.Period.Valid() { + httpapi.Write(ctx, rw, http.StatusBadRequest, codersdk.Response{ + Message: "Invalid chat usage limit period.", + Detail: "Period must be one of: day, week, month.", + }) + return + } + + params.Enabled = true + params.DefaultLimitMicros = *req.SpendLimitMicros + params.Period = string(req.Period) + } + + config, err := api.Database.UpsertChatUsageLimitConfig(ctx, params) + if err != nil { + httpapi.Write(ctx, rw, http.StatusInternalServerError, codersdk.Response{ + Message: "Failed to update chat usage limit config.", + Detail: err.Error(), + }) + return + } + + response := codersdk.ChatUsageLimitConfig{ + Period: codersdk.ChatUsageLimitPeriod(config.Period), + UpdatedAt: config.UpdatedAt, + } + if config.Enabled { + response.SpendLimitMicros = ptr.Ref(config.DefaultLimitMicros) + } + + httpapi.Write(ctx, rw, http.StatusOK, response) +} + +// @Summary Get my chat usage limit status +// @x-apidocgen {"skip": true} +// EXPERIMENTAL: this endpoint is experimental and is subject to change. +// +// getMyChatUsageLimitStatus returns the current usage-limit status for the +// authenticated user. No additional RBAC check is required because the +// endpoint always operates on the requesting user's own data via +// httpmw.APIKey(r).UserID. +// +//nolint:revive // HTTP handler writes to ResponseWriter. +func (api *API) getMyChatUsageLimitStatus(rw http.ResponseWriter, r *http.Request) { + ctx := r.Context() + status, err := chatd.ResolveUsageLimitStatus(ctx, api.Database, httpmw.APIKey(r).UserID, time.Now()) + if err != nil { + httpapi.Write(ctx, rw, http.StatusInternalServerError, codersdk.Response{ + Message: "Failed to get chat usage limit status.", + Detail: err.Error(), + }) + return + } + if status == nil { + httpapi.Write(ctx, rw, http.StatusOK, codersdk.ChatUsageLimitStatus{IsLimited: false}) + return + } + + httpapi.Write(ctx, rw, http.StatusOK, status) +} + +// @Summary Upsert chat usage limit override +// @x-apidocgen {"skip": true} +// EXPERIMENTAL: this endpoint is experimental and is subject to change. +func (api *API) upsertChatUsageLimitOverride(rw http.ResponseWriter, r *http.Request) { + ctx := r.Context() + if !api.Authorize(r, policy.ActionUpdate, rbac.ResourceDeploymentConfig) { + httpapi.Forbidden(rw) + return + } + + userID, ok := parseChatUsageLimitUserID(rw, r) + if !ok { + return + } + + var req codersdk.UpsertChatUsageLimitOverrideRequest + if !httpapi.Read(ctx, rw, r, &req) { + return + } + if req.SpendLimitMicros <= 0 { + httpapi.Write(ctx, rw, http.StatusBadRequest, codersdk.Response{ + Message: "Invalid chat usage limit override.", + Detail: "Spend limit must be greater than 0.", + }) + return + } + + user, err := api.Database.GetUserByID(ctx, userID) + if err != nil { + if errors.Is(err, sql.ErrNoRows) { + httpapi.Write(ctx, rw, http.StatusNotFound, codersdk.Response{ + Message: "User not found.", + }) + return + } + httpapi.Write(ctx, rw, http.StatusInternalServerError, codersdk.Response{ + Message: "Failed to look up chat usage limit user.", + Detail: err.Error(), + }) + return + } + + _, err = api.Database.UpsertChatUsageLimitUserOverride(ctx, database.UpsertChatUsageLimitUserOverrideParams{ + UserID: userID, + SpendLimitMicros: req.SpendLimitMicros, + }) + if err != nil { + httpapi.Write(ctx, rw, http.StatusInternalServerError, codersdk.Response{ + Message: "Failed to upsert chat usage limit override.", + Detail: err.Error(), + }) + return + } + + httpapi.Write(ctx, rw, http.StatusOK, codersdk.ChatUsageLimitOverride{ + UserID: user.ID, + Username: user.Username, + Name: user.Name, + AvatarURL: user.AvatarURL, + SpendLimitMicros: nullInt64Ptr(sql.NullInt64{Int64: req.SpendLimitMicros, Valid: true}), + }) +} + +// @Summary Delete chat usage limit override +// @x-apidocgen {"skip": true} +// EXPERIMENTAL: this endpoint is experimental and is subject to change. +func (api *API) deleteChatUsageLimitOverride(rw http.ResponseWriter, r *http.Request) { + ctx := r.Context() + if !api.Authorize(r, policy.ActionUpdate, rbac.ResourceDeploymentConfig) { + httpapi.Forbidden(rw) + return + } + + userID, ok := parseChatUsageLimitUserID(rw, r) + if !ok { + return + } + + if _, err := api.Database.GetUserByID(ctx, userID); err != nil { + if errors.Is(err, sql.ErrNoRows) { + writeChatUsageLimitUserNotFound(ctx, rw) + return + } + httpapi.Write(ctx, rw, http.StatusInternalServerError, codersdk.Response{ + Message: "Failed to look up chat usage limit user.", + Detail: err.Error(), + }) + return + } + if _, err := api.Database.GetChatUsageLimitUserOverride(ctx, userID); err != nil { + if errors.Is(err, sql.ErrNoRows) { + writeChatUsageLimitOverrideNotFound(ctx, rw) + return + } + httpapi.Write(ctx, rw, http.StatusInternalServerError, codersdk.Response{ + Message: "Failed to look up chat usage limit override.", + Detail: err.Error(), + }) + return + } + if err := api.Database.DeleteChatUsageLimitUserOverride(ctx, userID); err != nil { + httpapi.Write(ctx, rw, http.StatusInternalServerError, codersdk.Response{ + Message: "Failed to delete chat usage limit override.", + Detail: err.Error(), + }) + return + } + + rw.WriteHeader(http.StatusNoContent) +} + +// @Summary Upsert chat usage limit group override +// @x-apidocgen {"skip": true} +// EXPERIMENTAL: this endpoint is experimental and is subject to change. +func (api *API) upsertChatUsageLimitGroupOverride(rw http.ResponseWriter, r *http.Request) { + ctx := r.Context() + if !api.Authorize(r, policy.ActionUpdate, rbac.ResourceDeploymentConfig) { + httpapi.Forbidden(rw) + return + } + + groupIDStr := chi.URLParam(r, "group") + groupID, err := uuid.Parse(groupIDStr) + if err != nil { + httpapi.Write(ctx, rw, http.StatusBadRequest, codersdk.Response{ + Message: "Invalid group ID.", + Detail: err.Error(), + }) + return + } + + var req codersdk.UpdateChatUsageLimitGroupOverrideRequest + if !httpapi.Read(ctx, rw, r, &req) { + return + } + + if req.SpendLimitMicros <= 0 { + httpapi.Write(ctx, rw, http.StatusBadRequest, codersdk.Response{ + Message: "Invalid chat usage limit group override.", + Detail: "Spend limit (in microdollars) must be greater than 0.", + }) + return + } + + group, err := api.Database.GetGroupByID(ctx, groupID) + if err != nil { + if errors.Is(err, sql.ErrNoRows) { + httpapi.Write(ctx, rw, http.StatusNotFound, codersdk.Response{ + Message: "Group not found.", + }) + return + } + httpapi.Write(ctx, rw, http.StatusInternalServerError, codersdk.Response{ + Message: "Failed to look up group details.", + Detail: err.Error(), + }) + return + } + + _, err = api.Database.UpsertChatUsageLimitGroupOverride(ctx, database.UpsertChatUsageLimitGroupOverrideParams{ + GroupID: groupID, + SpendLimitMicros: req.SpendLimitMicros, + }) + if err != nil { + httpapi.Write(ctx, rw, http.StatusInternalServerError, codersdk.Response{ + Message: "Failed to upsert group usage limit override.", + Detail: err.Error(), + }) + return + } + + memberCount, err := api.Database.GetGroupMembersCountByGroupID(ctx, database.GetGroupMembersCountByGroupIDParams{ + GroupID: groupID, + IncludeSystem: false, + }) + if err != nil { + if errors.Is(err, sql.ErrNoRows) { + writeChatUsageLimitGroupNotFound(ctx, rw) + return + } + httpapi.Write(ctx, rw, http.StatusInternalServerError, codersdk.Response{ + Message: "Failed to fetch group member count.", + Detail: err.Error(), + }) + return + } + + httpapi.Write(ctx, rw, http.StatusOK, codersdk.ChatUsageLimitGroupOverride{ + GroupID: group.ID, + GroupName: group.Name, + GroupDisplayName: group.DisplayName, + GroupAvatarURL: group.AvatarURL, + MemberCount: memberCount, + SpendLimitMicros: nullInt64Ptr(sql.NullInt64{Int64: req.SpendLimitMicros, Valid: true}), + }) +} + +// @Summary Delete chat usage limit group override +// @x-apidocgen {"skip": true} +// EXPERIMENTAL: this endpoint is experimental and is subject to change. +func (api *API) deleteChatUsageLimitGroupOverride(rw http.ResponseWriter, r *http.Request) { + ctx := r.Context() + if !api.Authorize(r, policy.ActionUpdate, rbac.ResourceDeploymentConfig) { + httpapi.Forbidden(rw) + return + } + + groupIDStr := chi.URLParam(r, "group") + groupID, err := uuid.Parse(groupIDStr) + if err != nil { + httpapi.Write(ctx, rw, http.StatusBadRequest, codersdk.Response{ + Message: "Invalid group ID.", + Detail: err.Error(), + }) + return + } + + if _, err := api.Database.GetGroupByID(ctx, groupID); err != nil { + if errors.Is(err, sql.ErrNoRows) { + writeChatUsageLimitGroupNotFound(ctx, rw) + return + } + httpapi.Write(ctx, rw, http.StatusInternalServerError, codersdk.Response{ + Message: "Failed to look up group details.", + Detail: err.Error(), + }) + return + } + if _, err := api.Database.GetChatUsageLimitGroupOverride(ctx, groupID); err != nil { + if errors.Is(err, sql.ErrNoRows) { + writeChatUsageLimitGroupOverrideNotFound(ctx, rw) + return + } + httpapi.Write(ctx, rw, http.StatusInternalServerError, codersdk.Response{ + Message: "Failed to look up group usage limit override.", + Detail: err.Error(), + }) + return + } + if err := api.Database.DeleteChatUsageLimitGroupOverride(ctx, groupID); err != nil { + httpapi.Write(ctx, rw, http.StatusInternalServerError, codersdk.Response{ + Message: "Failed to delete group usage limit override.", + Detail: err.Error(), + }) + return + } + rw.WriteHeader(http.StatusNoContent) +} + // EXPERIMENTAL: this endpoint is experimental and is subject to change. // //nolint:revive // HTTP handler writes to ResponseWriter. @@ -996,6 +1472,9 @@ func (api *API) postChatMessages(rw http.ResponseWriter, r *http.Request) { }, ) if sendErr != nil { + if maybeWriteLimitErr(ctx, rw, sendErr) { + return + } if xerrors.Is(sendErr, chatd.ErrMessageQueueFull) { httpapi.Write(ctx, rw, http.StatusTooManyRequests, codersdk.Response{ Message: "Message queue is full.", @@ -1068,6 +1547,10 @@ func (api *API) patchChatMessage(rw http.ResponseWriter, r *http.Request) { Content: contentBlocks, }) if editErr != nil { + if maybeWriteLimitErr(ctx, rw, editErr) { + return + } + switch { case xerrors.Is(editErr, chatd.ErrEditedMessageNotFound): httpapi.Write(ctx, rw, http.StatusNotFound, codersdk.Response{ @@ -1158,6 +1641,9 @@ func (api *API) promoteChatQueuedMessage(rw http.ResponseWriter, r *http.Request }) if txErr != nil { + if maybeWriteLimitErr(ctx, rw, txErr) { + return + } httpapi.Write(ctx, rw, http.StatusInternalServerError, codersdk.Response{ Message: "Failed to promote queued message.", Detail: txErr.Error(), @@ -3355,6 +3841,49 @@ func chatModelConfigToUpdateParams( } } +func nullInt64Ptr(n sql.NullInt64) *int64 { + if !n.Valid { + return nil + } + return &n.Int64 +} + +func writeChatUsageLimitUserNotFound(ctx context.Context, rw http.ResponseWriter) { + httpapi.Write(ctx, rw, http.StatusBadRequest, codersdk.Response{ + Message: "User not found.", + }) +} + +func writeChatUsageLimitOverrideNotFound(ctx context.Context, rw http.ResponseWriter) { + httpapi.Write(ctx, rw, http.StatusBadRequest, codersdk.Response{ + Message: "Chat usage limit override not found.", + }) +} + +func writeChatUsageLimitGroupOverrideNotFound(ctx context.Context, rw http.ResponseWriter) { + httpapi.Write(ctx, rw, http.StatusBadRequest, codersdk.Response{ + Message: "Chat usage limit group override not found.", + }) +} + +func writeChatUsageLimitGroupNotFound(ctx context.Context, rw http.ResponseWriter) { + httpapi.Write(ctx, rw, http.StatusBadRequest, codersdk.Response{ + Message: "Group not found.", + }) +} + +func parseChatUsageLimitUserID(rw http.ResponseWriter, r *http.Request) (uuid.UUID, bool) { + userID, err := uuid.Parse(chi.URLParam(r, "user")) + if err != nil { + httpapi.Write(r.Context(), rw, http.StatusBadRequest, codersdk.Response{ + Message: "Invalid chat usage limit user ID.", + Detail: err.Error(), + }) + return uuid.Nil, false + } + return userID, true +} + func parseChatProviderID(rw http.ResponseWriter, r *http.Request) (uuid.UUID, bool) { providerID, err := uuid.Parse(chi.URLParam(r, "providerConfig")) if err != nil { diff --git a/coderd/chats_test.go b/coderd/chats_test.go index 1f6d6e6594..7ca7c93cd6 100644 --- a/coderd/chats_test.go +++ b/coderd/chats_test.go @@ -2,6 +2,7 @@ package coderd_test import ( "bytes" + "context" "database/sql" "encoding/json" "fmt" @@ -16,12 +17,15 @@ import ( "github.com/shopspring/decimal" "github.com/stretchr/testify/require" + "github.com/coder/coder/v2/coderd/chatd" + "github.com/coder/coder/v2/coderd/chatd/chatprompt" "github.com/coder/coder/v2/coderd/coderdtest" "github.com/coder/coder/v2/coderd/coderdtest/oidctest" "github.com/coder/coder/v2/coderd/database" "github.com/coder/coder/v2/coderd/database/db2sdk" "github.com/coder/coder/v2/coderd/database/dbauthz" "github.com/coder/coder/v2/coderd/database/dbfake" + "github.com/coder/coder/v2/coderd/database/dbgen" "github.com/coder/coder/v2/coderd/externalauth" coderdpubsub "github.com/coder/coder/v2/coderd/pubsub" "github.com/coder/coder/v2/coderd/rbac" @@ -55,6 +59,93 @@ func newChatClientWithDatabase(t testing.TB) (*codersdk.Client, database.Store) }) } +func requireChatUsageLimitExceededError( + t *testing.T, + err error, + wantSpentMicros int64, + wantLimitMicros int64, + wantResetsAt time.Time, +) *codersdk.ChatUsageLimitExceededResponse { + t.Helper() + + sdkErr, ok := codersdk.AsError(err) + require.True(t, ok) + require.Equal(t, http.StatusConflict, sdkErr.StatusCode()) + require.Equal(t, "Chat usage limit exceeded.", sdkErr.Message) + + limitErr := codersdk.ChatUsageLimitExceededFrom(err) + require.NotNil(t, limitErr) + require.Equal(t, "Chat usage limit exceeded.", limitErr.Message) + require.Equal(t, wantSpentMicros, limitErr.SpentMicros) + require.Equal(t, wantLimitMicros, limitErr.LimitMicros) + require.True( + t, + limitErr.ResetsAt.Equal(wantResetsAt), + "expected resets_at %s, got %s", + wantResetsAt.UTC().Format(time.RFC3339), + limitErr.ResetsAt.UTC().Format(time.RFC3339), + ) + + return limitErr +} + +func enableDailyChatUsageLimit( + ctx context.Context, + t *testing.T, + db database.Store, + limitMicros int64, +) time.Time { + t.Helper() + + _, err := db.UpsertChatUsageLimitConfig( + dbauthz.AsSystemRestricted(ctx), + database.UpsertChatUsageLimitConfigParams{ + Enabled: true, + DefaultLimitMicros: limitMicros, + Period: string(codersdk.ChatUsageLimitPeriodDay), + }, + ) + require.NoError(t, err) + + _, periodEnd := chatd.ComputeUsagePeriodBounds(time.Now(), codersdk.ChatUsageLimitPeriodDay) + return periodEnd +} + +func insertAssistantCostMessage( + ctx context.Context, + t *testing.T, + db database.Store, + chatID uuid.UUID, + modelConfigID uuid.UUID, + totalCostMicros int64, +) { + t.Helper() + + assistantContent, err := chatprompt.MarshalParts([]codersdk.ChatMessagePart{ + codersdk.ChatMessageText("assistant"), + }) + require.NoError(t, err) + + _, err = db.InsertChatMessage(dbauthz.AsSystemRestricted(ctx), database.InsertChatMessageParams{ + ChatID: chatID, + ModelConfigID: uuid.NullUUID{UUID: modelConfigID, Valid: true}, + Role: database.ChatMessageRoleAssistant, + ContentVersion: chatprompt.CurrentContentVersion, + Content: assistantContent, + Visibility: database.ChatMessageVisibilityBoth, + InputTokens: sql.NullInt64{}, + OutputTokens: sql.NullInt64{}, + TotalTokens: sql.NullInt64{}, + ReasoningTokens: sql.NullInt64{}, + CacheCreationTokens: sql.NullInt64{}, + CacheReadTokens: sql.NullInt64{}, + ContextLimit: sql.NullInt64{}, + Compressed: sql.NullBool{}, + TotalCostMicros: sql.NullInt64{Int64: totalCostMicros, Valid: true}, + }) + require.NoError(t, err) +} + func TestPostChats(t *testing.T) { t.Parallel() @@ -325,6 +416,33 @@ func TestPostChats(t *testing.T) { require.Equal(t, "Invalid input part.", sdkErr.Message) require.Equal(t, `content[0].type "image" is not supported.`, sdkErr.Detail) }) + + t.Run("UsageLimitExceeded", func(t *testing.T) { + t.Parallel() + + ctx := testutil.Context(t, testutil.WaitLong) + client, db := newChatClientWithDatabase(t) + user := coderdtest.CreateFirstUser(t, client) + modelConfig := createChatModelConfig(t, client) + wantResetsAt := enableDailyChatUsageLimit(ctx, t, db, 100) + + existingChat, err := db.InsertChat(dbauthz.AsSystemRestricted(ctx), database.InsertChatParams{ + OwnerID: user.UserID, + LastModelConfigID: modelConfig.ID, + Title: "existing-limit-chat", + }) + require.NoError(t, err) + + insertAssistantCostMessage(ctx, t, db, existingChat.ID, modelConfig.ID, 100) + + _, err = client.CreateChat(ctx, codersdk.CreateChatRequest{ + Content: []codersdk.ChatInputPart{{ + Type: codersdk.ChatInputPartTypeText, + Text: "over limit", + }}, + }) + requireChatUsageLimitExceededError(t, err, 100, 100, wantResetsAt) + }) } func TestListChats(t *testing.T) { @@ -1883,6 +2001,34 @@ func TestPostChatMessages(t *testing.T) { require.Equal(t, "content[0].text cannot be empty.", sdkErr.Detail) }) + t.Run("UsageLimitExceeded", func(t *testing.T) { + t.Parallel() + + ctx := testutil.Context(t, testutil.WaitLong) + client, db := newChatClientWithDatabase(t) + _ = coderdtest.CreateFirstUser(t, client) + modelConfig := createChatModelConfig(t, client) + + chat, err := client.CreateChat(ctx, codersdk.CreateChatRequest{ + Content: []codersdk.ChatInputPart{{ + Type: codersdk.ChatInputPartTypeText, + Text: "initial message for usage-limit test", + }}, + }) + require.NoError(t, err) + + wantResetsAt := enableDailyChatUsageLimit(ctx, t, db, 100) + insertAssistantCostMessage(ctx, t, db, chat.ID, modelConfig.ID, 100) + + _, err = client.CreateChatMessage(ctx, chat.ID, codersdk.CreateChatMessageRequest{ + Content: []codersdk.ChatInputPart{{ + Type: codersdk.ChatInputPartTypeText, + Text: "over limit", + }}, + }) + requireChatUsageLimitExceededError(t, err, 100, 100, wantResetsAt) + }) + t.Run("ChatNotFound", func(t *testing.T) { t.Parallel() @@ -2643,6 +2789,46 @@ func TestPatchChatMessage(t *testing.T) { require.True(t, foundFileInChat, "chat should preserve file_id after edit") }) + t.Run("UsageLimitExceeded", func(t *testing.T) { + t.Parallel() + + ctx := testutil.Context(t, testutil.WaitLong) + client, db := newChatClientWithDatabase(t) + _ = coderdtest.CreateFirstUser(t, client) + modelConfig := createChatModelConfig(t, client) + + chat, err := client.CreateChat(ctx, codersdk.CreateChatRequest{ + Content: []codersdk.ChatInputPart{{ + Type: codersdk.ChatInputPartTypeText, + Text: "hello before edit", + }}, + }) + require.NoError(t, err) + + messagesResult, err := client.GetChatMessages(ctx, chat.ID, nil) + require.NoError(t, err) + + var userMessageID int64 + for _, message := range messagesResult.Messages { + if message.Role == codersdk.ChatMessageRoleUser { + userMessageID = message.ID + break + } + } + require.NotZero(t, userMessageID) + + wantResetsAt := enableDailyChatUsageLimit(ctx, t, db, 100) + insertAssistantCostMessage(ctx, t, db, chat.ID, modelConfig.ID, 100) + + _, err = client.EditChatMessage(ctx, chat.ID, userMessageID, codersdk.EditChatMessageRequest{ + Content: []codersdk.ChatInputPart{{ + Type: codersdk.ChatInputPartTypeText, + Text: "edited over limit", + }}, + }) + requireChatUsageLimitExceededError(t, err, 100, 100, wantResetsAt) + }) + t.Run("MessageNotFound", func(t *testing.T) { t.Parallel() @@ -3352,6 +3538,81 @@ func TestPromoteChatQueuedMessage(t *testing.T) { } }) + t.Run("PromotesAlreadyQueuedMessageAfterLimitReached", func(t *testing.T) { + t.Parallel() + + ctx := testutil.Context(t, testutil.WaitLong) + client, db := newChatClientWithDatabase(t) + user := coderdtest.CreateFirstUser(t, client) + modelConfig := createChatModelConfig(t, client) + enableDailyChatUsageLimit(ctx, t, db, 100) + + chat, err := db.InsertChat(dbauthz.AsSystemRestricted(ctx), database.InsertChatParams{ + OwnerID: user.UserID, + LastModelConfigID: modelConfig.ID, + Title: "promote queued usage limit", + }) + require.NoError(t, err) + + const queuedText = "queued message for promote route" + queuedContent, err := json.Marshal([]codersdk.ChatMessagePart{ + codersdk.ChatMessageText(queuedText), + }) + require.NoError(t, err) + queuedMessage, err := db.InsertChatQueuedMessage( + dbauthz.AsSystemRestricted(ctx), + database.InsertChatQueuedMessageParams{ + ChatID: chat.ID, + Content: queuedContent, + }, + ) + require.NoError(t, err) + + insertAssistantCostMessage(ctx, t, db, chat.ID, modelConfig.ID, 100) + + _, err = db.UpdateChatStatus(dbauthz.AsSystemRestricted(ctx), database.UpdateChatStatusParams{ + ID: chat.ID, + Status: database.ChatStatusWaiting, + WorkerID: uuid.NullUUID{}, + StartedAt: sql.NullTime{}, + HeartbeatAt: sql.NullTime{}, + LastError: sql.NullString{}, + }) + require.NoError(t, err) + + promoteRes, err := client.Request( + ctx, + http.MethodPost, + fmt.Sprintf("/api/experimental/chats/%s/queue/%d/promote", chat.ID, queuedMessage.ID), + nil, + ) + require.NoError(t, err) + defer promoteRes.Body.Close() + require.Equal(t, http.StatusOK, promoteRes.StatusCode) + + var promoted codersdk.ChatMessage + err = json.NewDecoder(promoteRes.Body).Decode(&promoted) + require.NoError(t, err) + require.NotZero(t, promoted.ID) + require.Equal(t, chat.ID, promoted.ChatID) + require.Equal(t, codersdk.ChatMessageRoleUser, promoted.Role) + + foundPromotedText := false + for _, part := range promoted.Content { + if part.Type == codersdk.ChatMessagePartTypeText && part.Text == queuedText { + foundPromotedText = true + break + } + } + require.True(t, foundPromotedText) + + queuedMessages, err := db.GetChatQueuedMessages(dbauthz.AsSystemRestricted(ctx), chat.ID) + require.NoError(t, err) + for _, queued := range queuedMessages { + require.NotEqual(t, queuedMessage.ID, queued.ID) + } + }) + t.Run("InvalidQueuedMessageID", func(t *testing.T) { t.Parallel() @@ -3383,6 +3644,133 @@ func TestPromoteChatQueuedMessage(t *testing.T) { }) } +func TestChatUsageLimitOverrideRoutes(t *testing.T) { + t.Parallel() + + t.Run("UpsertUserOverrideRequiresPositiveSpendLimit", func(t *testing.T) { + t.Parallel() + + ctx := testutil.Context(t, testutil.WaitLong) + client, _ := newChatClientWithDatabase(t) + firstUser := coderdtest.CreateFirstUser(t, client) + _, member := coderdtest.CreateAnotherUser(t, client, firstUser.OrganizationID) + + res, err := client.Request( + ctx, + http.MethodPut, + fmt.Sprintf("/api/experimental/chats/usage-limits/overrides/%s", member.ID), + map[string]any{}, + ) + require.NoError(t, err) + defer res.Body.Close() + + err = codersdk.ReadBodyAsError(res) + sdkErr := requireSDKError(t, err, http.StatusBadRequest) + require.Equal(t, "Invalid chat usage limit override.", sdkErr.Message) + require.Equal(t, "Spend limit must be greater than 0.", sdkErr.Detail) + }) + + t.Run("UpsertUserOverrideMissingUser", func(t *testing.T) { + t.Parallel() + + ctx := testutil.Context(t, testutil.WaitLong) + client := newChatClient(t) + _ = coderdtest.CreateFirstUser(t, client) + + _, err := client.UpsertChatUsageLimitOverride(ctx, uuid.New(), codersdk.UpsertChatUsageLimitOverrideRequest{ + SpendLimitMicros: 7_000_000, + }) + sdkErr := requireSDKError(t, err, http.StatusNotFound) + require.Equal(t, "User not found.", sdkErr.Message) + }) + + t.Run("DeleteUserOverrideMissingUser", func(t *testing.T) { + t.Parallel() + + ctx := testutil.Context(t, testutil.WaitLong) + client := newChatClient(t) + _ = coderdtest.CreateFirstUser(t, client) + + err := client.DeleteChatUsageLimitOverride(ctx, uuid.New()) + sdkErr := requireSDKError(t, err, http.StatusBadRequest) + require.Equal(t, "User not found.", sdkErr.Message) + }) + + t.Run("DeleteUserOverrideMissingOverride", func(t *testing.T) { + t.Parallel() + + ctx := testutil.Context(t, testutil.WaitLong) + client := newChatClient(t) + firstUser := coderdtest.CreateFirstUser(t, client) + _, member := coderdtest.CreateAnotherUser(t, client, firstUser.OrganizationID) + + err := client.DeleteChatUsageLimitOverride(ctx, member.ID) + sdkErr := requireSDKError(t, err, http.StatusBadRequest) + require.Equal(t, "Chat usage limit override not found.", sdkErr.Message) + }) + + t.Run("UpsertGroupOverrideIncludesMemberCount", func(t *testing.T) { + t.Parallel() + + ctx := testutil.Context(t, testutil.WaitLong) + client, db := newChatClientWithDatabase(t) + firstUser := coderdtest.CreateFirstUser(t, client) + _, member := coderdtest.CreateAnotherUser(t, client, firstUser.OrganizationID) + group := dbgen.Group(t, db, database.Group{OrganizationID: firstUser.OrganizationID}) + dbgen.GroupMember(t, db, database.GroupMemberTable{GroupID: group.ID, UserID: member.ID}) + dbgen.GroupMember(t, db, database.GroupMemberTable{GroupID: group.ID, UserID: database.PrebuildsSystemUserID}) + + override, err := client.UpsertChatUsageLimitGroupOverride(ctx, group.ID, codersdk.UpsertChatUsageLimitGroupOverrideRequest{ + SpendLimitMicros: 7_000_000, + }) + require.NoError(t, err) + require.Equal(t, group.ID, override.GroupID) + require.EqualValues(t, 1, override.MemberCount) + require.NotNil(t, override.SpendLimitMicros) + require.EqualValues(t, 7_000_000, *override.SpendLimitMicros) + + config, err := client.GetChatUsageLimitConfig(ctx) + require.NoError(t, err) + + var listed *codersdk.ChatUsageLimitGroupOverride + for i := range config.GroupOverrides { + if config.GroupOverrides[i].GroupID == group.ID { + listed = &config.GroupOverrides[i] + break + } + } + require.NotNil(t, listed) + require.EqualValues(t, 1, listed.MemberCount) + }) + + t.Run("UpsertGroupOverrideMissingGroup", func(t *testing.T) { + t.Parallel() + + ctx := testutil.Context(t, testutil.WaitLong) + client := newChatClient(t) + _ = coderdtest.CreateFirstUser(t, client) + + _, err := client.UpsertChatUsageLimitGroupOverride(ctx, uuid.New(), codersdk.UpsertChatUsageLimitGroupOverrideRequest{ + SpendLimitMicros: 7_000_000, + }) + sdkErr := requireSDKError(t, err, http.StatusNotFound) + require.Equal(t, "Group not found.", sdkErr.Message) + }) + + t.Run("DeleteGroupOverrideMissingOverride", func(t *testing.T) { + t.Parallel() + + ctx := testutil.Context(t, testutil.WaitLong) + client, db := newChatClientWithDatabase(t) + firstUser := coderdtest.CreateFirstUser(t, client) + group := dbgen.Group(t, db, database.Group{OrganizationID: firstUser.OrganizationID}) + + err := client.DeleteChatUsageLimitGroupOverride(ctx, group.ID) + sdkErr := requireSDKError(t, err, http.StatusBadRequest) + require.Equal(t, "Chat usage limit group override not found.", sdkErr.Message) + }) +} + func TestPostChatFile(t *testing.T) { t.Parallel() diff --git a/coderd/coderd.go b/coderd/coderd.go index f6bdf09d2a..35f902207f 100644 --- a/coderd/coderd.go +++ b/coderd/coderd.go @@ -1178,6 +1178,19 @@ func New(options *Options) *API { r.Delete("/", api.deleteChatModelConfig) }) }) + r.Route("/usage-limits", func(r chi.Router) { + r.Get("/", api.getChatUsageLimitConfig) + r.Put("/", api.updateChatUsageLimitConfig) + r.Get("/status", api.getMyChatUsageLimitStatus) + r.Route("/overrides/{user}", func(r chi.Router) { + r.Put("/", api.upsertChatUsageLimitOverride) + r.Delete("/", api.deleteChatUsageLimitOverride) + }) + r.Route("/group-overrides/{group}", func(r chi.Router) { + r.Put("/", api.upsertChatUsageLimitGroupOverride) + r.Delete("/", api.deleteChatUsageLimitGroupOverride) + }) + }) r.Route("/{chat}", func(r chi.Router) { r.Use(httpmw.ExtractChatParam(options.Database)) r.Get("/", api.getChat) diff --git a/coderd/database/check_constraint.go b/coderd/database/check_constraint.go index 9b738411ef..af6d0fc248 100644 --- a/coderd/database/check_constraint.go +++ b/coderd/database/check_constraint.go @@ -6,22 +6,27 @@ type CheckConstraint string // CheckConstraint enums. const ( - CheckAPIKeysAllowListNotEmpty CheckConstraint = "api_keys_allow_list_not_empty" // api_keys - CheckChatModelConfigsCompressionThresholdCheck CheckConstraint = "chat_model_configs_compression_threshold_check" // chat_model_configs - CheckChatModelConfigsContextLimitCheck CheckConstraint = "chat_model_configs_context_limit_check" // chat_model_configs - CheckChatProvidersProviderCheck CheckConstraint = "chat_providers_provider_check" // chat_providers - CheckOrganizationIDNotZero CheckConstraint = "organization_id_not_zero" // custom_roles - CheckOneTimePasscodeSet CheckConstraint = "one_time_passcode_set" // users - CheckUsersEmailNotEmpty CheckConstraint = "users_email_not_empty" // users - CheckUsersServiceAccountLoginType CheckConstraint = "users_service_account_login_type" // users - CheckUsersUsernameMinLength CheckConstraint = "users_username_min_length" // users - 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 + CheckAPIKeysAllowListNotEmpty CheckConstraint = "api_keys_allow_list_not_empty" // api_keys + CheckChatModelConfigsCompressionThresholdCheck CheckConstraint = "chat_model_configs_compression_threshold_check" // chat_model_configs + CheckChatModelConfigsContextLimitCheck CheckConstraint = "chat_model_configs_context_limit_check" // chat_model_configs + CheckChatProvidersProviderCheck CheckConstraint = "chat_providers_provider_check" // chat_providers + 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 + CheckOrganizationIDNotZero CheckConstraint = "organization_id_not_zero" // custom_roles + CheckGroupsChatSpendLimitMicrosCheck CheckConstraint = "groups_chat_spend_limit_micros_check" // groups + 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 + 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 ) diff --git a/coderd/database/dbauthz/dbauthz.go b/coderd/database/dbauthz/dbauthz.go index 43f202d063..70ef6cff48 100644 --- a/coderd/database/dbauthz/dbauthz.go +++ b/coderd/database/dbauthz/dbauthz.go @@ -1726,6 +1726,13 @@ 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 @@ -1854,6 +1861,20 @@ func (q *querier) DeleteChatQueuedMessage(ctx context.Context, arg database.Dele return q.db.DeleteChatQueuedMessage(ctx, arg) } +func (q *querier) DeleteChatUsageLimitGroupOverride(ctx context.Context, groupID uuid.UUID) error { + if err := q.authorizeContext(ctx, policy.ActionUpdate, rbac.ResourceDeploymentConfig); err != nil { + return err + } + return q.db.DeleteChatUsageLimitGroupOverride(ctx, groupID) +} + +func (q *querier) DeleteChatUsageLimitUserOverride(ctx context.Context, userID uuid.UUID) error { + if err := q.authorizeContext(ctx, policy.ActionUpdate, rbac.ResourceDeploymentConfig); err != nil { + return err + } + return q.db.DeleteChatUsageLimitUserOverride(ctx, userID) +} + func (q *querier) DeleteCryptoKey(ctx context.Context, arg database.DeleteCryptoKeyParams) (database.CryptoKey, error) { if err := q.authorizeContext(ctx, policy.ActionDelete, rbac.ResourceCryptoKey); err != nil { return database.CryptoKey{}, err @@ -2611,6 +2632,27 @@ func (q *querier) GetChatSystemPrompt(ctx context.Context) (string, error) { return q.db.GetChatSystemPrompt(ctx) } +func (q *querier) GetChatUsageLimitConfig(ctx context.Context) (database.ChatUsageLimitConfig, error) { + if err := q.authorizeContext(ctx, policy.ActionRead, rbac.ResourceDeploymentConfig); err != nil { + return database.ChatUsageLimitConfig{}, err + } + return q.db.GetChatUsageLimitConfig(ctx) +} + +func (q *querier) GetChatUsageLimitGroupOverride(ctx context.Context, groupID uuid.UUID) (database.GetChatUsageLimitGroupOverrideRow, error) { + if err := q.authorizeContext(ctx, policy.ActionRead, rbac.ResourceDeploymentConfig); err != nil { + return database.GetChatUsageLimitGroupOverrideRow{}, err + } + return q.db.GetChatUsageLimitGroupOverride(ctx, groupID) +} + +func (q *querier) GetChatUsageLimitUserOverride(ctx context.Context, userID uuid.UUID) (database.GetChatUsageLimitUserOverrideRow, error) { + if err := q.authorizeContext(ctx, policy.ActionRead, rbac.ResourceDeploymentConfig); err != nil { + return database.GetChatUsageLimitUserOverrideRow{}, err + } + return q.db.GetChatUsageLimitUserOverride(ctx, userID) +} + func (q *querier) GetChatsByOwnerID(ctx context.Context, ownerID database.GetChatsByOwnerIDParams) ([]database.Chat, error) { return fetchWithPostFilter(q.auth, policy.ActionRead, q.db.GetChatsByOwnerID)(ctx, ownerID) } @@ -3765,6 +3807,13 @@ func (q *querier) GetUserChatCustomPrompt(ctx context.Context, userID uuid.UUID) return q.db.GetUserChatCustomPrompt(ctx, userID) } +func (q *querier) GetUserChatSpendInPeriod(ctx context.Context, arg database.GetUserChatSpendInPeriodParams) (int64, error) { + if err := q.authorizeContext(ctx, policy.ActionRead, rbac.ResourceChat.WithOwner(arg.UserID.String())); err != nil { + return 0, err + } + return q.db.GetUserChatSpendInPeriod(ctx, arg) +} + func (q *querier) GetUserCount(ctx context.Context, includeSystem bool) (int64, error) { if err := q.authorizeContext(ctx, policy.ActionRead, rbac.ResourceSystem); err != nil { return 0, err @@ -3772,6 +3821,13 @@ func (q *querier) GetUserCount(ctx context.Context, includeSystem bool) (int64, return q.db.GetUserCount(ctx, includeSystem) } +func (q *querier) GetUserGroupSpendLimit(ctx context.Context, userID uuid.UUID) (int64, error) { + if err := q.authorizeContext(ctx, policy.ActionRead, rbac.ResourceChat.WithOwner(userID.String())); err != nil { + return 0, err + } + return q.db.GetUserGroupSpendLimit(ctx, userID) +} + func (q *querier) GetUserLatencyInsights(ctx context.Context, arg database.GetUserLatencyInsightsParams) ([]database.GetUserLatencyInsightsRow, error) { // Used by insights endpoints. Need to check both for auditors and for regular users with template acl perms. if err := q.authorizeContext(ctx, policy.ActionViewInsights, rbac.ResourceTemplate); err != nil { @@ -5130,6 +5186,20 @@ func (q *querier) ListAIBridgeUserPromptsByInterceptionIDs(ctx context.Context, return q.db.ListAIBridgeUserPromptsByInterceptionIDs(ctx, interceptionIDs) } +func (q *querier) ListChatUsageLimitGroupOverrides(ctx context.Context) ([]database.ListChatUsageLimitGroupOverridesRow, error) { + if err := q.authorizeContext(ctx, policy.ActionRead, rbac.ResourceDeploymentConfig); err != nil { + return nil, err + } + return q.db.ListChatUsageLimitGroupOverrides(ctx) +} + +func (q *querier) ListChatUsageLimitOverrides(ctx context.Context) ([]database.ListChatUsageLimitOverridesRow, error) { + if err := q.authorizeContext(ctx, policy.ActionRead, rbac.ResourceDeploymentConfig); err != nil { + return nil, err + } + return q.db.ListChatUsageLimitOverrides(ctx) +} + func (q *querier) ListProvisionerKeysByOrganization(ctx context.Context, organizationID uuid.UUID) ([]database.ProvisionerKey, error) { return fetchWithPostFilter(q.auth, policy.ActionRead, q.db.ListProvisionerKeysByOrganization)(ctx, organizationID) } @@ -5249,6 +5319,13 @@ func (q *querier) RemoveUserFromGroups(ctx context.Context, arg database.RemoveU return q.db.RemoveUserFromGroups(ctx, arg) } +func (q *querier) ResolveUserChatSpendLimit(ctx context.Context, userID uuid.UUID) (int64, error) { + if err := q.authorizeContext(ctx, policy.ActionRead, rbac.ResourceChat.WithOwner(userID.String())); err != nil { + return 0, err + } + return q.db.ResolveUserChatSpendLimit(ctx, userID) +} + func (q *querier) RevokeDBCryptKey(ctx context.Context, activeKeyDigest string) error { if err := q.authorizeContext(ctx, policy.ActionUpdate, rbac.ResourceSystem); err != nil { return err @@ -6491,6 +6568,27 @@ func (q *querier) UpsertChatSystemPrompt(ctx context.Context, value string) erro return q.db.UpsertChatSystemPrompt(ctx, value) } +func (q *querier) UpsertChatUsageLimitConfig(ctx context.Context, arg database.UpsertChatUsageLimitConfigParams) (database.ChatUsageLimitConfig, error) { + if err := q.authorizeContext(ctx, policy.ActionUpdate, rbac.ResourceDeploymentConfig); err != nil { + return database.ChatUsageLimitConfig{}, err + } + return q.db.UpsertChatUsageLimitConfig(ctx, arg) +} + +func (q *querier) UpsertChatUsageLimitGroupOverride(ctx context.Context, arg database.UpsertChatUsageLimitGroupOverrideParams) (database.UpsertChatUsageLimitGroupOverrideRow, error) { + if err := q.authorizeContext(ctx, policy.ActionUpdate, rbac.ResourceDeploymentConfig); err != nil { + return database.UpsertChatUsageLimitGroupOverrideRow{}, err + } + return q.db.UpsertChatUsageLimitGroupOverride(ctx, arg) +} + +func (q *querier) UpsertChatUsageLimitUserOverride(ctx context.Context, arg database.UpsertChatUsageLimitUserOverrideParams) (database.UpsertChatUsageLimitUserOverrideRow, error) { + if err := q.authorizeContext(ctx, policy.ActionUpdate, rbac.ResourceDeploymentConfig); err != nil { + return database.UpsertChatUsageLimitUserOverrideRow{}, err + } + return q.db.UpsertChatUsageLimitUserOverride(ctx, arg) +} + func (q *querier) UpsertConnectionLog(ctx context.Context, arg database.UpsertConnectionLogParams) (database.ConnectionLog, error) { if err := q.authorizeContext(ctx, policy.ActionUpdate, rbac.ResourceConnectionLog); err != nil { return database.ConnectionLog{}, err diff --git a/coderd/database/dbauthz/dbauthz_test.go b/coderd/database/dbauthz/dbauthz_test.go index c5b5070ef7..0e4c068151 100644 --- a/coderd/database/dbauthz/dbauthz_test.go +++ b/coderd/database/dbauthz/dbauthz_test.go @@ -513,6 +513,10 @@ func (s *MethodTestSuite) TestChats() { dbm.EXPECT().GetChatCostSummary(gomock.Any(), arg).Return(row, nil).AnyTimes() check.Args(arg).Asserts(rbac.ResourceChat.WithOwner(arg.OwnerID.String()), 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}) @@ -841,6 +845,142 @@ func (s *MethodTestSuite) TestChats() { dbm.EXPECT().UpsertChatSystemPrompt(gomock.Any(), "").Return(nil).AnyTimes() check.Args("").Asserts(rbac.ResourceDeploymentConfig, policy.ActionUpdate) })) + s.Run("GetUserChatSpendInPeriod", s.Mocked(func(dbm *dbmock.MockStore, _ *gofakeit.Faker, check *expects) { + arg := database.GetUserChatSpendInPeriodParams{ + UserID: uuid.New(), + StartTime: time.Date(2025, 1, 1, 0, 0, 0, 0, time.UTC), + EndTime: time.Date(2025, 2, 1, 0, 0, 0, 0, time.UTC), + } + spend := int64(123) + dbm.EXPECT().GetUserChatSpendInPeriod(gomock.Any(), arg).Return(spend, nil).AnyTimes() + check.Args(arg).Asserts(rbac.ResourceChat.WithOwner(arg.UserID.String()), policy.ActionRead).Returns(spend) + })) + s.Run("GetUserGroupSpendLimit", s.Mocked(func(dbm *dbmock.MockStore, _ *gofakeit.Faker, check *expects) { + userID := uuid.New() + limit := int64(456) + dbm.EXPECT().GetUserGroupSpendLimit(gomock.Any(), userID).Return(limit, nil).AnyTimes() + check.Args(userID).Asserts(rbac.ResourceChat.WithOwner(userID.String()), policy.ActionRead).Returns(limit) + })) + s.Run("ResolveUserChatSpendLimit", s.Mocked(func(dbm *dbmock.MockStore, _ *gofakeit.Faker, check *expects) { + userID := uuid.New() + limit := int64(789) + dbm.EXPECT().ResolveUserChatSpendLimit(gomock.Any(), userID).Return(limit, nil).AnyTimes() + check.Args(userID).Asserts(rbac.ResourceChat.WithOwner(userID.String()), policy.ActionRead).Returns(limit) + })) + s.Run("GetChatUsageLimitConfig", s.Mocked(func(dbm *dbmock.MockStore, _ *gofakeit.Faker, check *expects) { + now := dbtime.Now() + config := database.ChatUsageLimitConfig{ + ID: 1, + Singleton: true, + Enabled: true, + DefaultLimitMicros: 1_000_000, + Period: "monthly", + CreatedAt: now, + UpdatedAt: now, + } + dbm.EXPECT().GetChatUsageLimitConfig(gomock.Any()).Return(config, nil).AnyTimes() + check.Args().Asserts(rbac.ResourceDeploymentConfig, policy.ActionRead).Returns(config) + })) + s.Run("GetChatUsageLimitGroupOverride", s.Mocked(func(dbm *dbmock.MockStore, _ *gofakeit.Faker, check *expects) { + groupID := uuid.New() + override := database.GetChatUsageLimitGroupOverrideRow{ + GroupID: groupID, + SpendLimitMicros: sql.NullInt64{Int64: 2_000_000, Valid: true}, + } + dbm.EXPECT().GetChatUsageLimitGroupOverride(gomock.Any(), groupID).Return(override, nil).AnyTimes() + check.Args(groupID).Asserts(rbac.ResourceDeploymentConfig, policy.ActionRead).Returns(override) + })) + s.Run("GetChatUsageLimitUserOverride", s.Mocked(func(dbm *dbmock.MockStore, _ *gofakeit.Faker, check *expects) { + userID := uuid.New() + override := database.GetChatUsageLimitUserOverrideRow{ + UserID: userID, + SpendLimitMicros: sql.NullInt64{Int64: 3_000_000, Valid: true}, + } + dbm.EXPECT().GetChatUsageLimitUserOverride(gomock.Any(), userID).Return(override, nil).AnyTimes() + check.Args(userID).Asserts(rbac.ResourceDeploymentConfig, policy.ActionRead).Returns(override) + })) + s.Run("ListChatUsageLimitGroupOverrides", s.Mocked(func(dbm *dbmock.MockStore, _ *gofakeit.Faker, check *expects) { + overrides := []database.ListChatUsageLimitGroupOverridesRow{{ + GroupID: uuid.New(), + GroupName: "group-name", + GroupDisplayName: "Group Name", + GroupAvatarUrl: "https://example.com/group.png", + SpendLimitMicros: sql.NullInt64{Int64: 4_000_000, Valid: true}, + MemberCount: 5, + }} + dbm.EXPECT().ListChatUsageLimitGroupOverrides(gomock.Any()).Return(overrides, nil).AnyTimes() + check.Args().Asserts(rbac.ResourceDeploymentConfig, policy.ActionRead).Returns(overrides) + })) + s.Run("ListChatUsageLimitOverrides", s.Mocked(func(dbm *dbmock.MockStore, _ *gofakeit.Faker, check *expects) { + overrides := []database.ListChatUsageLimitOverridesRow{{ + UserID: uuid.New(), + Username: "usage-limit-user", + Name: "Usage Limit User", + AvatarURL: "https://example.com/avatar.png", + SpendLimitMicros: sql.NullInt64{Int64: 5_000_000, Valid: true}, + }} + dbm.EXPECT().ListChatUsageLimitOverrides(gomock.Any()).Return(overrides, nil).AnyTimes() + check.Args().Asserts(rbac.ResourceDeploymentConfig, policy.ActionRead).Returns(overrides) + })) + s.Run("UpsertChatUsageLimitConfig", s.Mocked(func(dbm *dbmock.MockStore, _ *gofakeit.Faker, check *expects) { + now := dbtime.Now() + arg := database.UpsertChatUsageLimitConfigParams{ + Enabled: true, + DefaultLimitMicros: 6_000_000, + Period: "monthly", + } + config := database.ChatUsageLimitConfig{ + ID: 1, + Singleton: true, + Enabled: arg.Enabled, + DefaultLimitMicros: arg.DefaultLimitMicros, + Period: arg.Period, + CreatedAt: now, + UpdatedAt: now, + } + dbm.EXPECT().UpsertChatUsageLimitConfig(gomock.Any(), arg).Return(config, nil).AnyTimes() + check.Args(arg).Asserts(rbac.ResourceDeploymentConfig, policy.ActionUpdate).Returns(config) + })) + s.Run("UpsertChatUsageLimitGroupOverride", s.Mocked(func(dbm *dbmock.MockStore, _ *gofakeit.Faker, check *expects) { + arg := database.UpsertChatUsageLimitGroupOverrideParams{ + SpendLimitMicros: 7_000_000, + GroupID: uuid.New(), + } + override := database.UpsertChatUsageLimitGroupOverrideRow{ + GroupID: arg.GroupID, + Name: "group", + DisplayName: "Group", + AvatarURL: "", + SpendLimitMicros: sql.NullInt64{Int64: arg.SpendLimitMicros, Valid: true}, + } + dbm.EXPECT().UpsertChatUsageLimitGroupOverride(gomock.Any(), arg).Return(override, nil).AnyTimes() + check.Args(arg).Asserts(rbac.ResourceDeploymentConfig, policy.ActionUpdate).Returns(override) + })) + s.Run("UpsertChatUsageLimitUserOverride", s.Mocked(func(dbm *dbmock.MockStore, _ *gofakeit.Faker, check *expects) { + arg := database.UpsertChatUsageLimitUserOverrideParams{ + SpendLimitMicros: 8_000_000, + UserID: uuid.New(), + } + override := database.UpsertChatUsageLimitUserOverrideRow{ + UserID: arg.UserID, + Username: "user", + Name: "User", + AvatarURL: "", + SpendLimitMicros: sql.NullInt64{Int64: arg.SpendLimitMicros, Valid: true}, + } + dbm.EXPECT().UpsertChatUsageLimitUserOverride(gomock.Any(), arg).Return(override, nil).AnyTimes() + check.Args(arg).Asserts(rbac.ResourceDeploymentConfig, policy.ActionUpdate).Returns(override) + })) + s.Run("DeleteChatUsageLimitGroupOverride", s.Mocked(func(dbm *dbmock.MockStore, _ *gofakeit.Faker, check *expects) { + groupID := uuid.New() + dbm.EXPECT().DeleteChatUsageLimitGroupOverride(gomock.Any(), groupID).Return(nil).AnyTimes() + check.Args(groupID).Asserts(rbac.ResourceDeploymentConfig, policy.ActionUpdate) + })) + s.Run("DeleteChatUsageLimitUserOverride", s.Mocked(func(dbm *dbmock.MockStore, _ *gofakeit.Faker, check *expects) { + userID := uuid.New() + dbm.EXPECT().DeleteChatUsageLimitUserOverride(gomock.Any(), userID).Return(nil).AnyTimes() + check.Args(userID).Asserts(rbac.ResourceDeploymentConfig, policy.ActionUpdate) + })) } func (s *MethodTestSuite) TestFile() { diff --git a/coderd/database/dbmetrics/querymetrics.go b/coderd/database/dbmetrics/querymetrics.go index f1b275f23c..a971546694 100644 --- a/coderd/database/dbmetrics/querymetrics.go +++ b/coderd/database/dbmetrics/querymetrics.go @@ -288,6 +288,14 @@ 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) @@ -408,6 +416,22 @@ func (m queryMetricsStore) DeleteChatQueuedMessage(ctx context.Context, arg data return r0 } +func (m queryMetricsStore) DeleteChatUsageLimitGroupOverride(ctx context.Context, groupID uuid.UUID) error { + start := time.Now() + r0 := m.s.DeleteChatUsageLimitGroupOverride(ctx, groupID) + m.queryLatencies.WithLabelValues("DeleteChatUsageLimitGroupOverride").Observe(time.Since(start).Seconds()) + m.queryCounts.WithLabelValues(httpmw.ExtractHTTPRoute(ctx), httpmw.ExtractHTTPMethod(ctx), "DeleteChatUsageLimitGroupOverride").Inc() + return r0 +} + +func (m queryMetricsStore) DeleteChatUsageLimitUserOverride(ctx context.Context, userID uuid.UUID) error { + start := time.Now() + r0 := m.s.DeleteChatUsageLimitUserOverride(ctx, userID) + m.queryLatencies.WithLabelValues("DeleteChatUsageLimitUserOverride").Observe(time.Since(start).Seconds()) + m.queryCounts.WithLabelValues(httpmw.ExtractHTTPRoute(ctx), httpmw.ExtractHTTPMethod(ctx), "DeleteChatUsageLimitUserOverride").Inc() + return r0 +} + func (m queryMetricsStore) DeleteCryptoKey(ctx context.Context, arg database.DeleteCryptoKeyParams) (database.CryptoKey, error) { start := time.Now() r0, r1 := m.s.DeleteCryptoKey(ctx, arg) @@ -1143,6 +1167,30 @@ func (m queryMetricsStore) GetChatSystemPrompt(ctx context.Context) (string, err return r0, r1 } +func (m queryMetricsStore) GetChatUsageLimitConfig(ctx context.Context) (database.ChatUsageLimitConfig, error) { + start := time.Now() + r0, r1 := m.s.GetChatUsageLimitConfig(ctx) + m.queryLatencies.WithLabelValues("GetChatUsageLimitConfig").Observe(time.Since(start).Seconds()) + m.queryCounts.WithLabelValues(httpmw.ExtractHTTPRoute(ctx), httpmw.ExtractHTTPMethod(ctx), "GetChatUsageLimitConfig").Inc() + return r0, r1 +} + +func (m queryMetricsStore) GetChatUsageLimitGroupOverride(ctx context.Context, groupID uuid.UUID) (database.GetChatUsageLimitGroupOverrideRow, error) { + start := time.Now() + r0, r1 := m.s.GetChatUsageLimitGroupOverride(ctx, groupID) + m.queryLatencies.WithLabelValues("GetChatUsageLimitGroupOverride").Observe(time.Since(start).Seconds()) + m.queryCounts.WithLabelValues(httpmw.ExtractHTTPRoute(ctx), httpmw.ExtractHTTPMethod(ctx), "GetChatUsageLimitGroupOverride").Inc() + return r0, r1 +} + +func (m queryMetricsStore) GetChatUsageLimitUserOverride(ctx context.Context, userID uuid.UUID) (database.GetChatUsageLimitUserOverrideRow, error) { + start := time.Now() + r0, r1 := m.s.GetChatUsageLimitUserOverride(ctx, userID) + m.queryLatencies.WithLabelValues("GetChatUsageLimitUserOverride").Observe(time.Since(start).Seconds()) + m.queryCounts.WithLabelValues(httpmw.ExtractHTTPRoute(ctx), httpmw.ExtractHTTPMethod(ctx), "GetChatUsageLimitUserOverride").Inc() + return r0, r1 +} + func (m queryMetricsStore) GetChatsByOwnerID(ctx context.Context, ownerID database.GetChatsByOwnerIDParams) ([]database.Chat, error) { start := time.Now() r0, r1 := m.s.GetChatsByOwnerID(ctx, ownerID) @@ -2271,6 +2319,14 @@ func (m queryMetricsStore) GetUserChatCustomPrompt(ctx context.Context, userID u return r0, r1 } +func (m queryMetricsStore) GetUserChatSpendInPeriod(ctx context.Context, arg database.GetUserChatSpendInPeriodParams) (int64, error) { + start := time.Now() + r0, r1 := m.s.GetUserChatSpendInPeriod(ctx, arg) + m.queryLatencies.WithLabelValues("GetUserChatSpendInPeriod").Observe(time.Since(start).Seconds()) + m.queryCounts.WithLabelValues(httpmw.ExtractHTTPRoute(ctx), httpmw.ExtractHTTPMethod(ctx), "GetUserChatSpendInPeriod").Inc() + return r0, r1 +} + func (m queryMetricsStore) GetUserCount(ctx context.Context, includeSystem bool) (int64, error) { start := time.Now() r0, r1 := m.s.GetUserCount(ctx, includeSystem) @@ -2279,6 +2335,14 @@ func (m queryMetricsStore) GetUserCount(ctx context.Context, includeSystem bool) return r0, r1 } +func (m queryMetricsStore) GetUserGroupSpendLimit(ctx context.Context, userID uuid.UUID) (int64, error) { + start := time.Now() + r0, r1 := m.s.GetUserGroupSpendLimit(ctx, userID) + m.queryLatencies.WithLabelValues("GetUserGroupSpendLimit").Observe(time.Since(start).Seconds()) + m.queryCounts.WithLabelValues(httpmw.ExtractHTTPRoute(ctx), httpmw.ExtractHTTPMethod(ctx), "GetUserGroupSpendLimit").Inc() + return r0, r1 +} + func (m queryMetricsStore) GetUserLatencyInsights(ctx context.Context, arg database.GetUserLatencyInsightsParams) ([]database.GetUserLatencyInsightsRow, error) { start := time.Now() r0, r1 := m.s.GetUserLatencyInsights(ctx, arg) @@ -3511,6 +3575,22 @@ func (m queryMetricsStore) ListAIBridgeUserPromptsByInterceptionIDs(ctx context. return r0, r1 } +func (m queryMetricsStore) ListChatUsageLimitGroupOverrides(ctx context.Context) ([]database.ListChatUsageLimitGroupOverridesRow, error) { + start := time.Now() + r0, r1 := m.s.ListChatUsageLimitGroupOverrides(ctx) + m.queryLatencies.WithLabelValues("ListChatUsageLimitGroupOverrides").Observe(time.Since(start).Seconds()) + m.queryCounts.WithLabelValues(httpmw.ExtractHTTPRoute(ctx), httpmw.ExtractHTTPMethod(ctx), "ListChatUsageLimitGroupOverrides").Inc() + return r0, r1 +} + +func (m queryMetricsStore) ListChatUsageLimitOverrides(ctx context.Context) ([]database.ListChatUsageLimitOverridesRow, error) { + start := time.Now() + r0, r1 := m.s.ListChatUsageLimitOverrides(ctx) + m.queryLatencies.WithLabelValues("ListChatUsageLimitOverrides").Observe(time.Since(start).Seconds()) + m.queryCounts.WithLabelValues(httpmw.ExtractHTTPRoute(ctx), httpmw.ExtractHTTPMethod(ctx), "ListChatUsageLimitOverrides").Inc() + return r0, r1 +} + func (m queryMetricsStore) ListProvisionerKeysByOrganization(ctx context.Context, organizationID uuid.UUID) ([]database.ProvisionerKey, error) { start := time.Now() r0, r1 := m.s.ListProvisionerKeysByOrganization(ctx, organizationID) @@ -3623,6 +3703,14 @@ func (m queryMetricsStore) RemoveUserFromGroups(ctx context.Context, arg databas return r0, r1 } +func (m queryMetricsStore) ResolveUserChatSpendLimit(ctx context.Context, userID uuid.UUID) (int64, error) { + start := time.Now() + r0, r1 := m.s.ResolveUserChatSpendLimit(ctx, userID) + m.queryLatencies.WithLabelValues("ResolveUserChatSpendLimit").Observe(time.Since(start).Seconds()) + m.queryCounts.WithLabelValues(httpmw.ExtractHTTPRoute(ctx), httpmw.ExtractHTTPMethod(ctx), "ResolveUserChatSpendLimit").Inc() + return r0, r1 +} + func (m queryMetricsStore) RevokeDBCryptKey(ctx context.Context, activeKeyDigest string) error { start := time.Now() r0 := m.s.RevokeDBCryptKey(ctx, activeKeyDigest) @@ -4478,6 +4566,30 @@ func (m queryMetricsStore) UpsertChatSystemPrompt(ctx context.Context, value str return r0 } +func (m queryMetricsStore) UpsertChatUsageLimitConfig(ctx context.Context, arg database.UpsertChatUsageLimitConfigParams) (database.ChatUsageLimitConfig, error) { + start := time.Now() + r0, r1 := m.s.UpsertChatUsageLimitConfig(ctx, arg) + m.queryLatencies.WithLabelValues("UpsertChatUsageLimitConfig").Observe(time.Since(start).Seconds()) + m.queryCounts.WithLabelValues(httpmw.ExtractHTTPRoute(ctx), httpmw.ExtractHTTPMethod(ctx), "UpsertChatUsageLimitConfig").Inc() + return r0, r1 +} + +func (m queryMetricsStore) UpsertChatUsageLimitGroupOverride(ctx context.Context, arg database.UpsertChatUsageLimitGroupOverrideParams) (database.UpsertChatUsageLimitGroupOverrideRow, error) { + start := time.Now() + r0, r1 := m.s.UpsertChatUsageLimitGroupOverride(ctx, arg) + m.queryLatencies.WithLabelValues("UpsertChatUsageLimitGroupOverride").Observe(time.Since(start).Seconds()) + m.queryCounts.WithLabelValues(httpmw.ExtractHTTPRoute(ctx), httpmw.ExtractHTTPMethod(ctx), "UpsertChatUsageLimitGroupOverride").Inc() + return r0, r1 +} + +func (m queryMetricsStore) UpsertChatUsageLimitUserOverride(ctx context.Context, arg database.UpsertChatUsageLimitUserOverrideParams) (database.UpsertChatUsageLimitUserOverrideRow, error) { + start := time.Now() + r0, r1 := m.s.UpsertChatUsageLimitUserOverride(ctx, arg) + m.queryLatencies.WithLabelValues("UpsertChatUsageLimitUserOverride").Observe(time.Since(start).Seconds()) + m.queryCounts.WithLabelValues(httpmw.ExtractHTTPRoute(ctx), httpmw.ExtractHTTPMethod(ctx), "UpsertChatUsageLimitUserOverride").Inc() + return r0, r1 +} + func (m queryMetricsStore) UpsertConnectionLog(ctx context.Context, arg database.UpsertConnectionLogParams) (database.ConnectionLog, error) { start := time.Now() r0, r1 := m.s.UpsertConnectionLog(ctx, arg) diff --git a/coderd/database/dbmock/dbmock.go b/coderd/database/dbmock/dbmock.go index 1c33c22eeb..584d9737f8 100644 --- a/coderd/database/dbmock/dbmock.go +++ b/coderd/database/dbmock/dbmock.go @@ -424,6 +424,21 @@ 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() @@ -639,6 +654,34 @@ func (mr *MockStoreMockRecorder) DeleteChatQueuedMessage(ctx, arg any) *gomock.C return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "DeleteChatQueuedMessage", reflect.TypeOf((*MockStore)(nil).DeleteChatQueuedMessage), ctx, arg) } +// DeleteChatUsageLimitGroupOverride mocks base method. +func (m *MockStore) DeleteChatUsageLimitGroupOverride(ctx context.Context, groupID uuid.UUID) error { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "DeleteChatUsageLimitGroupOverride", ctx, groupID) + ret0, _ := ret[0].(error) + return ret0 +} + +// DeleteChatUsageLimitGroupOverride indicates an expected call of DeleteChatUsageLimitGroupOverride. +func (mr *MockStoreMockRecorder) DeleteChatUsageLimitGroupOverride(ctx, groupID any) *gomock.Call { + mr.mock.ctrl.T.Helper() + return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "DeleteChatUsageLimitGroupOverride", reflect.TypeOf((*MockStore)(nil).DeleteChatUsageLimitGroupOverride), ctx, groupID) +} + +// DeleteChatUsageLimitUserOverride mocks base method. +func (m *MockStore) DeleteChatUsageLimitUserOverride(ctx context.Context, userID uuid.UUID) error { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "DeleteChatUsageLimitUserOverride", ctx, userID) + ret0, _ := ret[0].(error) + return ret0 +} + +// DeleteChatUsageLimitUserOverride indicates an expected call of DeleteChatUsageLimitUserOverride. +func (mr *MockStoreMockRecorder) DeleteChatUsageLimitUserOverride(ctx, userID any) *gomock.Call { + mr.mock.ctrl.T.Helper() + return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "DeleteChatUsageLimitUserOverride", reflect.TypeOf((*MockStore)(nil).DeleteChatUsageLimitUserOverride), ctx, userID) +} + // DeleteCryptoKey mocks base method. func (m *MockStore) DeleteCryptoKey(ctx context.Context, arg database.DeleteCryptoKeyParams) (database.CryptoKey, error) { m.ctrl.T.Helper() @@ -2078,6 +2121,51 @@ func (mr *MockStoreMockRecorder) GetChatSystemPrompt(ctx any) *gomock.Call { return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetChatSystemPrompt", reflect.TypeOf((*MockStore)(nil).GetChatSystemPrompt), ctx) } +// GetChatUsageLimitConfig mocks base method. +func (m *MockStore) GetChatUsageLimitConfig(ctx context.Context) (database.ChatUsageLimitConfig, error) { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "GetChatUsageLimitConfig", ctx) + ret0, _ := ret[0].(database.ChatUsageLimitConfig) + ret1, _ := ret[1].(error) + return ret0, ret1 +} + +// GetChatUsageLimitConfig indicates an expected call of GetChatUsageLimitConfig. +func (mr *MockStoreMockRecorder) GetChatUsageLimitConfig(ctx any) *gomock.Call { + mr.mock.ctrl.T.Helper() + return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetChatUsageLimitConfig", reflect.TypeOf((*MockStore)(nil).GetChatUsageLimitConfig), ctx) +} + +// GetChatUsageLimitGroupOverride mocks base method. +func (m *MockStore) GetChatUsageLimitGroupOverride(ctx context.Context, groupID uuid.UUID) (database.GetChatUsageLimitGroupOverrideRow, error) { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "GetChatUsageLimitGroupOverride", ctx, groupID) + ret0, _ := ret[0].(database.GetChatUsageLimitGroupOverrideRow) + ret1, _ := ret[1].(error) + return ret0, ret1 +} + +// GetChatUsageLimitGroupOverride indicates an expected call of GetChatUsageLimitGroupOverride. +func (mr *MockStoreMockRecorder) GetChatUsageLimitGroupOverride(ctx, groupID any) *gomock.Call { + mr.mock.ctrl.T.Helper() + return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetChatUsageLimitGroupOverride", reflect.TypeOf((*MockStore)(nil).GetChatUsageLimitGroupOverride), ctx, groupID) +} + +// GetChatUsageLimitUserOverride mocks base method. +func (m *MockStore) GetChatUsageLimitUserOverride(ctx context.Context, userID uuid.UUID) (database.GetChatUsageLimitUserOverrideRow, error) { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "GetChatUsageLimitUserOverride", ctx, userID) + ret0, _ := ret[0].(database.GetChatUsageLimitUserOverrideRow) + ret1, _ := ret[1].(error) + return ret0, ret1 +} + +// GetChatUsageLimitUserOverride indicates an expected call of GetChatUsageLimitUserOverride. +func (mr *MockStoreMockRecorder) GetChatUsageLimitUserOverride(ctx, userID any) *gomock.Call { + mr.mock.ctrl.T.Helper() + return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetChatUsageLimitUserOverride", reflect.TypeOf((*MockStore)(nil).GetChatUsageLimitUserOverride), ctx, userID) +} + // GetChatsByOwnerID mocks base method. func (m *MockStore) GetChatsByOwnerID(ctx context.Context, arg database.GetChatsByOwnerIDParams) ([]database.Chat, error) { m.ctrl.T.Helper() @@ -4223,6 +4311,21 @@ func (mr *MockStoreMockRecorder) GetUserChatCustomPrompt(ctx, userID any) *gomoc return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetUserChatCustomPrompt", reflect.TypeOf((*MockStore)(nil).GetUserChatCustomPrompt), ctx, userID) } +// GetUserChatSpendInPeriod mocks base method. +func (m *MockStore) GetUserChatSpendInPeriod(ctx context.Context, arg database.GetUserChatSpendInPeriodParams) (int64, error) { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "GetUserChatSpendInPeriod", ctx, arg) + ret0, _ := ret[0].(int64) + ret1, _ := ret[1].(error) + return ret0, ret1 +} + +// GetUserChatSpendInPeriod indicates an expected call of GetUserChatSpendInPeriod. +func (mr *MockStoreMockRecorder) GetUserChatSpendInPeriod(ctx, arg any) *gomock.Call { + mr.mock.ctrl.T.Helper() + return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetUserChatSpendInPeriod", reflect.TypeOf((*MockStore)(nil).GetUserChatSpendInPeriod), ctx, arg) +} + // GetUserCount mocks base method. func (m *MockStore) GetUserCount(ctx context.Context, includeSystem bool) (int64, error) { m.ctrl.T.Helper() @@ -4238,6 +4341,21 @@ func (mr *MockStoreMockRecorder) GetUserCount(ctx, includeSystem any) *gomock.Ca return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetUserCount", reflect.TypeOf((*MockStore)(nil).GetUserCount), ctx, includeSystem) } +// GetUserGroupSpendLimit mocks base method. +func (m *MockStore) GetUserGroupSpendLimit(ctx context.Context, userID uuid.UUID) (int64, error) { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "GetUserGroupSpendLimit", ctx, userID) + ret0, _ := ret[0].(int64) + ret1, _ := ret[1].(error) + return ret0, ret1 +} + +// GetUserGroupSpendLimit indicates an expected call of GetUserGroupSpendLimit. +func (mr *MockStoreMockRecorder) GetUserGroupSpendLimit(ctx, userID any) *gomock.Call { + mr.mock.ctrl.T.Helper() + return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetUserGroupSpendLimit", reflect.TypeOf((*MockStore)(nil).GetUserGroupSpendLimit), ctx, userID) +} + // GetUserLatencyInsights mocks base method. func (m *MockStore) GetUserLatencyInsights(ctx context.Context, arg database.GetUserLatencyInsightsParams) ([]database.GetUserLatencyInsightsRow, error) { m.ctrl.T.Helper() @@ -6577,6 +6695,36 @@ func (mr *MockStoreMockRecorder) ListAuthorizedAIBridgeModels(ctx, arg, prepared return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "ListAuthorizedAIBridgeModels", reflect.TypeOf((*MockStore)(nil).ListAuthorizedAIBridgeModels), ctx, arg, prepared) } +// ListChatUsageLimitGroupOverrides mocks base method. +func (m *MockStore) ListChatUsageLimitGroupOverrides(ctx context.Context) ([]database.ListChatUsageLimitGroupOverridesRow, error) { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "ListChatUsageLimitGroupOverrides", ctx) + ret0, _ := ret[0].([]database.ListChatUsageLimitGroupOverridesRow) + ret1, _ := ret[1].(error) + return ret0, ret1 +} + +// ListChatUsageLimitGroupOverrides indicates an expected call of ListChatUsageLimitGroupOverrides. +func (mr *MockStoreMockRecorder) ListChatUsageLimitGroupOverrides(ctx any) *gomock.Call { + mr.mock.ctrl.T.Helper() + return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "ListChatUsageLimitGroupOverrides", reflect.TypeOf((*MockStore)(nil).ListChatUsageLimitGroupOverrides), ctx) +} + +// ListChatUsageLimitOverrides mocks base method. +func (m *MockStore) ListChatUsageLimitOverrides(ctx context.Context) ([]database.ListChatUsageLimitOverridesRow, error) { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "ListChatUsageLimitOverrides", ctx) + ret0, _ := ret[0].([]database.ListChatUsageLimitOverridesRow) + ret1, _ := ret[1].(error) + return ret0, ret1 +} + +// ListChatUsageLimitOverrides indicates an expected call of ListChatUsageLimitOverrides. +func (mr *MockStoreMockRecorder) ListChatUsageLimitOverrides(ctx any) *gomock.Call { + mr.mock.ctrl.T.Helper() + return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "ListChatUsageLimitOverrides", reflect.TypeOf((*MockStore)(nil).ListChatUsageLimitOverrides), ctx) +} + // ListProvisionerKeysByOrganization mocks base method. func (m *MockStore) ListProvisionerKeysByOrganization(ctx context.Context, organizationID uuid.UUID) ([]database.ProvisionerKey, error) { m.ctrl.T.Helper() @@ -6815,6 +6963,21 @@ func (mr *MockStoreMockRecorder) RemoveUserFromGroups(ctx, arg any) *gomock.Call return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "RemoveUserFromGroups", reflect.TypeOf((*MockStore)(nil).RemoveUserFromGroups), ctx, arg) } +// ResolveUserChatSpendLimit mocks base method. +func (m *MockStore) ResolveUserChatSpendLimit(ctx context.Context, userID uuid.UUID) (int64, error) { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "ResolveUserChatSpendLimit", ctx, userID) + ret0, _ := ret[0].(int64) + ret1, _ := ret[1].(error) + return ret0, ret1 +} + +// ResolveUserChatSpendLimit indicates an expected call of ResolveUserChatSpendLimit. +func (mr *MockStoreMockRecorder) ResolveUserChatSpendLimit(ctx, userID any) *gomock.Call { + mr.mock.ctrl.T.Helper() + return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "ResolveUserChatSpendLimit", reflect.TypeOf((*MockStore)(nil).ResolveUserChatSpendLimit), ctx, userID) +} + // RevokeDBCryptKey mocks base method. func (m *MockStore) RevokeDBCryptKey(ctx context.Context, activeKeyDigest string) error { m.ctrl.T.Helper() @@ -8361,6 +8524,51 @@ func (mr *MockStoreMockRecorder) UpsertChatSystemPrompt(ctx, value any) *gomock. return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "UpsertChatSystemPrompt", reflect.TypeOf((*MockStore)(nil).UpsertChatSystemPrompt), ctx, value) } +// UpsertChatUsageLimitConfig mocks base method. +func (m *MockStore) UpsertChatUsageLimitConfig(ctx context.Context, arg database.UpsertChatUsageLimitConfigParams) (database.ChatUsageLimitConfig, error) { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "UpsertChatUsageLimitConfig", ctx, arg) + ret0, _ := ret[0].(database.ChatUsageLimitConfig) + ret1, _ := ret[1].(error) + return ret0, ret1 +} + +// UpsertChatUsageLimitConfig indicates an expected call of UpsertChatUsageLimitConfig. +func (mr *MockStoreMockRecorder) UpsertChatUsageLimitConfig(ctx, arg any) *gomock.Call { + mr.mock.ctrl.T.Helper() + return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "UpsertChatUsageLimitConfig", reflect.TypeOf((*MockStore)(nil).UpsertChatUsageLimitConfig), ctx, arg) +} + +// UpsertChatUsageLimitGroupOverride mocks base method. +func (m *MockStore) UpsertChatUsageLimitGroupOverride(ctx context.Context, arg database.UpsertChatUsageLimitGroupOverrideParams) (database.UpsertChatUsageLimitGroupOverrideRow, error) { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "UpsertChatUsageLimitGroupOverride", ctx, arg) + ret0, _ := ret[0].(database.UpsertChatUsageLimitGroupOverrideRow) + ret1, _ := ret[1].(error) + return ret0, ret1 +} + +// UpsertChatUsageLimitGroupOverride indicates an expected call of UpsertChatUsageLimitGroupOverride. +func (mr *MockStoreMockRecorder) UpsertChatUsageLimitGroupOverride(ctx, arg any) *gomock.Call { + mr.mock.ctrl.T.Helper() + return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "UpsertChatUsageLimitGroupOverride", reflect.TypeOf((*MockStore)(nil).UpsertChatUsageLimitGroupOverride), ctx, arg) +} + +// UpsertChatUsageLimitUserOverride mocks base method. +func (m *MockStore) UpsertChatUsageLimitUserOverride(ctx context.Context, arg database.UpsertChatUsageLimitUserOverrideParams) (database.UpsertChatUsageLimitUserOverrideRow, error) { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "UpsertChatUsageLimitUserOverride", ctx, arg) + ret0, _ := ret[0].(database.UpsertChatUsageLimitUserOverrideRow) + ret1, _ := ret[1].(error) + return ret0, ret1 +} + +// UpsertChatUsageLimitUserOverride indicates an expected call of UpsertChatUsageLimitUserOverride. +func (mr *MockStoreMockRecorder) UpsertChatUsageLimitUserOverride(ctx, arg any) *gomock.Call { + mr.mock.ctrl.T.Helper() + return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "UpsertChatUsageLimitUserOverride", reflect.TypeOf((*MockStore)(nil).UpsertChatUsageLimitUserOverride), ctx, arg) +} + // UpsertConnectionLog mocks base method. func (m *MockStore) UpsertConnectionLog(ctx context.Context, arg database.UpsertConnectionLogParams) (database.ConnectionLog, error) { m.ctrl.T.Helper() diff --git a/coderd/database/dump.sql b/coderd/database/dump.sql index e81cf62f72..fc2a274cb0 100644 --- a/coderd/database/dump.sql +++ b/coderd/database/dump.sql @@ -1318,6 +1318,28 @@ CREATE SEQUENCE chat_queued_messages_id_seq ALTER SEQUENCE chat_queued_messages_id_seq OWNED BY chat_queued_messages.id; +CREATE TABLE chat_usage_limit_config ( + id bigint NOT NULL, + singleton boolean DEFAULT true NOT NULL, + enabled boolean DEFAULT false NOT NULL, + default_limit_micros bigint DEFAULT 0 NOT NULL, + period text DEFAULT 'month'::text NOT NULL, + created_at timestamp with time zone DEFAULT now() NOT NULL, + updated_at timestamp with time zone DEFAULT now() NOT NULL, + CONSTRAINT chat_usage_limit_config_default_limit_micros_check CHECK ((default_limit_micros >= 0)), + CONSTRAINT chat_usage_limit_config_period_check CHECK ((period = ANY (ARRAY['day'::text, 'week'::text, 'month'::text]))), + CONSTRAINT chat_usage_limit_config_singleton_check CHECK (singleton) +); + +CREATE SEQUENCE chat_usage_limit_config_id_seq + START WITH 1 + INCREMENT BY 1 + NO MINVALUE + NO MAXVALUE + CACHE 1; + +ALTER SEQUENCE chat_usage_limit_config_id_seq OWNED BY chat_usage_limit_config.id; + CREATE TABLE chats ( id uuid DEFAULT gen_random_uuid() NOT NULL, owner_id uuid NOT NULL, @@ -1474,7 +1496,9 @@ CREATE TABLE groups ( avatar_url text DEFAULT ''::text NOT NULL, quota_allowance integer DEFAULT 0 NOT NULL, display_name text DEFAULT ''::text NOT NULL, - source group_source DEFAULT 'user'::group_source NOT NULL + source group_source DEFAULT 'user'::group_source NOT NULL, + chat_spend_limit_micros bigint, + CONSTRAINT groups_chat_spend_limit_micros_check CHECK (((chat_spend_limit_micros IS NULL) OR (chat_spend_limit_micros > 0))) ); COMMENT ON COLUMN groups.display_name IS 'Display name is a custom, human-friendly group name that user can set. This is not required to be unique and can be the empty string.'; @@ -1509,7 +1533,9 @@ CREATE TABLE users ( one_time_passcode_expires_at timestamp with time zone, is_system boolean DEFAULT false NOT NULL, is_service_account boolean DEFAULT false NOT NULL, + chat_spend_limit_micros bigint, CONSTRAINT one_time_passcode_set CHECK ((((hashed_one_time_passcode IS NULL) AND (one_time_passcode_expires_at IS NULL)) OR ((hashed_one_time_passcode IS NOT NULL) AND (one_time_passcode_expires_at IS NOT NULL)))), + CONSTRAINT users_chat_spend_limit_micros_check CHECK (((chat_spend_limit_micros IS NULL) OR (chat_spend_limit_micros > 0))), CONSTRAINT users_email_not_empty CHECK (((is_service_account = true) = (email = ''::text))), CONSTRAINT users_service_account_login_type CHECK (((is_service_account = false) OR (login_type = 'none'::login_type))), CONSTRAINT users_username_min_length CHECK ((length(username) >= 1)) @@ -3156,6 +3182,8 @@ ALTER TABLE ONLY chat_messages ALTER COLUMN id SET DEFAULT nextval('chat_message ALTER TABLE ONLY chat_queued_messages ALTER COLUMN id SET DEFAULT nextval('chat_queued_messages_id_seq'::regclass); +ALTER TABLE ONLY chat_usage_limit_config ALTER COLUMN id SET DEFAULT nextval('chat_usage_limit_config_id_seq'::regclass); + ALTER TABLE ONLY licenses ALTER COLUMN id SET DEFAULT nextval('licenses_id_seq'::regclass); ALTER TABLE ONLY provisioner_job_logs ALTER COLUMN id SET DEFAULT nextval('provisioner_job_logs_id_seq'::regclass); @@ -3216,6 +3244,12 @@ ALTER TABLE ONLY chat_providers ALTER TABLE ONLY chat_queued_messages ADD CONSTRAINT chat_queued_messages_pkey PRIMARY KEY (id); +ALTER TABLE ONLY chat_usage_limit_config + ADD CONSTRAINT chat_usage_limit_config_pkey PRIMARY KEY (id); + +ALTER TABLE ONLY chat_usage_limit_config + ADD CONSTRAINT chat_usage_limit_config_singleton_key UNIQUE (singleton); + ALTER TABLE ONLY chats ADD CONSTRAINT chats_pkey PRIMARY KEY (id); @@ -3568,6 +3602,8 @@ CREATE INDEX idx_chat_messages_compressed_summary_boundary ON chat_messages USIN CREATE INDEX idx_chat_messages_created_at ON chat_messages USING btree (created_at); +CREATE INDEX idx_chat_messages_owner_spend ON chat_messages USING btree (chat_id, created_at) WHERE (total_cost_micros IS NOT NULL); + CREATE INDEX idx_chat_model_configs_enabled ON chat_model_configs USING btree (enabled); CREATE INDEX idx_chat_model_configs_provider ON chat_model_configs USING btree (provider); diff --git a/coderd/database/migrations/000441_chat_usage_limits.down.sql b/coderd/database/migrations/000441_chat_usage_limits.down.sql new file mode 100644 index 0000000000..56ce07e91e --- /dev/null +++ b/coderd/database/migrations/000441_chat_usage_limits.down.sql @@ -0,0 +1,4 @@ +DROP INDEX IF EXISTS idx_chat_messages_owner_spend; +ALTER TABLE groups DROP COLUMN IF EXISTS chat_spend_limit_micros; +ALTER TABLE users DROP COLUMN IF EXISTS chat_spend_limit_micros; +DROP TABLE IF EXISTS chat_usage_limit_config; diff --git a/coderd/database/migrations/000441_chat_usage_limits.up.sql b/coderd/database/migrations/000441_chat_usage_limits.up.sql new file mode 100644 index 0000000000..2dbfdb7a55 --- /dev/null +++ b/coderd/database/migrations/000441_chat_usage_limits.up.sql @@ -0,0 +1,32 @@ +-- 1. Singleton config table +CREATE TABLE chat_usage_limit_config ( + id BIGSERIAL PRIMARY KEY, + -- Only one row allowed (enforced by CHECK). + singleton BOOLEAN NOT NULL DEFAULT TRUE CHECK (singleton), + UNIQUE (singleton), + enabled BOOLEAN NOT NULL DEFAULT FALSE, + -- Limit per user per period, in micro-dollars (1 USD = 1,000,000). + default_limit_micros BIGINT NOT NULL DEFAULT 0 + CHECK (default_limit_micros >= 0), + -- Period length: 'day', 'week', or 'month'. + period TEXT NOT NULL DEFAULT 'month' + CHECK (period IN ('day', 'week', 'month')), + created_at TIMESTAMPTZ NOT NULL DEFAULT NOW(), + updated_at TIMESTAMPTZ NOT NULL DEFAULT NOW() +); + +-- Seed a single disabled row so reads never return empty. +INSERT INTO chat_usage_limit_config (singleton) VALUES (TRUE); + +-- 2. Per-user overrides (inline on users table). +ALTER TABLE users ADD COLUMN chat_spend_limit_micros BIGINT DEFAULT NULL + CHECK (chat_spend_limit_micros IS NULL OR chat_spend_limit_micros > 0); + +-- 3. Per-group overrides (inline on groups table). +ALTER TABLE groups ADD COLUMN chat_spend_limit_micros BIGINT DEFAULT NULL + CHECK (chat_spend_limit_micros IS NULL OR chat_spend_limit_micros > 0); + +-- Speed up per-user spend aggregation in the usage-limit hot path. +CREATE INDEX idx_chat_messages_owner_spend + ON chat_messages (chat_id, created_at) + WHERE total_cost_micros IS NOT NULL; diff --git a/coderd/database/migrations/testdata/fixtures/000441_chat_usage_limits.up.sql b/coderd/database/migrations/testdata/fixtures/000441_chat_usage_limits.up.sql new file mode 100644 index 0000000000..a01dbc8862 --- /dev/null +++ b/coderd/database/migrations/testdata/fixtures/000441_chat_usage_limits.up.sql @@ -0,0 +1,5 @@ +UPDATE users SET chat_spend_limit_micros = 5000000 +WHERE id = 'fc1511ef-4fcf-4a3b-98a1-8df64160e35a'; + +UPDATE groups SET chat_spend_limit_micros = 10000000 +WHERE id = 'bb640d07-ca8a-4869-b6bc-ae61ebb2fda1'; diff --git a/coderd/database/modelqueries.go b/coderd/database/modelqueries.go index dbad0a1329..5f0de02c56 100644 --- a/coderd/database/modelqueries.go +++ b/coderd/database/modelqueries.go @@ -451,6 +451,7 @@ func (q *sqlQuerier) GetAuthorizedUsers(ctx context.Context, arg GetUsersParams, &i.OneTimePasscodeExpiresAt, &i.IsSystem, &i.IsServiceAccount, + &i.ChatSpendLimitMicros, &i.Count, ); err != nil { return nil, err diff --git a/coderd/database/models.go b/coderd/database/models.go index 2ac42f9277..9beaaf9f85 100644 --- a/coderd/database/models.go +++ b/coderd/database/models.go @@ -4196,6 +4196,16 @@ type ChatQueuedMessage struct { CreatedAt time.Time `db:"created_at" json:"created_at"` } +type ChatUsageLimitConfig struct { + ID int64 `db:"id" json:"id"` + Singleton bool `db:"singleton" json:"singleton"` + Enabled bool `db:"enabled" json:"enabled"` + DefaultLimitMicros int64 `db:"default_limit_micros" json:"default_limit_micros"` + Period string `db:"period" json:"period"` + CreatedAt time.Time `db:"created_at" json:"created_at"` + UpdatedAt time.Time `db:"updated_at" json:"updated_at"` +} + type ConnectionLog struct { ID uuid.UUID `db:"id" json:"id"` ConnectTime time.Time `db:"connect_time" json:"connect_time"` @@ -4308,7 +4318,8 @@ type Group struct { // Display name is a custom, human-friendly group name that user can set. This is not required to be unique and can be the empty string. DisplayName string `db:"display_name" json:"display_name"` // Source indicates how the group was created. It can be created by a user manually, or through some system process like OIDC group sync. - Source GroupSource `db:"source" json:"source"` + Source GroupSource `db:"source" json:"source"` + ChatSpendLimitMicros sql.NullInt64 `db:"chat_spend_limit_micros" json:"chat_spend_limit_micros"` } // Joins group members with user information, organization ID, group name. Includes both regular group members and organization members (as part of the "Everyone" group). @@ -5078,7 +5089,8 @@ type User struct { // Determines if a user is a system user, and therefore cannot login or perform normal actions IsSystem bool `db:"is_system" json:"is_system"` // Determines if a user is an admin-managed account that cannot login - IsServiceAccount bool `db:"is_service_account" json:"is_service_account"` + IsServiceAccount bool `db:"is_service_account" json:"is_service_account"` + ChatSpendLimitMicros sql.NullInt64 `db:"chat_spend_limit_micros" json:"chat_spend_limit_micros"` } type UserConfig struct { diff --git a/coderd/database/querier.go b/coderd/database/querier.go index dfb368c03d..5759cc5cc3 100644 --- a/coderd/database/querier.go +++ b/coderd/database/querier.go @@ -77,6 +77,9 @@ type sqlcQuerier interface { CountAIBridgeInterceptions(ctx context.Context, arg CountAIBridgeInterceptionsParams) (int64, error) CountAuditLogs(ctx context.Context, arg CountAuditLogsParams) (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) @@ -99,6 +102,8 @@ type sqlcQuerier interface { DeleteChatModelConfigByID(ctx context.Context, id uuid.UUID) error DeleteChatProviderByID(ctx context.Context, id uuid.UUID) error DeleteChatQueuedMessage(ctx context.Context, arg DeleteChatQueuedMessageParams) error + DeleteChatUsageLimitGroupOverride(ctx context.Context, groupID uuid.UUID) error + DeleteChatUsageLimitUserOverride(ctx context.Context, userID uuid.UUID) error DeleteCryptoKey(ctx context.Context, arg DeleteCryptoKeyParams) (CryptoKey, error) DeleteCustomRole(ctx context.Context, arg DeleteCustomRoleParams) error DeleteExpiredAPIKeys(ctx context.Context, arg DeleteExpiredAPIKeysParams) (int64, error) @@ -244,6 +249,9 @@ type sqlcQuerier interface { GetChatProviders(ctx context.Context) ([]ChatProvider, error) GetChatQueuedMessages(ctx context.Context, chatID uuid.UUID) ([]ChatQueuedMessage, error) GetChatSystemPrompt(ctx context.Context) (string, error) + GetChatUsageLimitConfig(ctx context.Context) (ChatUsageLimitConfig, error) + GetChatUsageLimitGroupOverride(ctx context.Context, groupID uuid.UUID) (GetChatUsageLimitGroupOverrideRow, error) + GetChatUsageLimitUserOverride(ctx context.Context, userID uuid.UUID) (GetChatUsageLimitUserOverrideRow, error) GetChatsByOwnerID(ctx context.Context, arg GetChatsByOwnerIDParams) ([]Chat, error) GetConnectionLogsOffset(ctx context.Context, arg GetConnectionLogsOffsetParams) ([]GetConnectionLogsOffsetRow, error) GetCryptoKeyByFeatureAndSequence(ctx context.Context, arg GetCryptoKeyByFeatureAndSequenceParams) (CryptoKey, error) @@ -496,7 +504,11 @@ type sqlcQuerier interface { GetUserByEmailOrUsername(ctx context.Context, arg GetUserByEmailOrUsernameParams) (User, error) GetUserByID(ctx context.Context, id uuid.UUID) (User, error) GetUserChatCustomPrompt(ctx context.Context, userID uuid.UUID) (string, error) + GetUserChatSpendInPeriod(ctx context.Context, arg GetUserChatSpendInPeriodParams) (int64, error) GetUserCount(ctx context.Context, includeSystem bool) (int64, error) + // Returns the minimum (most restrictive) group limit for a user. + // Returns -1 if the user has no group limits applied. + GetUserGroupSpendLimit(ctx context.Context, userID uuid.UUID) (int64, error) // GetUserLatencyInsights returns the median and 95th percentile connection // latency that users have experienced. The result can be filtered on // template_ids, meaning only user data from workspaces based on those templates @@ -696,6 +708,8 @@ type sqlcQuerier interface { ListAIBridgeTokenUsagesByInterceptionIDs(ctx context.Context, interceptionIds []uuid.UUID) ([]AIBridgeTokenUsage, error) ListAIBridgeToolUsagesByInterceptionIDs(ctx context.Context, interceptionIds []uuid.UUID) ([]AIBridgeToolUsage, error) ListAIBridgeUserPromptsByInterceptionIDs(ctx context.Context, interceptionIds []uuid.UUID) ([]AIBridgeUserPrompt, error) + ListChatUsageLimitGroupOverrides(ctx context.Context) ([]ListChatUsageLimitGroupOverridesRow, error) + ListChatUsageLimitOverrides(ctx context.Context) ([]ListChatUsageLimitOverridesRow, error) ListProvisionerKeysByOrganization(ctx context.Context, organizationID uuid.UUID) ([]ProvisionerKey, error) ListProvisionerKeysByOrganizationExcludeReserved(ctx context.Context, organizationID uuid.UUID) ([]ProvisionerKey, error) ListTasks(ctx context.Context, arg ListTasksParams) ([]Task, error) @@ -716,6 +730,12 @@ type sqlcQuerier interface { ReduceWorkspaceAgentShareLevelToAuthenticatedByTemplate(ctx context.Context, templateID uuid.UUID) error RegisterWorkspaceProxy(ctx context.Context, arg RegisterWorkspaceProxyParams) (WorkspaceProxy, error) RemoveUserFromGroups(ctx context.Context, arg RemoveUserFromGroupsParams) ([]uuid.UUID, error) + // Resolves the effective spend limit for a user using the hierarchy: + // 1. Individual user override (highest priority) + // 2. Minimum group limit across all user's groups + // 3. Global default from config + // Returns -1 if limits are not enabled. + ResolveUserChatSpendLimit(ctx context.Context, userID uuid.UUID) (int64, error) RevokeDBCryptKey(ctx context.Context, activeKeyDigest string) error // Note that this selects from the CTE, not the original table. The CTE is named // the same as the original table to trick sqlc into reusing the existing struct @@ -847,6 +867,9 @@ type sqlcQuerier interface { UpsertChatDiffStatus(ctx context.Context, arg UpsertChatDiffStatusParams) (ChatDiffStatus, error) UpsertChatDiffStatusReference(ctx context.Context, arg UpsertChatDiffStatusReferenceParams) (ChatDiffStatus, error) UpsertChatSystemPrompt(ctx context.Context, value string) error + UpsertChatUsageLimitConfig(ctx context.Context, arg UpsertChatUsageLimitConfigParams) (ChatUsageLimitConfig, error) + UpsertChatUsageLimitGroupOverride(ctx context.Context, arg UpsertChatUsageLimitGroupOverrideParams) (UpsertChatUsageLimitGroupOverrideRow, error) + UpsertChatUsageLimitUserOverride(ctx context.Context, arg UpsertChatUsageLimitUserOverrideParams) (UpsertChatUsageLimitUserOverrideRow, error) UpsertConnectionLog(ctx context.Context, arg UpsertConnectionLogParams) (ConnectionLog, error) // The default proxy is implied and not actually stored in the database. // So we need to store it's configuration here for display purposes. diff --git a/coderd/database/queries.sql.go b/coderd/database/queries.sql.go index 7376d01004..e64529fd3c 100644 --- a/coderd/database/queries.sql.go +++ b/coderd/database/queries.sql.go @@ -3210,6 +3210,30 @@ func (q *sqlQuerier) BackoffChatDiffStatus(ctx context.Context, arg BackoffChatD return 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 deleteAllChatQueuedMessages = `-- name: DeleteAllChatQueuedMessages :exec DELETE FROM chat_queued_messages WHERE chat_id = $1 ` @@ -3251,6 +3275,24 @@ func (q *sqlQuerier) DeleteChatQueuedMessage(ctx context.Context, arg DeleteChat return err } +const deleteChatUsageLimitGroupOverride = `-- name: DeleteChatUsageLimitGroupOverride :exec +UPDATE groups SET chat_spend_limit_micros = NULL WHERE id = $1::uuid +` + +func (q *sqlQuerier) DeleteChatUsageLimitGroupOverride(ctx context.Context, groupID uuid.UUID) error { + _, err := q.db.ExecContext(ctx, deleteChatUsageLimitGroupOverride, groupID) + return err +} + +const deleteChatUsageLimitUserOverride = `-- name: DeleteChatUsageLimitUserOverride :exec +UPDATE users SET chat_spend_limit_micros = NULL WHERE id = $1::uuid +` + +func (q *sqlQuerier) DeleteChatUsageLimitUserOverride(ctx context.Context, userID uuid.UUID) error { + _, err := q.db.ExecContext(ctx, deleteChatUsageLimitUserOverride, userID) + return err +} + const getChatByID = `-- name: GetChatByID :one SELECT id, owner_id, workspace_id, title, status, worker_id, started_at, heartbeat_at, created_at, updated_at, parent_chat_id, root_chat_id, last_model_config_id, archived, last_error, mode @@ -4077,6 +4119,61 @@ func (q *sqlQuerier) GetChatQueuedMessages(ctx context.Context, chatID uuid.UUID return items, nil } +const getChatUsageLimitConfig = `-- name: GetChatUsageLimitConfig :one +SELECT id, singleton, enabled, default_limit_micros, period, created_at, updated_at FROM chat_usage_limit_config WHERE singleton = TRUE LIMIT 1 +` + +func (q *sqlQuerier) GetChatUsageLimitConfig(ctx context.Context) (ChatUsageLimitConfig, error) { + row := q.db.QueryRowContext(ctx, getChatUsageLimitConfig) + var i ChatUsageLimitConfig + err := row.Scan( + &i.ID, + &i.Singleton, + &i.Enabled, + &i.DefaultLimitMicros, + &i.Period, + &i.CreatedAt, + &i.UpdatedAt, + ) + return i, err +} + +const getChatUsageLimitGroupOverride = `-- name: GetChatUsageLimitGroupOverride :one +SELECT id AS group_id, chat_spend_limit_micros AS spend_limit_micros +FROM groups +WHERE id = $1::uuid AND chat_spend_limit_micros IS NOT NULL +` + +type GetChatUsageLimitGroupOverrideRow struct { + GroupID uuid.UUID `db:"group_id" json:"group_id"` + SpendLimitMicros sql.NullInt64 `db:"spend_limit_micros" json:"spend_limit_micros"` +} + +func (q *sqlQuerier) GetChatUsageLimitGroupOverride(ctx context.Context, groupID uuid.UUID) (GetChatUsageLimitGroupOverrideRow, error) { + row := q.db.QueryRowContext(ctx, getChatUsageLimitGroupOverride, groupID) + var i GetChatUsageLimitGroupOverrideRow + err := row.Scan(&i.GroupID, &i.SpendLimitMicros) + return i, err +} + +const getChatUsageLimitUserOverride = `-- name: GetChatUsageLimitUserOverride :one +SELECT id AS user_id, chat_spend_limit_micros AS spend_limit_micros +FROM users +WHERE id = $1::uuid AND chat_spend_limit_micros IS NOT NULL +` + +type GetChatUsageLimitUserOverrideRow struct { + UserID uuid.UUID `db:"user_id" json:"user_id"` + SpendLimitMicros sql.NullInt64 `db:"spend_limit_micros" json:"spend_limit_micros"` +} + +func (q *sqlQuerier) GetChatUsageLimitUserOverride(ctx context.Context, userID uuid.UUID) (GetChatUsageLimitUserOverrideRow, error) { + row := q.db.QueryRowContext(ctx, getChatUsageLimitUserOverride, userID) + var i GetChatUsageLimitUserOverrideRow + err := row.Scan(&i.UserID, &i.SpendLimitMicros) + return i, err +} + const getChatsByOwnerID = `-- name: GetChatsByOwnerID :many SELECT id, owner_id, workspace_id, title, status, worker_id, started_at, heartbeat_at, created_at, updated_at, parent_chat_id, root_chat_id, last_model_config_id, archived, last_error, mode @@ -4268,6 +4365,46 @@ func (q *sqlQuerier) GetStaleChats(ctx context.Context, staleThreshold time.Time return items, nil } +const getUserChatSpendInPeriod = `-- name: GetUserChatSpendInPeriod :one +SELECT COALESCE(SUM(cm.total_cost_micros), 0)::bigint AS total_spend_micros +FROM chat_messages cm +JOIN chats c ON c.id = cm.chat_id +WHERE c.owner_id = $1::uuid + AND cm.created_at >= $2::timestamptz + AND cm.created_at < $3::timestamptz + AND cm.total_cost_micros IS NOT NULL +` + +type GetUserChatSpendInPeriodParams struct { + UserID uuid.UUID `db:"user_id" json:"user_id"` + StartTime time.Time `db:"start_time" json:"start_time"` + EndTime time.Time `db:"end_time" json:"end_time"` +} + +func (q *sqlQuerier) GetUserChatSpendInPeriod(ctx context.Context, arg GetUserChatSpendInPeriodParams) (int64, error) { + row := q.db.QueryRowContext(ctx, getUserChatSpendInPeriod, arg.UserID, arg.StartTime, arg.EndTime) + var total_spend_micros int64 + err := row.Scan(&total_spend_micros) + return total_spend_micros, err +} + +const getUserGroupSpendLimit = `-- name: GetUserGroupSpendLimit :one +SELECT COALESCE(MIN(g.chat_spend_limit_micros), -1)::bigint AS limit_micros +FROM groups g +JOIN group_members_expanded gme ON gme.group_id = g.id +WHERE gme.user_id = $1::uuid + AND g.chat_spend_limit_micros IS NOT NULL +` + +// Returns the minimum (most restrictive) group limit for a user. +// Returns -1 if the user has no group limits applied. +func (q *sqlQuerier) GetUserGroupSpendLimit(ctx context.Context, userID uuid.UUID) (int64, error) { + row := q.db.QueryRowContext(ctx, getUserGroupSpendLimit, userID) + var limit_micros int64 + err := row.Scan(&limit_micros) + return limit_micros, err +} + const insertChat = `-- name: InsertChat :one INSERT INTO chats ( owner_id, @@ -4467,6 +4604,106 @@ func (q *sqlQuerier) InsertChatQueuedMessage(ctx context.Context, arg InsertChat return i, err } +const listChatUsageLimitGroupOverrides = `-- name: ListChatUsageLimitGroupOverrides :many +SELECT + g.id AS group_id, + g.name AS group_name, + g.display_name AS group_display_name, + g.avatar_url AS group_avatar_url, + g.chat_spend_limit_micros AS spend_limit_micros, + (SELECT COUNT(*) + FROM group_members_expanded gme + WHERE gme.group_id = g.id + AND gme.user_is_system = FALSE) AS member_count +FROM groups g +WHERE g.chat_spend_limit_micros IS NOT NULL +ORDER BY g.name ASC +` + +type ListChatUsageLimitGroupOverridesRow struct { + GroupID uuid.UUID `db:"group_id" json:"group_id"` + GroupName string `db:"group_name" json:"group_name"` + GroupDisplayName string `db:"group_display_name" json:"group_display_name"` + GroupAvatarUrl string `db:"group_avatar_url" json:"group_avatar_url"` + SpendLimitMicros sql.NullInt64 `db:"spend_limit_micros" json:"spend_limit_micros"` + MemberCount int64 `db:"member_count" json:"member_count"` +} + +func (q *sqlQuerier) ListChatUsageLimitGroupOverrides(ctx context.Context) ([]ListChatUsageLimitGroupOverridesRow, error) { + rows, err := q.db.QueryContext(ctx, listChatUsageLimitGroupOverrides) + if err != nil { + return nil, err + } + defer rows.Close() + var items []ListChatUsageLimitGroupOverridesRow + for rows.Next() { + var i ListChatUsageLimitGroupOverridesRow + if err := rows.Scan( + &i.GroupID, + &i.GroupName, + &i.GroupDisplayName, + &i.GroupAvatarUrl, + &i.SpendLimitMicros, + &i.MemberCount, + ); 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 listChatUsageLimitOverrides = `-- name: ListChatUsageLimitOverrides :many +SELECT u.id AS user_id, u.username, u.name, u.avatar_url, + u.chat_spend_limit_micros AS spend_limit_micros +FROM users u +WHERE u.chat_spend_limit_micros IS NOT NULL +ORDER BY u.username ASC +` + +type ListChatUsageLimitOverridesRow 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"` + SpendLimitMicros sql.NullInt64 `db:"spend_limit_micros" json:"spend_limit_micros"` +} + +func (q *sqlQuerier) ListChatUsageLimitOverrides(ctx context.Context) ([]ListChatUsageLimitOverridesRow, error) { + rows, err := q.db.QueryContext(ctx, listChatUsageLimitOverrides) + if err != nil { + return nil, err + } + defer rows.Close() + var items []ListChatUsageLimitOverridesRow + for rows.Next() { + var i ListChatUsageLimitOverridesRow + if err := rows.Scan( + &i.UserID, + &i.Username, + &i.Name, + &i.AvatarURL, + &i.SpendLimitMicros, + ); 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 popNextQueuedMessage = `-- name: PopNextQueuedMessage :one DELETE FROM chat_queued_messages WHERE id = ( @@ -4490,6 +4727,42 @@ func (q *sqlQuerier) PopNextQueuedMessage(ctx context.Context, chatID uuid.UUID) return i, err } +const resolveUserChatSpendLimit = `-- name: ResolveUserChatSpendLimit :one +SELECT CASE + -- If limits are disabled, return -1. + WHEN NOT cfg.enabled THEN -1 + -- Individual override takes priority. + WHEN u.chat_spend_limit_micros IS NOT NULL THEN u.chat_spend_limit_micros + -- Group limit (minimum across all user's groups) is next. + WHEN gl.limit_micros IS NOT NULL THEN gl.limit_micros + -- Fall back to global default. + ELSE cfg.default_limit_micros +END::bigint AS effective_limit_micros +FROM chat_usage_limit_config cfg +CROSS JOIN users u +LEFT JOIN LATERAL ( + SELECT MIN(g.chat_spend_limit_micros) AS limit_micros + FROM groups g + JOIN group_members_expanded gme ON gme.group_id = g.id + WHERE gme.user_id = $1::uuid + AND g.chat_spend_limit_micros IS NOT NULL +) gl ON TRUE +WHERE u.id = $1::uuid +LIMIT 1 +` + +// Resolves the effective spend limit for a user using the hierarchy: +// 1. Individual user override (highest priority) +// 2. Minimum group limit across all user's groups +// 3. Global default from config +// Returns -1 if limits are not enabled. +func (q *sqlQuerier) ResolveUserChatSpendLimit(ctx context.Context, userID uuid.UUID) (int64, error) { + row := q.db.QueryRowContext(ctx, resolveUserChatSpendLimit, userID) + var effective_limit_micros int64 + err := row.Scan(&effective_limit_micros) + return effective_limit_micros, err +} + const unarchiveChatByID = `-- name: UnarchiveChatByID :exec UPDATE chats SET archived = false, updated_at = NOW() WHERE id = $1::uuid ` @@ -4926,6 +5199,104 @@ func (q *sqlQuerier) UpsertChatDiffStatusReference(ctx context.Context, arg Upse return i, err } +const upsertChatUsageLimitConfig = `-- name: UpsertChatUsageLimitConfig :one +INSERT INTO chat_usage_limit_config (singleton, enabled, default_limit_micros, period, updated_at) +VALUES (TRUE, $1::boolean, $2::bigint, $3::text, NOW()) +ON CONFLICT (singleton) DO UPDATE SET + enabled = EXCLUDED.enabled, + default_limit_micros = EXCLUDED.default_limit_micros, + period = EXCLUDED.period, + updated_at = NOW() +RETURNING id, singleton, enabled, default_limit_micros, period, created_at, updated_at +` + +type UpsertChatUsageLimitConfigParams struct { + Enabled bool `db:"enabled" json:"enabled"` + DefaultLimitMicros int64 `db:"default_limit_micros" json:"default_limit_micros"` + Period string `db:"period" json:"period"` +} + +func (q *sqlQuerier) UpsertChatUsageLimitConfig(ctx context.Context, arg UpsertChatUsageLimitConfigParams) (ChatUsageLimitConfig, error) { + row := q.db.QueryRowContext(ctx, upsertChatUsageLimitConfig, arg.Enabled, arg.DefaultLimitMicros, arg.Period) + var i ChatUsageLimitConfig + err := row.Scan( + &i.ID, + &i.Singleton, + &i.Enabled, + &i.DefaultLimitMicros, + &i.Period, + &i.CreatedAt, + &i.UpdatedAt, + ) + return i, err +} + +const upsertChatUsageLimitGroupOverride = `-- name: UpsertChatUsageLimitGroupOverride :one +UPDATE groups +SET chat_spend_limit_micros = $1::bigint +WHERE id = $2::uuid +RETURNING id AS group_id, name, display_name, avatar_url, chat_spend_limit_micros AS spend_limit_micros +` + +type UpsertChatUsageLimitGroupOverrideParams struct { + SpendLimitMicros int64 `db:"spend_limit_micros" json:"spend_limit_micros"` + GroupID uuid.UUID `db:"group_id" json:"group_id"` +} + +type UpsertChatUsageLimitGroupOverrideRow struct { + GroupID uuid.UUID `db:"group_id" json:"group_id"` + Name string `db:"name" json:"name"` + DisplayName string `db:"display_name" json:"display_name"` + AvatarURL string `db:"avatar_url" json:"avatar_url"` + SpendLimitMicros sql.NullInt64 `db:"spend_limit_micros" json:"spend_limit_micros"` +} + +func (q *sqlQuerier) UpsertChatUsageLimitGroupOverride(ctx context.Context, arg UpsertChatUsageLimitGroupOverrideParams) (UpsertChatUsageLimitGroupOverrideRow, error) { + row := q.db.QueryRowContext(ctx, upsertChatUsageLimitGroupOverride, arg.SpendLimitMicros, arg.GroupID) + var i UpsertChatUsageLimitGroupOverrideRow + err := row.Scan( + &i.GroupID, + &i.Name, + &i.DisplayName, + &i.AvatarURL, + &i.SpendLimitMicros, + ) + return i, err +} + +const upsertChatUsageLimitUserOverride = `-- name: UpsertChatUsageLimitUserOverride :one +UPDATE users +SET chat_spend_limit_micros = $1::bigint +WHERE id = $2::uuid +RETURNING id AS user_id, username, name, avatar_url, chat_spend_limit_micros AS spend_limit_micros +` + +type UpsertChatUsageLimitUserOverrideParams struct { + SpendLimitMicros int64 `db:"spend_limit_micros" json:"spend_limit_micros"` + UserID uuid.UUID `db:"user_id" json:"user_id"` +} + +type UpsertChatUsageLimitUserOverrideRow 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"` + SpendLimitMicros sql.NullInt64 `db:"spend_limit_micros" json:"spend_limit_micros"` +} + +func (q *sqlQuerier) UpsertChatUsageLimitUserOverride(ctx context.Context, arg UpsertChatUsageLimitUserOverrideParams) (UpsertChatUsageLimitUserOverrideRow, error) { + row := q.db.QueryRowContext(ctx, upsertChatUsageLimitUserOverride, arg.SpendLimitMicros, arg.UserID) + var i UpsertChatUsageLimitUserOverrideRow + err := row.Scan( + &i.UserID, + &i.Username, + &i.Name, + &i.AvatarURL, + &i.SpendLimitMicros, + ) + return i, err +} + const countConnectionLogs = `-- name: CountConnectionLogs :one SELECT COUNT(*) AS count @@ -6533,7 +6904,7 @@ func (q *sqlQuerier) DeleteGroupByID(ctx context.Context, id uuid.UUID) error { const getGroupByID = `-- name: GetGroupByID :one SELECT - id, name, organization_id, avatar_url, quota_allowance, display_name, source + id, name, organization_id, avatar_url, quota_allowance, display_name, source, chat_spend_limit_micros FROM groups WHERE @@ -6553,13 +6924,14 @@ func (q *sqlQuerier) GetGroupByID(ctx context.Context, id uuid.UUID) (Group, err &i.QuotaAllowance, &i.DisplayName, &i.Source, + &i.ChatSpendLimitMicros, ) return i, err } const getGroupByOrgAndName = `-- name: GetGroupByOrgAndName :one SELECT - id, name, organization_id, avatar_url, quota_allowance, display_name, source + id, name, organization_id, avatar_url, quota_allowance, display_name, source, chat_spend_limit_micros FROM groups WHERE @@ -6586,13 +6958,14 @@ func (q *sqlQuerier) GetGroupByOrgAndName(ctx context.Context, arg GetGroupByOrg &i.QuotaAllowance, &i.DisplayName, &i.Source, + &i.ChatSpendLimitMicros, ) return i, err } const getGroups = `-- name: GetGroups :many SELECT - groups.id, groups.name, groups.organization_id, groups.avatar_url, groups.quota_allowance, groups.display_name, groups.source, + groups.id, groups.name, groups.organization_id, groups.avatar_url, groups.quota_allowance, groups.display_name, groups.source, groups.chat_spend_limit_micros, organizations.name AS organization_name, organizations.display_name AS organization_display_name FROM @@ -6667,6 +7040,7 @@ func (q *sqlQuerier) GetGroups(ctx context.Context, arg GetGroupsParams) ([]GetG &i.Group.QuotaAllowance, &i.Group.DisplayName, &i.Group.Source, + &i.Group.ChatSpendLimitMicros, &i.OrganizationName, &i.OrganizationDisplayName, ); err != nil { @@ -6690,7 +7064,7 @@ INSERT INTO groups ( organization_id ) VALUES - ($1, 'Everyone', $1) RETURNING id, name, organization_id, avatar_url, quota_allowance, display_name, source + ($1, 'Everyone', $1) RETURNING id, name, organization_id, avatar_url, quota_allowance, display_name, source, chat_spend_limit_micros ` // We use the organization_id as the id @@ -6707,6 +7081,7 @@ func (q *sqlQuerier) InsertAllUsersGroup(ctx context.Context, organizationID uui &i.QuotaAllowance, &i.DisplayName, &i.Source, + &i.ChatSpendLimitMicros, ) return i, err } @@ -6721,7 +7096,7 @@ INSERT INTO groups ( quota_allowance ) VALUES - ($1, $2, $3, $4, $5, $6) RETURNING id, name, organization_id, avatar_url, quota_allowance, display_name, source + ($1, $2, $3, $4, $5, $6) RETURNING id, name, organization_id, avatar_url, quota_allowance, display_name, source, chat_spend_limit_micros ` type InsertGroupParams struct { @@ -6751,6 +7126,7 @@ func (q *sqlQuerier) InsertGroup(ctx context.Context, arg InsertGroupParams) (Gr &i.QuotaAllowance, &i.DisplayName, &i.Source, + &i.ChatSpendLimitMicros, ) return i, err } @@ -6770,7 +7146,7 @@ SELECT FROM UNNEST($3 :: text[]) AS group_name ON CONFLICT DO NOTHING -RETURNING id, name, organization_id, avatar_url, quota_allowance, display_name, source +RETURNING id, name, organization_id, avatar_url, quota_allowance, display_name, source, chat_spend_limit_micros ` type InsertMissingGroupsParams struct { @@ -6800,6 +7176,7 @@ func (q *sqlQuerier) InsertMissingGroups(ctx context.Context, arg InsertMissingG &i.QuotaAllowance, &i.DisplayName, &i.Source, + &i.ChatSpendLimitMicros, ); err != nil { return nil, err } @@ -6824,7 +7201,7 @@ SET quota_allowance = $4 WHERE id = $5 -RETURNING id, name, organization_id, avatar_url, quota_allowance, display_name, source +RETURNING id, name, organization_id, avatar_url, quota_allowance, display_name, source, chat_spend_limit_micros ` type UpdateGroupByIDParams struct { @@ -6852,6 +7229,7 @@ func (q *sqlQuerier) UpdateGroupByID(ctx context.Context, arg UpdateGroupByIDPar &i.QuotaAllowance, &i.DisplayName, &i.Source, + &i.ChatSpendLimitMicros, ) return i, err } @@ -19003,7 +19381,7 @@ func (q *sqlQuerier) GetAuthorizationUserRoles(ctx context.Context, userID uuid. const getUserByEmailOrUsername = `-- name: GetUserByEmailOrUsername :one SELECT - id, email, username, hashed_password, created_at, updated_at, status, rbac_roles, login_type, avatar_url, deleted, last_seen_at, quiet_hours_schedule, name, github_com_user_id, hashed_one_time_passcode, one_time_passcode_expires_at, is_system, is_service_account + id, email, username, hashed_password, created_at, updated_at, status, rbac_roles, login_type, avatar_url, deleted, last_seen_at, quiet_hours_schedule, name, github_com_user_id, hashed_one_time_passcode, one_time_passcode_expires_at, is_system, is_service_account, chat_spend_limit_micros FROM users WHERE @@ -19041,13 +19419,14 @@ func (q *sqlQuerier) GetUserByEmailOrUsername(ctx context.Context, arg GetUserBy &i.OneTimePasscodeExpiresAt, &i.IsSystem, &i.IsServiceAccount, + &i.ChatSpendLimitMicros, ) return i, err } const getUserByID = `-- name: GetUserByID :one SELECT - id, email, username, hashed_password, created_at, updated_at, status, rbac_roles, login_type, avatar_url, deleted, last_seen_at, quiet_hours_schedule, name, github_com_user_id, hashed_one_time_passcode, one_time_passcode_expires_at, is_system, is_service_account + id, email, username, hashed_password, created_at, updated_at, status, rbac_roles, login_type, avatar_url, deleted, last_seen_at, quiet_hours_schedule, name, github_com_user_id, hashed_one_time_passcode, one_time_passcode_expires_at, is_system, is_service_account, chat_spend_limit_micros FROM users WHERE @@ -19079,6 +19458,7 @@ func (q *sqlQuerier) GetUserByID(ctx context.Context, id uuid.UUID) (User, error &i.OneTimePasscodeExpiresAt, &i.IsSystem, &i.IsServiceAccount, + &i.ChatSpendLimitMicros, ) return i, err } @@ -19170,7 +19550,7 @@ func (q *sqlQuerier) GetUserThemePreference(ctx context.Context, userID uuid.UUI const getUsers = `-- name: GetUsers :many SELECT - id, email, username, hashed_password, created_at, updated_at, status, rbac_roles, login_type, avatar_url, deleted, last_seen_at, quiet_hours_schedule, name, github_com_user_id, hashed_one_time_passcode, one_time_passcode_expires_at, is_system, is_service_account, COUNT(*) OVER() AS count + id, email, username, hashed_password, created_at, updated_at, status, rbac_roles, login_type, avatar_url, deleted, last_seen_at, quiet_hours_schedule, name, github_com_user_id, hashed_one_time_passcode, one_time_passcode_expires_at, is_system, is_service_account, chat_spend_limit_micros, COUNT(*) OVER() AS count FROM users WHERE @@ -19312,6 +19692,7 @@ type GetUsersRow struct { OneTimePasscodeExpiresAt sql.NullTime `db:"one_time_passcode_expires_at" json:"one_time_passcode_expires_at"` IsSystem bool `db:"is_system" json:"is_system"` IsServiceAccount bool `db:"is_service_account" json:"is_service_account"` + ChatSpendLimitMicros sql.NullInt64 `db:"chat_spend_limit_micros" json:"chat_spend_limit_micros"` Count int64 `db:"count" json:"count"` } @@ -19360,6 +19741,7 @@ func (q *sqlQuerier) GetUsers(ctx context.Context, arg GetUsersParams) ([]GetUse &i.OneTimePasscodeExpiresAt, &i.IsSystem, &i.IsServiceAccount, + &i.ChatSpendLimitMicros, &i.Count, ); err != nil { return nil, err @@ -19376,7 +19758,7 @@ func (q *sqlQuerier) GetUsers(ctx context.Context, arg GetUsersParams) ([]GetUse } const getUsersByIDs = `-- name: GetUsersByIDs :many -SELECT id, email, username, hashed_password, created_at, updated_at, status, rbac_roles, login_type, avatar_url, deleted, last_seen_at, quiet_hours_schedule, name, github_com_user_id, hashed_one_time_passcode, one_time_passcode_expires_at, is_system, is_service_account FROM users WHERE id = ANY($1 :: uuid [ ]) +SELECT id, email, username, hashed_password, created_at, updated_at, status, rbac_roles, login_type, avatar_url, deleted, last_seen_at, quiet_hours_schedule, name, github_com_user_id, hashed_one_time_passcode, one_time_passcode_expires_at, is_system, is_service_account, chat_spend_limit_micros FROM users WHERE id = ANY($1 :: uuid [ ]) ` // This shouldn't check for deleted, because it's frequently used @@ -19411,6 +19793,7 @@ func (q *sqlQuerier) GetUsersByIDs(ctx context.Context, ids []uuid.UUID) ([]User &i.OneTimePasscodeExpiresAt, &i.IsSystem, &i.IsServiceAccount, + &i.ChatSpendLimitMicros, ); err != nil { return nil, err } @@ -19446,7 +19829,7 @@ VALUES -- we were doing before. COALESCE(NULLIF($10::text, '')::user_status, 'dormant'::user_status), $11::bool - ) RETURNING id, email, username, hashed_password, created_at, updated_at, status, rbac_roles, login_type, avatar_url, deleted, last_seen_at, quiet_hours_schedule, name, github_com_user_id, hashed_one_time_passcode, one_time_passcode_expires_at, is_system, is_service_account + ) RETURNING id, email, username, hashed_password, created_at, updated_at, status, rbac_roles, login_type, avatar_url, deleted, last_seen_at, quiet_hours_schedule, name, github_com_user_id, hashed_one_time_passcode, one_time_passcode_expires_at, is_system, is_service_account, chat_spend_limit_micros ` type InsertUserParams struct { @@ -19498,6 +19881,7 @@ func (q *sqlQuerier) InsertUser(ctx context.Context, arg InsertUserParams) (User &i.OneTimePasscodeExpiresAt, &i.IsSystem, &i.IsServiceAccount, + &i.ChatSpendLimitMicros, ) return i, err } @@ -19664,7 +20048,7 @@ SET last_seen_at = $2, updated_at = $3 WHERE - id = $1 RETURNING id, email, username, hashed_password, created_at, updated_at, status, rbac_roles, login_type, avatar_url, deleted, last_seen_at, quiet_hours_schedule, name, github_com_user_id, hashed_one_time_passcode, one_time_passcode_expires_at, is_system, is_service_account + id = $1 RETURNING id, email, username, hashed_password, created_at, updated_at, status, rbac_roles, login_type, avatar_url, deleted, last_seen_at, quiet_hours_schedule, name, github_com_user_id, hashed_one_time_passcode, one_time_passcode_expires_at, is_system, is_service_account, chat_spend_limit_micros ` type UpdateUserLastSeenAtParams struct { @@ -19696,6 +20080,7 @@ func (q *sqlQuerier) UpdateUserLastSeenAt(ctx context.Context, arg UpdateUserLas &i.OneTimePasscodeExpiresAt, &i.IsSystem, &i.IsServiceAccount, + &i.ChatSpendLimitMicros, ) return i, err } @@ -19715,7 +20100,7 @@ SET WHERE id = $2 AND NOT is_system -RETURNING id, email, username, hashed_password, created_at, updated_at, status, rbac_roles, login_type, avatar_url, deleted, last_seen_at, quiet_hours_schedule, name, github_com_user_id, hashed_one_time_passcode, one_time_passcode_expires_at, is_system, is_service_account +RETURNING id, email, username, hashed_password, created_at, updated_at, status, rbac_roles, login_type, avatar_url, deleted, last_seen_at, quiet_hours_schedule, name, github_com_user_id, hashed_one_time_passcode, one_time_passcode_expires_at, is_system, is_service_account, chat_spend_limit_micros ` type UpdateUserLoginTypeParams struct { @@ -19746,6 +20131,7 @@ func (q *sqlQuerier) UpdateUserLoginType(ctx context.Context, arg UpdateUserLogi &i.OneTimePasscodeExpiresAt, &i.IsSystem, &i.IsServiceAccount, + &i.ChatSpendLimitMicros, ) return i, err } @@ -19761,7 +20147,7 @@ SET name = $6 WHERE id = $1 -RETURNING id, email, username, hashed_password, created_at, updated_at, status, rbac_roles, login_type, avatar_url, deleted, last_seen_at, quiet_hours_schedule, name, github_com_user_id, hashed_one_time_passcode, one_time_passcode_expires_at, is_system, is_service_account +RETURNING id, email, username, hashed_password, created_at, updated_at, status, rbac_roles, login_type, avatar_url, deleted, last_seen_at, quiet_hours_schedule, name, github_com_user_id, hashed_one_time_passcode, one_time_passcode_expires_at, is_system, is_service_account, chat_spend_limit_micros ` type UpdateUserProfileParams struct { @@ -19803,6 +20189,7 @@ func (q *sqlQuerier) UpdateUserProfile(ctx context.Context, arg UpdateUserProfil &i.OneTimePasscodeExpiresAt, &i.IsSystem, &i.IsServiceAccount, + &i.ChatSpendLimitMicros, ) return i, err } @@ -19814,7 +20201,7 @@ SET quiet_hours_schedule = $2 WHERE id = $1 -RETURNING id, email, username, hashed_password, created_at, updated_at, status, rbac_roles, login_type, avatar_url, deleted, last_seen_at, quiet_hours_schedule, name, github_com_user_id, hashed_one_time_passcode, one_time_passcode_expires_at, is_system, is_service_account +RETURNING id, email, username, hashed_password, created_at, updated_at, status, rbac_roles, login_type, avatar_url, deleted, last_seen_at, quiet_hours_schedule, name, github_com_user_id, hashed_one_time_passcode, one_time_passcode_expires_at, is_system, is_service_account, chat_spend_limit_micros ` type UpdateUserQuietHoursScheduleParams struct { @@ -19845,6 +20232,7 @@ func (q *sqlQuerier) UpdateUserQuietHoursSchedule(ctx context.Context, arg Updat &i.OneTimePasscodeExpiresAt, &i.IsSystem, &i.IsServiceAccount, + &i.ChatSpendLimitMicros, ) return i, err } @@ -19857,7 +20245,7 @@ SET rbac_roles = ARRAY(SELECT DISTINCT UNNEST($1 :: text[])) WHERE id = $2 -RETURNING id, email, username, hashed_password, created_at, updated_at, status, rbac_roles, login_type, avatar_url, deleted, last_seen_at, quiet_hours_schedule, name, github_com_user_id, hashed_one_time_passcode, one_time_passcode_expires_at, is_system, is_service_account +RETURNING id, email, username, hashed_password, created_at, updated_at, status, rbac_roles, login_type, avatar_url, deleted, last_seen_at, quiet_hours_schedule, name, github_com_user_id, hashed_one_time_passcode, one_time_passcode_expires_at, is_system, is_service_account, chat_spend_limit_micros ` type UpdateUserRolesParams struct { @@ -19888,6 +20276,7 @@ func (q *sqlQuerier) UpdateUserRoles(ctx context.Context, arg UpdateUserRolesPar &i.OneTimePasscodeExpiresAt, &i.IsSystem, &i.IsServiceAccount, + &i.ChatSpendLimitMicros, ) return i, err } @@ -19901,7 +20290,7 @@ SET -- If the user is logging in, set last_seen_at to updated_at. last_seen_at = CASE WHEN $4 :: boolean THEN $3 :: timestamptz ELSE last_seen_at END WHERE - id = $1 RETURNING id, email, username, hashed_password, created_at, updated_at, status, rbac_roles, login_type, avatar_url, deleted, last_seen_at, quiet_hours_schedule, name, github_com_user_id, hashed_one_time_passcode, one_time_passcode_expires_at, is_system, is_service_account + id = $1 RETURNING id, email, username, hashed_password, created_at, updated_at, status, rbac_roles, login_type, avatar_url, deleted, last_seen_at, quiet_hours_schedule, name, github_com_user_id, hashed_one_time_passcode, one_time_passcode_expires_at, is_system, is_service_account, chat_spend_limit_micros ` type UpdateUserStatusParams struct { @@ -19939,6 +20328,7 @@ func (q *sqlQuerier) UpdateUserStatus(ctx context.Context, arg UpdateUserStatusP &i.OneTimePasscodeExpiresAt, &i.IsSystem, &i.IsServiceAccount, + &i.ChatSpendLimitMicros, ) return i, err } diff --git a/coderd/database/queries/chats.sql b/coderd/database/queries/chats.sql index 6b994fe3e2..9a3e1c8c35 100644 --- a/coderd/database/queries/chats.sql +++ b/coderd/database/queries/chats.sql @@ -700,3 +700,128 @@ LIMIT sqlc.arg('page_limit')::int OFFSET sqlc.arg('page_offset')::int; + +-- name: GetChatUsageLimitConfig :one +SELECT * FROM chat_usage_limit_config WHERE singleton = TRUE LIMIT 1; + +-- name: UpsertChatUsageLimitConfig :one +INSERT INTO chat_usage_limit_config (singleton, enabled, default_limit_micros, period, updated_at) +VALUES (TRUE, @enabled::boolean, @default_limit_micros::bigint, @period::text, NOW()) +ON CONFLICT (singleton) DO UPDATE SET + enabled = EXCLUDED.enabled, + default_limit_micros = EXCLUDED.default_limit_micros, + period = EXCLUDED.period, + updated_at = NOW() +RETURNING *; + +-- name: ListChatUsageLimitOverrides :many +SELECT u.id AS user_id, u.username, u.name, u.avatar_url, + u.chat_spend_limit_micros AS spend_limit_micros +FROM users u +WHERE u.chat_spend_limit_micros IS NOT NULL +ORDER BY u.username ASC; + +-- name: UpsertChatUsageLimitUserOverride :one +UPDATE users +SET chat_spend_limit_micros = @spend_limit_micros::bigint +WHERE id = @user_id::uuid +RETURNING id AS user_id, username, name, avatar_url, chat_spend_limit_micros AS spend_limit_micros; + +-- name: DeleteChatUsageLimitUserOverride :exec +UPDATE users SET chat_spend_limit_micros = NULL WHERE id = @user_id::uuid; + +-- name: GetChatUsageLimitUserOverride :one +SELECT id AS user_id, chat_spend_limit_micros AS spend_limit_micros +FROM users +WHERE id = @user_id::uuid AND chat_spend_limit_micros IS NOT NULL; + +-- name: GetUserChatSpendInPeriod :one +SELECT COALESCE(SUM(cm.total_cost_micros), 0)::bigint AS total_spend_micros +FROM chat_messages cm +JOIN chats c ON c.id = cm.chat_id +WHERE c.owner_id = @user_id::uuid + AND cm.created_at >= @start_time::timestamptz + AND cm.created_at < @end_time::timestamptz + AND cm.total_cost_micros 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: ListChatUsageLimitGroupOverrides :many +SELECT + g.id AS group_id, + g.name AS group_name, + g.display_name AS group_display_name, + g.avatar_url AS group_avatar_url, + g.chat_spend_limit_micros AS spend_limit_micros, + (SELECT COUNT(*) + FROM group_members_expanded gme + WHERE gme.group_id = g.id + AND gme.user_is_system = FALSE) AS member_count +FROM groups g +WHERE g.chat_spend_limit_micros IS NOT NULL +ORDER BY g.name ASC; + +-- name: UpsertChatUsageLimitGroupOverride :one +UPDATE groups +SET chat_spend_limit_micros = @spend_limit_micros::bigint +WHERE id = @group_id::uuid +RETURNING id AS group_id, name, display_name, avatar_url, chat_spend_limit_micros AS spend_limit_micros; + +-- name: DeleteChatUsageLimitGroupOverride :exec +UPDATE groups SET chat_spend_limit_micros = NULL WHERE id = @group_id::uuid; + +-- name: GetChatUsageLimitGroupOverride :one +SELECT id AS group_id, chat_spend_limit_micros AS spend_limit_micros +FROM groups +WHERE id = @group_id::uuid AND chat_spend_limit_micros IS NOT NULL; + +-- name: GetUserGroupSpendLimit :one +-- Returns the minimum (most restrictive) group limit for a user. +-- Returns -1 if the user has no group limits applied. +SELECT COALESCE(MIN(g.chat_spend_limit_micros), -1)::bigint AS limit_micros +FROM groups g +JOIN group_members_expanded gme ON gme.group_id = g.id +WHERE gme.user_id = @user_id::uuid + AND g.chat_spend_limit_micros IS NOT NULL; + +-- name: ResolveUserChatSpendLimit :one +-- Resolves the effective spend limit for a user using the hierarchy: +-- 1. Individual user override (highest priority) +-- 2. Minimum group limit across all user's groups +-- 3. Global default from config +-- Returns -1 if limits are not enabled. +SELECT CASE + -- If limits are disabled, return -1. + WHEN NOT cfg.enabled THEN -1 + -- Individual override takes priority. + WHEN u.chat_spend_limit_micros IS NOT NULL THEN u.chat_spend_limit_micros + -- Group limit (minimum across all user's groups) is next. + WHEN gl.limit_micros IS NOT NULL THEN gl.limit_micros + -- Fall back to global default. + ELSE cfg.default_limit_micros +END::bigint AS effective_limit_micros +FROM chat_usage_limit_config cfg +CROSS JOIN users u +LEFT JOIN LATERAL ( + SELECT MIN(g.chat_spend_limit_micros) AS limit_micros + FROM groups g + JOIN group_members_expanded gme ON gme.group_id = g.id + WHERE gme.user_id = @user_id::uuid + AND g.chat_spend_limit_micros IS NOT NULL +) gl ON TRUE +WHERE u.id = @user_id::uuid +LIMIT 1; diff --git a/coderd/database/unique_constraint.go b/coderd/database/unique_constraint.go index 6066e4ea50..35f40d7c5f 100644 --- a/coderd/database/unique_constraint.go +++ b/coderd/database/unique_constraint.go @@ -22,6 +22,8 @@ const ( UniqueChatProvidersPkey UniqueConstraint = "chat_providers_pkey" // ALTER TABLE ONLY chat_providers ADD CONSTRAINT chat_providers_pkey PRIMARY KEY (id); UniqueChatProvidersProviderKey UniqueConstraint = "chat_providers_provider_key" // ALTER TABLE ONLY chat_providers ADD CONSTRAINT chat_providers_provider_key UNIQUE (provider); UniqueChatQueuedMessagesPkey UniqueConstraint = "chat_queued_messages_pkey" // ALTER TABLE ONLY chat_queued_messages ADD CONSTRAINT chat_queued_messages_pkey PRIMARY KEY (id); + UniqueChatUsageLimitConfigPkey UniqueConstraint = "chat_usage_limit_config_pkey" // ALTER TABLE ONLY chat_usage_limit_config ADD CONSTRAINT chat_usage_limit_config_pkey PRIMARY KEY (id); + UniqueChatUsageLimitConfigSingletonKey UniqueConstraint = "chat_usage_limit_config_singleton_key" // ALTER TABLE ONLY chat_usage_limit_config ADD CONSTRAINT chat_usage_limit_config_singleton_key UNIQUE (singleton); UniqueChatsPkey UniqueConstraint = "chats_pkey" // ALTER TABLE ONLY chats ADD CONSTRAINT chats_pkey PRIMARY KEY (id); UniqueConnectionLogsPkey UniqueConstraint = "connection_logs_pkey" // ALTER TABLE ONLY connection_logs ADD CONSTRAINT connection_logs_pkey PRIMARY KEY (id); UniqueCryptoKeysPkey UniqueConstraint = "crypto_keys_pkey" // ALTER TABLE ONLY crypto_keys ADD CONSTRAINT crypto_keys_pkey PRIMARY KEY (feature, sequence); diff --git a/codersdk/chats.go b/codersdk/chats.go index 35dad9ca41..ef0cc4378a 100644 --- a/codersdk/chats.go +++ b/codersdk/chats.go @@ -1,8 +1,10 @@ package codersdk import ( + "bytes" "context" "encoding/json" + "errors" "fmt" "io" "mime" @@ -14,6 +16,7 @@ import ( "github.com/google/uuid" "github.com/shopspring/decimal" + "golang.org/x/xerrors" "github.com/coder/websocket" "github.com/coder/websocket/wsjson" @@ -746,6 +749,7 @@ type ChatCostSummary struct { TotalCacheCreationTokens int64 `json:"total_cache_creation_tokens"` ByModel []ChatCostModelBreakdown `json:"by_model"` ByChat []ChatCostChatBreakdown `json:"by_chat"` + UsageLimit *ChatUsageLimitStatus `json:"usage_limit,omitempty"` } // ChatCostModelBreakdown contains per-model cost aggregation. @@ -797,6 +801,223 @@ type ChatCostUsersResponse struct { Users []ChatCostUserRollup `json:"users"` } +// ChatUsageLimitExceededResponse is the 409 response body returned when a +// chat operation exceeds the caller's usage limit. The structured fields let +// frontends render user-friendly spend, limit, and reset information without +// parsing debug text. +type ChatUsageLimitExceededResponse struct { + Response + SpentMicros int64 `json:"spent_micros"` + LimitMicros int64 `json:"limit_micros"` + ResetsAt time.Time `json:"resets_at" format:"date-time"` +} + +type chatUsageLimitExceededError struct { + err *Error + response ChatUsageLimitExceededResponse +} + +func (e *chatUsageLimitExceededError) Error() string { + if e.err == nil { + return e.response.Message + } + return e.err.Error() +} + +func (e *chatUsageLimitExceededError) Unwrap() error { + return e.err +} + +func readBodyAsChatUsageLimitError(res *http.Response) error { + if res == nil || res.StatusCode != http.StatusConflict { + return ReadBodyAsError(res) + } + defer res.Body.Close() + + rawBody, err := io.ReadAll(res.Body) + if err != nil { + return xerrors.Errorf("read body: %w", err) + } + + if mimeErr := ExpectJSONMime(res); mimeErr != nil { + return readRawBodyAsError(res, rawBody) + } + + var payload ChatUsageLimitExceededResponse + if err := json.NewDecoder(bytes.NewReader(rawBody)).Decode(&payload); err == nil && isChatUsageLimitExceededResponse(payload) { + return &chatUsageLimitExceededError{ + err: newResponseError(res, payload.Response), + response: payload, + } + } + + return readRawBodyAsError(res, rawBody) +} + +func isChatUsageLimitExceededResponse(resp ChatUsageLimitExceededResponse) bool { + return resp.Message != "" && !resp.ResetsAt.IsZero() +} + +func readRawBodyAsError(res *http.Response, rawBody []byte) error { + if mimeErr := ExpectJSONMime(res); mimeErr != nil { + if len(rawBody) > 2048 { + rawBody = append(rawBody[:2048], []byte("...")...) + } + if len(rawBody) == 0 { + rawBody = []byte("no response body") + } + return newResponseError(res, Response{ + Message: mimeErr.Error(), + Detail: string(rawBody), + }) + } + + var response Response + if err := json.NewDecoder(bytes.NewReader(rawBody)).Decode(&response); err != nil { + if errors.Is(err, io.EOF) { + return newResponseError(res, Response{Message: "empty response body"}) + } + return xerrors.Errorf("decode body: %w", err) + } + if response.Message == "" { + if len(rawBody) > 1024 { + rawBody = append(rawBody[:1024], []byte("...")...) + } + response.Message = fmt.Sprintf( + "unexpected status code %d, response has no message", + res.StatusCode, + ) + response.Detail = string(rawBody) + } + return newResponseError(res, response) +} + +func newResponseError(res *http.Response, response Response) *Error { + if res == nil { + return &Error{Response: response} + } + + var requestMethod, requestURL string + if res.Request != nil { + requestMethod = res.Request.Method + if res.Request.URL != nil { + requestURL = res.Request.URL.String() + } + } + + var helpMessage string + if res.StatusCode == http.StatusUnauthorized { + helpMessage = "Try logging in using 'coder login'." + } + + return &Error{ + Response: response, + statusCode: res.StatusCode, + method: requestMethod, + url: requestURL, + Helper: helpMessage, + } +} + +// ChatUsageLimitExceededFrom extracts a structured chat usage limit response +// from an SDK error returned by chat mutation methods. +func ChatUsageLimitExceededFrom(err error) *ChatUsageLimitExceededResponse { + var limitErr *chatUsageLimitExceededError + if !errors.As(err, &limitErr) { + return nil + } + return &limitErr.response +} + +// ChatUsageLimitPeriod represents the time window for usage limits. +type ChatUsageLimitPeriod string + +const ( + ChatUsageLimitPeriodDay ChatUsageLimitPeriod = "day" + ChatUsageLimitPeriodWeek ChatUsageLimitPeriod = "week" + ChatUsageLimitPeriodMonth ChatUsageLimitPeriod = "month" +) + +// Valid reports whether p is a supported chat usage limit period. +func (p ChatUsageLimitPeriod) Valid() bool { + switch p { + case ChatUsageLimitPeriodDay, ChatUsageLimitPeriodWeek, ChatUsageLimitPeriodMonth: + return true + default: + return false + } +} + +// ChatUsageLimitConfig is the deployment-wide default usage limit config. +type ChatUsageLimitConfig struct { + // Nil in the API means no default limit is set. The DB stores 0 when + // limiting is disabled. + SpendLimitMicros *int64 `json:"spend_limit_micros"` + Period ChatUsageLimitPeriod `json:"period"` + UpdatedAt time.Time `json:"updated_at" format:"date-time"` +} + +// ChatUsageLimitOverride is a per-user override of the deployment default. +type ChatUsageLimitOverride struct { + UserID uuid.UUID `json:"user_id" format:"uuid"` + Username string `json:"username"` + Name string `json:"name"` + AvatarURL string `json:"avatar_url"` + // Nil in the API means no user override is set. Persisted override rows + // store positive values. + SpendLimitMicros *int64 `json:"spend_limit_micros"` +} + +// ChatUsageLimitGroupOverride represents a group-scoped spend limit override. +type ChatUsageLimitGroupOverride struct { + GroupID uuid.UUID `json:"group_id" format:"uuid"` + GroupName string `json:"group_name"` + GroupDisplayName string `json:"group_display_name"` + GroupAvatarURL string `json:"group_avatar_url"` + MemberCount int64 `json:"member_count"` + // Nil in the API means no group override is set. Persisted override rows + // store positive values. + SpendLimitMicros *int64 `json:"spend_limit_micros"` +} + +// UpsertChatUsageLimitOverrideRequest is the body for creating/updating a +// per-user usage limit override. +type UpsertChatUsageLimitOverrideRequest struct { + SpendLimitMicros int64 `json:"spend_limit_micros"` // Must be greater than 0. +} + +// UpdateChatUsageLimitOverrideRequest is kept as a compatibility alias. +type UpdateChatUsageLimitOverrideRequest = UpsertChatUsageLimitOverrideRequest + +// UpsertChatUsageLimitGroupOverrideRequest is the request to create or update +// a group-level spend limit override. +type UpsertChatUsageLimitGroupOverrideRequest struct { + SpendLimitMicros int64 `json:"spend_limit_micros"` // Must be greater than 0. +} + +// UpdateChatUsageLimitGroupOverrideRequest is kept as a compatibility alias. +type UpdateChatUsageLimitGroupOverrideRequest = UpsertChatUsageLimitGroupOverrideRequest + +// ChatUsageLimitStatus represents the current spend status for a user +// within their active limit period. +type ChatUsageLimitStatus struct { + IsLimited bool `json:"is_limited"` + Period ChatUsageLimitPeriod `json:"period,omitempty"` + SpendLimitMicros *int64 `json:"spend_limit_micros,omitempty"` + CurrentSpend int64 `json:"current_spend"` + PeriodStart time.Time `json:"period_start,omitempty" format:"date-time"` + PeriodEnd time.Time `json:"period_end,omitempty" format:"date-time"` +} + +// ChatUsageLimitConfigResponse is returned from the admin config endpoint +// and includes the config plus a count of models without pricing. +type ChatUsageLimitConfigResponse struct { + ChatUsageLimitConfig + UnpricedModelCount int64 `json:"unpriced_model_count"` + Overrides []ChatUsageLimitOverride `json:"overrides"` + GroupOverrides []ChatUsageLimitGroupOverride `json:"group_overrides"` +} + // ListChatsOptions are optional parameters for ListChats. type ListChatsOptions struct { Query string @@ -1085,10 +1306,10 @@ func (c *Client) CreateChat(ctx context.Context, req CreateChatRequest) (Chat, e if err != nil { return Chat{}, err } - defer res.Body.Close() if res.StatusCode != http.StatusCreated { - return Chat{}, ReadBodyAsError(res) + return Chat{}, readBodyAsChatUsageLimitError(res) } + defer res.Body.Close() var chat Chat return chat, json.NewDecoder(res.Body).Decode(&chat) } @@ -1308,10 +1529,10 @@ func (c *Client) CreateChatMessage(ctx context.Context, chatID uuid.UUID, req Cr if err != nil { return CreateChatMessageResponse{}, err } - defer res.Body.Close() if res.StatusCode != http.StatusOK { - return CreateChatMessageResponse{}, ReadBodyAsError(res) + return CreateChatMessageResponse{}, readBodyAsChatUsageLimitError(res) } + defer res.Body.Close() var resp CreateChatMessageResponse return resp, json.NewDecoder(res.Body).Decode(&resp) } @@ -1332,10 +1553,10 @@ func (c *Client) EditChatMessage( if err != nil { return ChatMessage{}, err } - defer res.Body.Close() if res.StatusCode != http.StatusOK { - return ChatMessage{}, ReadBodyAsError(res) + return ChatMessage{}, readBodyAsChatUsageLimitError(res) } + defer res.Body.Close() var message ChatMessage return message, json.NewDecoder(res.Body).Decode(&message) } @@ -1418,6 +1639,120 @@ func (c *Client) GetChatFile(ctx context.Context, fileID uuid.UUID) ([]byte, str return data, res.Header.Get("Content-Type"), nil } +// GetChatUsageLimitConfig returns the deployment-wide chat usage limit config. +func (c *Client) GetChatUsageLimitConfig(ctx context.Context) (ChatUsageLimitConfigResponse, error) { + res, err := c.Request(ctx, http.MethodGet, "/api/experimental/chats/usage-limits", nil) + if err != nil { + return ChatUsageLimitConfigResponse{}, err + } + defer res.Body.Close() + if res.StatusCode != http.StatusOK { + return ChatUsageLimitConfigResponse{}, ReadBodyAsError(res) + } + var resp ChatUsageLimitConfigResponse + return resp, json.NewDecoder(res.Body).Decode(&resp) +} + +// UpdateChatUsageLimitConfig updates the deployment-wide usage limit config. +func (c *Client) UpdateChatUsageLimitConfig(ctx context.Context, req ChatUsageLimitConfig) (ChatUsageLimitConfig, error) { + res, err := c.Request(ctx, http.MethodPut, "/api/experimental/chats/usage-limits", req) + if err != nil { + return ChatUsageLimitConfig{}, err + } + defer res.Body.Close() + if res.StatusCode != http.StatusOK { + return ChatUsageLimitConfig{}, ReadBodyAsError(res) + } + var resp ChatUsageLimitConfig + return resp, json.NewDecoder(res.Body).Decode(&resp) +} + +// UpsertChatUsageLimitOverride creates or updates a per-user usage limit override. +func (c *Client) UpsertChatUsageLimitOverride(ctx context.Context, userID uuid.UUID, req UpsertChatUsageLimitOverrideRequest) (ChatUsageLimitOverride, error) { + res, err := c.Request(ctx, http.MethodPut, fmt.Sprintf("/api/experimental/chats/usage-limits/overrides/%s", userID), req) + if err != nil { + return ChatUsageLimitOverride{}, err + } + defer res.Body.Close() + if res.StatusCode != http.StatusOK { + return ChatUsageLimitOverride{}, ReadBodyAsError(res) + } + var resp ChatUsageLimitOverride + return resp, json.NewDecoder(res.Body).Decode(&resp) +} + +// UpdateChatUserUsageLimitOverride creates or updates a per-user usage limit override. +func (c *Client) UpdateChatUserUsageLimitOverride(ctx context.Context, userID uuid.UUID, req UpdateChatUsageLimitOverrideRequest) (ChatUsageLimitOverride, error) { + return c.UpsertChatUsageLimitOverride(ctx, userID, req) +} + +// DeleteChatUsageLimitOverride removes a per-user usage limit override. +func (c *Client) DeleteChatUsageLimitOverride(ctx context.Context, userID uuid.UUID) error { + res, err := c.Request(ctx, http.MethodDelete, fmt.Sprintf("/api/experimental/chats/usage-limits/overrides/%s", userID), nil) + if err != nil { + return err + } + defer res.Body.Close() + if res.StatusCode != http.StatusNoContent { + return ReadBodyAsError(res) + } + return nil +} + +// DeleteChatUserUsageLimitOverride removes a per-user usage limit override. +func (c *Client) DeleteChatUserUsageLimitOverride(ctx context.Context, userID uuid.UUID) error { + return c.DeleteChatUsageLimitOverride(ctx, userID) +} + +// UpsertChatUsageLimitGroupOverride creates or updates a group-level +// spend limit override. EXPERIMENTAL: This API is subject to change. +func (c *Client) UpsertChatUsageLimitGroupOverride(ctx context.Context, groupID uuid.UUID, req UpsertChatUsageLimitGroupOverrideRequest) (ChatUsageLimitGroupOverride, error) { + res, err := c.Request(ctx, http.MethodPut, + fmt.Sprintf("/api/experimental/chats/usage-limits/group-overrides/%s", groupID), + req, + ) + if err != nil { + return ChatUsageLimitGroupOverride{}, err + } + defer res.Body.Close() + if res.StatusCode != http.StatusOK { + return ChatUsageLimitGroupOverride{}, ReadBodyAsError(res) + } + var override ChatUsageLimitGroupOverride + return override, json.NewDecoder(res.Body).Decode(&override) +} + +// DeleteChatUsageLimitGroupOverride removes a group-level spend limit +// override. EXPERIMENTAL: This API is subject to change. +func (c *Client) DeleteChatUsageLimitGroupOverride(ctx context.Context, groupID uuid.UUID) error { + res, err := c.Request(ctx, http.MethodDelete, + fmt.Sprintf("/api/experimental/chats/usage-limits/group-overrides/%s", groupID), + nil, + ) + if err != nil { + return err + } + defer res.Body.Close() + if res.StatusCode != http.StatusNoContent { + return ReadBodyAsError(res) + } + return nil +} + +// GetMyChatUsageLimitStatus returns the current user's chat usage limit status. +func (c *Client) GetMyChatUsageLimitStatus(ctx context.Context) (ChatUsageLimitStatus, error) { + res, err := c.Request(ctx, http.MethodGet, "/api/experimental/chats/usage-limits/status", nil) + if err != nil { + return ChatUsageLimitStatus{}, err + } + defer res.Body.Close() + if res.StatusCode != http.StatusOK { + return ChatUsageLimitStatus{}, ReadBodyAsError(res) + } + var resp ChatUsageLimitStatus + return resp, json.NewDecoder(res.Body).Decode(&resp) +} + func formatChatStreamResponseError(response Response) string { message := strings.TrimSpace(response.Message) detail := strings.TrimSpace(response.Detail) diff --git a/codersdk/chats_test.go b/codersdk/chats_test.go index 7526ea23ce..b2494562ce 100644 --- a/codersdk/chats_test.go +++ b/codersdk/chats_test.go @@ -1,8 +1,13 @@ package codersdk_test import ( + "context" "encoding/json" + "net/http" + "net/http/httptest" + "net/url" "testing" + "time" "github.com/google/uuid" "github.com/shopspring/decimal" @@ -55,6 +60,81 @@ func TestChatModelProviderOptions_UnmarshalJSON_ParsesPlainProviderPayloads(t *t ) } +func TestChatUsageLimitExceededFrom(t *testing.T) { + t.Parallel() + + t.Run("ExtractsTyped409", func(t *testing.T) { + t.Parallel() + + want := codersdk.ChatUsageLimitExceededResponse{ + Response: codersdk.Response{Message: "Chat usage limit exceeded."}, + SpentMicros: 123, + LimitMicros: 456, + ResetsAt: time.Date(2026, time.March, 16, 12, 0, 0, 0, time.UTC), + } + + srv := httptest.NewServer(http.HandlerFunc(func(rw http.ResponseWriter, r *http.Request) { + require.Equal(t, http.MethodPost, r.Method) + require.Equal(t, "/api/experimental/chats", r.URL.Path) + rw.Header().Set("Content-Type", "application/json") + rw.WriteHeader(http.StatusConflict) + require.NoError(t, json.NewEncoder(rw).Encode(want)) + })) + defer srv.Close() + + serverURL, err := url.Parse(srv.URL) + require.NoError(t, err) + + client := codersdk.New(serverURL) + _, err = client.CreateChat(context.Background(), codersdk.CreateChatRequest{ + Content: []codersdk.ChatInputPart{{ + Type: codersdk.ChatInputPartTypeText, + Text: "hello", + }}, + }) + require.Error(t, err) + + sdkErr, ok := codersdk.AsError(err) + require.True(t, ok) + require.Equal(t, http.StatusConflict, sdkErr.StatusCode()) + require.Equal(t, want.Message, sdkErr.Message) + + limitErr := codersdk.ChatUsageLimitExceededFrom(err) + require.NotNil(t, limitErr) + require.Equal(t, want, *limitErr) + }) + + t.Run("ReturnsNilForNonLimitErrors", func(t *testing.T) { + t.Parallel() + + require.Nil(t, codersdk.ChatUsageLimitExceededFrom(codersdk.NewError(http.StatusConflict, codersdk.Response{Message: "plain conflict"}))) + + srv := httptest.NewServer(http.HandlerFunc(func(rw http.ResponseWriter, r *http.Request) { + rw.Header().Set("Content-Type", "application/json") + rw.WriteHeader(http.StatusBadRequest) + require.NoError(t, json.NewEncoder(rw).Encode(codersdk.Response{Message: "Invalid request."})) + })) + defer srv.Close() + + serverURL, err := url.Parse(srv.URL) + require.NoError(t, err) + + client := codersdk.New(serverURL) + _, err = client.CreateChat(context.Background(), codersdk.CreateChatRequest{ + Content: []codersdk.ChatInputPart{{ + Type: codersdk.ChatInputPartTypeText, + Text: "hello", + }}, + }) + require.Error(t, err) + + sdkErr, ok := codersdk.AsError(err) + require.True(t, ok) + require.Equal(t, http.StatusBadRequest, sdkErr.StatusCode()) + require.Nil(t, codersdk.ChatUsageLimitExceededFrom(err)) + }) +} + func TestChatMessagePart_StripInternal(t *testing.T) { t.Parallel() diff --git a/docs/admin/security/audit-logs.md b/docs/admin/security/audit-logs.md index aed4b96291..421eee64de 100644 --- a/docs/admin/security/audit-logs.md +++ b/docs/admin/security/audit-logs.md @@ -18,7 +18,7 @@ We track the following resources: | APIKey
login, logout, register, create, write, delete | |
FieldTracked
allow_listfalse
created_attrue
expires_attrue
hashed_secretfalse
idfalse
ip_addressfalse
last_usedtrue
lifetime_secondsfalse
login_typefalse
scopesfalse
token_namefalse
updated_atfalse
user_idtrue
| | AiSeatState
create | |
FieldTracked
first_used_attrue
last_event_descriptiontrue
last_event_typetrue
last_used_atfalse
updated_atfalse
user_idtrue
| | AuditOAuthConvertState
| |
FieldTracked
created_attrue
expires_attrue
from_login_typetrue
to_login_typetrue
user_idtrue
| -| Group
create, write, delete | |
FieldTracked
avatar_urltrue
display_nametrue
idtrue
memberstrue
nametrue
organization_idfalse
quota_allowancetrue
sourcefalse
| +| Group
create, write, delete | |
FieldTracked
avatar_urltrue
chat_spend_limit_microstrue
display_nametrue
idtrue
memberstrue
nametrue
organization_idfalse
quota_allowancetrue
sourcefalse
| | AuditableOrganizationMember
| |
FieldTracked
created_attrue
organization_idfalse
rolestrue
updated_attrue
user_idtrue
usernametrue
| | CustomRole
| |
FieldTracked
created_atfalse
display_nametrue
idfalse
is_systemfalse
member_permissionstrue
nametrue
org_permissionstrue
organization_idfalse
site_permissionstrue
updated_atfalse
user_permissionstrue
| | GitSSHKey
create | |
FieldTracked
created_atfalse
private_keytrue
public_keytrue
updated_atfalse
user_idtrue
| @@ -36,7 +36,7 @@ We track the following resources: | TaskTable
| |
FieldTracked
created_atfalse
deleted_atfalse
display_nametrue
idtrue
nametrue
organization_idfalse
owner_idtrue
prompttrue
template_parameterstrue
template_version_idtrue
workspace_idtrue
| | Template
write, delete | |
FieldTracked
active_version_idtrue
activity_bumptrue
allow_user_autostarttrue
allow_user_autostoptrue
allow_user_cancel_workspace_jobstrue
autostart_block_days_of_weektrue
autostop_requirement_days_of_weektrue
autostop_requirement_weekstrue
cors_behaviortrue
created_atfalse
created_bytrue
created_by_avatar_urlfalse
created_by_namefalse
created_by_usernamefalse
default_ttltrue
deletedfalse
deprecatedtrue
descriptiontrue
disable_module_cachetrue
display_nametrue
failure_ttltrue
group_acltrue
icontrue
idtrue
max_port_sharing_leveltrue
nametrue
organization_display_namefalse
organization_iconfalse
organization_idfalse
organization_namefalse
provisionertrue
require_active_versiontrue
time_til_dormanttrue
time_til_dormant_autodeletetrue
updated_atfalse
use_classic_parameter_flowtrue
user_acltrue
| | TemplateVersion
create, write | |
FieldTracked
archivedtrue
created_atfalse
created_bytrue
created_by_avatar_urlfalse
created_by_namefalse
created_by_usernamefalse
external_auth_providersfalse
has_ai_taskfalse
has_external_agentfalse
idtrue
job_idfalse
messagefalse
nametrue
organization_idfalse
readmetrue
source_example_idfalse
template_idtrue
updated_atfalse
| -| User
create, write, delete | |
FieldTracked
avatar_urlfalse
created_atfalse
deletedtrue
emailtrue
github_com_user_idfalse
hashed_one_time_passcodefalse
hashed_passwordtrue
idtrue
is_service_accounttrue
is_systemtrue
last_seen_atfalse
login_typetrue
nametrue
one_time_passcode_expires_attrue
quiet_hours_scheduletrue
rbac_rolestrue
statustrue
updated_atfalse
usernametrue
| +| User
create, write, delete | |
FieldTracked
avatar_urlfalse
chat_spend_limit_microstrue
created_atfalse
deletedtrue
emailtrue
github_com_user_idfalse
hashed_one_time_passcodefalse
hashed_passwordtrue
idtrue
is_service_accounttrue
is_systemtrue
last_seen_atfalse
login_typetrue
nametrue
one_time_passcode_expires_attrue
quiet_hours_scheduletrue
rbac_rolestrue
statustrue
updated_atfalse
usernametrue
| | WorkspaceBuild
start, stop | |
FieldTracked
build_numberfalse
created_atfalse
daily_costfalse
deadlinefalse
has_ai_taskfalse
has_external_agentfalse
idfalse
initiator_by_avatar_urlfalse
initiator_by_namefalse
initiator_by_usernamefalse
initiator_idfalse
job_idfalse
max_deadlinefalse
reasonfalse
template_version_idtrue
template_version_preset_idfalse
transitionfalse
updated_atfalse
workspace_idfalse
| | WorkspaceProxy
| |
FieldTracked
created_attrue
deletedfalse
derp_enabledtrue
derp_onlytrue
display_nametrue
icontrue
idtrue
nametrue
region_idtrue
token_hashed_secrettrue
updated_atfalse
urltrue
versiontrue
wildcard_hostnametrue
| | WorkspaceTable
| |
FieldTracked
automatic_updatestrue
autostart_scheduletrue
created_atfalse
deletedfalse
deleting_attrue
dormant_attrue
favoritetrue
group_acltrue
idtrue
last_used_atfalse
nametrue
next_start_attrue
organization_idfalse
owner_idtrue
template_idtrue
ttltrue
updated_atfalse
user_acltrue
| diff --git a/enterprise/audit/table.go b/enterprise/audit/table.go index d556e47d44..9d7901d5ac 100644 --- a/enterprise/audit/table.go +++ b/enterprise/audit/table.go @@ -162,6 +162,7 @@ var auditableResourcesTypes = map[any]map[string]Action{ "one_time_passcode_expires_at": ActionTrack, "is_system": ActionTrack, // Should never change, but track it anyway. "is_service_account": ActionTrack, // Should never change, but track it anyway. + "chat_spend_limit_micros": ActionTrack, }, &database.WorkspaceTable{}: { "id": ActionTrack, @@ -205,14 +206,15 @@ var auditableResourcesTypes = map[any]map[string]Action{ "has_external_agent": ActionIgnore, // Never changes. }, &database.AuditableGroup{}: { - "id": ActionTrack, - "name": ActionTrack, - "display_name": ActionTrack, - "organization_id": ActionIgnore, // Never changes. - "avatar_url": ActionTrack, - "quota_allowance": ActionTrack, - "members": ActionTrack, - "source": ActionIgnore, + "id": ActionTrack, + "name": ActionTrack, + "display_name": ActionTrack, + "organization_id": ActionIgnore, // Never changes. + "avatar_url": ActionTrack, + "quota_allowance": ActionTrack, + "members": ActionTrack, + "source": ActionIgnore, + "chat_spend_limit_micros": ActionTrack, }, &database.APIKey{}: { "id": ActionIgnore, diff --git a/provisioner/terraform/testdata/resources/devcontainer-multiple-agents/converted_state.state.golden b/provisioner/terraform/testdata/resources/devcontainer-multiple-agents/converted_state.state.golden index ad4ca3b762..3f3144c17c 100644 --- a/provisioner/terraform/testdata/resources/devcontainer-multiple-agents/converted_state.state.golden +++ b/provisioner/terraform/testdata/resources/devcontainer-multiple-agents/converted_state.state.golden @@ -43,7 +43,7 @@ "workspace_folder": "/other", "name": "other", "id": "8e5a16da-e98c-4a6f-b24c-3c0cbd6bb9df", - "subagent_id": "cfa0a8fc-29cf-44e8-8ae2-42637057826a" + "subagent_id": "bffaad51-64f5-4da4-9a08-ffab24d04c7f" } ], "api_key_scope": "all" diff --git a/provisioner/terraform/testdata/resources/devcontainer-multiple-agents/devcontainer-multiple-agents.tfstate.json b/provisioner/terraform/testdata/resources/devcontainer-multiple-agents/devcontainer-multiple-agents.tfstate.json index 5b37b294fc..51dd3f843a 100644 --- a/provisioner/terraform/testdata/resources/devcontainer-multiple-agents/devcontainer-multiple-agents.tfstate.json +++ b/provisioner/terraform/testdata/resources/devcontainer-multiple-agents/devcontainer-multiple-agents.tfstate.json @@ -157,7 +157,7 @@ "agent_id": "37a1bd80-851e-48cf-bd36-af4aab414203", "config_path": null, "id": "8e5a16da-e98c-4a6f-b24c-3c0cbd6bb9df", - "subagent_id": "cfa0a8fc-29cf-44e8-8ae2-42637057826a", + "subagent_id": "bffaad51-64f5-4da4-9a08-ffab24d04c7f", "workspace_folder": "/other" }, "sensitive_values": {}, diff --git a/provisioner/terraform/testdata/resources/duplicate-env-keys/converted_state.plan.golden b/provisioner/terraform/testdata/resources/duplicate-env-keys/converted_state.plan.golden index 9e956c13d7..e95d6e977e 100644 --- a/provisioner/terraform/testdata/resources/duplicate-env-keys/converted_state.plan.golden +++ b/provisioner/terraform/testdata/resources/duplicate-env-keys/converted_state.plan.golden @@ -5,12 +5,11 @@ "type": "null_resource", "agents": [ { - "id": "aaaaaaaa-1111-2222-3333-444444444444", "name": "dev", "operating_system": "linux", "architecture": "amd64", "Auth": { - "Token": "11111111-2222-3333-4444-555555555555" + "Token": "" }, "connection_timeout_seconds": 120, "display_apps": { diff --git a/provisioner/terraform/testdata/resources/duplicate-env-keys/duplicate-env-keys.tfplan.json b/provisioner/terraform/testdata/resources/duplicate-env-keys/duplicate-env-keys.tfplan.json index f216d9258e..5e985716df 100644 --- a/provisioner/terraform/testdata/resources/duplicate-env-keys/duplicate-env-keys.tfplan.json +++ b/provisioner/terraform/testdata/resources/duplicate-env-keys/duplicate-env-keys.tfplan.json @@ -17,18 +17,7 @@ "auth": "token", "connection_timeout": 120, "dir": null, - "display_apps": [ - { - "port_forwarding_helper": true, - "ssh_helper": true, - "vscode": true, - "vscode_insiders": false, - "web_terminal": true - } - ], "env": null, - "id": "aaaaaaaa-1111-2222-3333-444444444444", - "init_script": "", "metadata": [], "motd_file": null, "order": null, @@ -37,13 +26,10 @@ "shutdown_script": null, "startup_script": null, "startup_script_behavior": "non-blocking", - "token": "11111111-2222-3333-4444-555555555555", "troubleshooting_url": null }, "sensitive_values": { - "display_apps": [ - {} - ], + "display_apps": [], "metadata": [], "resources_monitoring": [], "token": true @@ -57,15 +43,10 @@ "provider_name": "registry.terraform.io/coder/coder", "schema_version": 1, "values": { - "agent_id": "aaaaaaaa-1111-2222-3333-444444444444", - "id": "bbbbbbbb-1111-2222-3333-444444444444", "name": "PATH", "value": "/a/bin" }, - "sensitive_values": {}, - "depends_on": [ - "coder_agent.dev" - ] + "sensitive_values": {} }, { "address": "coder_env.path_b", @@ -75,15 +56,10 @@ "provider_name": "registry.terraform.io/coder/coder", "schema_version": 1, "values": { - "agent_id": "aaaaaaaa-1111-2222-3333-444444444444", - "id": "cccccccc-1111-2222-3333-444444444444", "name": "PATH", "value": "/b/bin" }, - "sensitive_values": {}, - "depends_on": [ - "coder_agent.dev" - ] + "sensitive_values": {} }, { "address": "coder_env.unique_env", @@ -93,15 +69,10 @@ "provider_name": "registry.terraform.io/coder/coder", "schema_version": 1, "values": { - "agent_id": "aaaaaaaa-1111-2222-3333-444444444444", - "id": "dddddddd-1111-2222-3333-444444444444", "name": "UNIQUE", "value": "unique_value" }, - "sensitive_values": {}, - "depends_on": [ - "coder_agent.dev" - ] + "sensitive_values": {} }, { "address": "null_resource.dev", @@ -111,15 +82,270 @@ "provider_name": "registry.terraform.io/hashicorp/null", "schema_version": 0, "values": { - "id": "1234567890123456789", "triggers": null }, - "sensitive_values": {}, + "sensitive_values": {} + } + ] + } + }, + "resource_changes": [ + { + "address": "coder_agent.dev", + "mode": "managed", + "type": "coder_agent", + "name": "dev", + "provider_name": "registry.terraform.io/coder/coder", + "change": { + "actions": [ + "create" + ], + "before": null, + "after": { + "api_key_scope": "all", + "arch": "amd64", + "auth": "token", + "connection_timeout": 120, + "dir": null, + "env": null, + "metadata": [], + "motd_file": null, + "order": null, + "os": "linux", + "resources_monitoring": [], + "shutdown_script": null, + "startup_script": null, + "startup_script_behavior": "non-blocking", + "troubleshooting_url": null + }, + "after_unknown": { + "display_apps": true, + "id": true, + "init_script": true, + "metadata": [], + "resources_monitoring": [], + "token": true + }, + "before_sensitive": false, + "after_sensitive": { + "display_apps": [], + "metadata": [], + "resources_monitoring": [], + "token": true + } + } + }, + { + "address": "coder_env.path_a", + "mode": "managed", + "type": "coder_env", + "name": "path_a", + "provider_name": "registry.terraform.io/coder/coder", + "change": { + "actions": [ + "create" + ], + "before": null, + "after": { + "name": "PATH", + "value": "/a/bin" + }, + "after_unknown": { + "agent_id": true, + "id": true + }, + "before_sensitive": false, + "after_sensitive": {} + } + }, + { + "address": "coder_env.path_b", + "mode": "managed", + "type": "coder_env", + "name": "path_b", + "provider_name": "registry.terraform.io/coder/coder", + "change": { + "actions": [ + "create" + ], + "before": null, + "after": { + "name": "PATH", + "value": "/b/bin" + }, + "after_unknown": { + "agent_id": true, + "id": true + }, + "before_sensitive": false, + "after_sensitive": {} + } + }, + { + "address": "coder_env.unique_env", + "mode": "managed", + "type": "coder_env", + "name": "unique_env", + "provider_name": "registry.terraform.io/coder/coder", + "change": { + "actions": [ + "create" + ], + "before": null, + "after": { + "name": "UNIQUE", + "value": "unique_value" + }, + "after_unknown": { + "agent_id": true, + "id": true + }, + "before_sensitive": false, + "after_sensitive": {} + } + }, + { + "address": "null_resource.dev", + "mode": "managed", + "type": "null_resource", + "name": "dev", + "provider_name": "registry.terraform.io/hashicorp/null", + "change": { + "actions": [ + "create" + ], + "before": null, + "after": { + "triggers": null + }, + "after_unknown": { + "id": true + }, + "before_sensitive": false, + "after_sensitive": {} + } + } + ], + "configuration": { + "provider_config": { + "coder": { + "name": "coder", + "full_name": "registry.terraform.io/coder/coder", + "version_constraint": ">= 2.0.0" + }, + "null": { + "name": "null", + "full_name": "registry.terraform.io/hashicorp/null" + } + }, + "root_module": { + "resources": [ + { + "address": "coder_agent.dev", + "mode": "managed", + "type": "coder_agent", + "name": "dev", + "provider_config_key": "coder", + "expressions": { + "arch": { + "constant_value": "amd64" + }, + "os": { + "constant_value": "linux" + } + }, + "schema_version": 1 + }, + { + "address": "coder_env.path_a", + "mode": "managed", + "type": "coder_env", + "name": "path_a", + "provider_config_key": "coder", + "expressions": { + "agent_id": { + "references": [ + "coder_agent.dev.id", + "coder_agent.dev" + ] + }, + "name": { + "constant_value": "PATH" + }, + "value": { + "constant_value": "/a/bin" + } + }, + "schema_version": 1 + }, + { + "address": "coder_env.path_b", + "mode": "managed", + "type": "coder_env", + "name": "path_b", + "provider_config_key": "coder", + "expressions": { + "agent_id": { + "references": [ + "coder_agent.dev.id", + "coder_agent.dev" + ] + }, + "name": { + "constant_value": "PATH" + }, + "value": { + "constant_value": "/b/bin" + } + }, + "schema_version": 1 + }, + { + "address": "coder_env.unique_env", + "mode": "managed", + "type": "coder_env", + "name": "unique_env", + "provider_config_key": "coder", + "expressions": { + "agent_id": { + "references": [ + "coder_agent.dev.id", + "coder_agent.dev" + ] + }, + "name": { + "constant_value": "UNIQUE" + }, + "value": { + "constant_value": "unique_value" + } + }, + "schema_version": 1 + }, + { + "address": "null_resource.dev", + "mode": "managed", + "type": "null_resource", + "name": "dev", + "provider_config_key": "null", + "schema_version": 0, "depends_on": [ "coder_agent.dev" ] } ] } - } + }, + "relevant_attributes": [ + { + "resource": "coder_agent.dev", + "attribute": [ + "id" + ] + } + ], + "timestamp": "2026-03-16T15:54:16Z", + "applyable": true, + "complete": true, + "errored": false } diff --git a/site/src/api/typesGenerated.ts b/site/src/api/typesGenerated.ts index f12a8cc36e..047a7ff833 100644 --- a/site/src/api/typesGenerated.ts +++ b/site/src/api/typesGenerated.ts @@ -1119,6 +1119,7 @@ export interface ChatCostSummary { readonly total_cache_creation_tokens: number; readonly by_model: readonly ChatCostModelBreakdown[]; readonly by_chat: readonly ChatCostChatBreakdown[]; + readonly usage_limit?: ChatUsageLimitStatus; } // From codersdk/chats.go @@ -1782,6 +1783,100 @@ export interface ChatSystemPromptResponse { readonly system_prompt: string; } +// From codersdk/chats.go +/** + * ChatUsageLimitConfig is the deployment-wide default usage limit config. + */ +export interface ChatUsageLimitConfig { + /** + * Nil in the API means no default limit is set. The DB stores 0 when + * limiting is disabled. + */ + readonly spend_limit_micros: number | null; + readonly period: ChatUsageLimitPeriod; + readonly updated_at: string; +} + +// From codersdk/chats.go +/** + * ChatUsageLimitConfigResponse is returned from the admin config endpoint + * and includes the config plus a count of models without pricing. + */ +export interface ChatUsageLimitConfigResponse extends ChatUsageLimitConfig { + readonly unpriced_model_count: number; + readonly overrides: readonly ChatUsageLimitOverride[]; + readonly group_overrides: readonly ChatUsageLimitGroupOverride[]; +} + +// From codersdk/chats.go +/** + * ChatUsageLimitExceededResponse is the 409 response body returned when a + * chat operation exceeds the caller's usage limit. The structured fields let + * frontends render user-friendly spend, limit, and reset information without + * parsing debug text. + */ +export interface ChatUsageLimitExceededResponse extends Response { + readonly spent_micros: number; + readonly limit_micros: number; + readonly resets_at: string; +} + +// From codersdk/chats.go +/** + * ChatUsageLimitGroupOverride represents a group-scoped spend limit override. + */ +export interface ChatUsageLimitGroupOverride { + readonly group_id: string; + readonly group_name: string; + readonly group_display_name: string; + readonly group_avatar_url: string; + readonly member_count: number; + /** + * Nil in the API means no group override is set. Persisted override rows + * store positive values. + */ + readonly spend_limit_micros: number | null; +} + +// From codersdk/chats.go +/** + * ChatUsageLimitOverride is a per-user override of the deployment default. + */ +export interface ChatUsageLimitOverride { + readonly user_id: string; + readonly username: string; + readonly name: string; + readonly avatar_url: string; + /** + * Nil in the API means no user override is set. Persisted override rows + * store positive values. + */ + readonly spend_limit_micros: number | null; +} + +// From codersdk/chats.go +export type ChatUsageLimitPeriod = "day" | "month" | "week"; + +export const ChatUsageLimitPeriods: ChatUsageLimitPeriod[] = [ + "day", + "month", + "week", +]; + +// From codersdk/chats.go +/** + * ChatUsageLimitStatus represents the current spend status for a user + * within their active limit period. + */ +export interface ChatUsageLimitStatus { + readonly is_limited: boolean; + readonly period?: ChatUsageLimitPeriod; + readonly spend_limit_micros?: number; + readonly current_spend: number; + readonly period_start?: string; + readonly period_end?: string; +} + // From codersdk/client.go /** * CoderDesktopTelemetryHeader contains a JSON-encoded representation of Desktop telemetry @@ -6524,6 +6619,22 @@ export interface UpdateChatSystemPromptRequest { readonly system_prompt: string; } +// From codersdk/chats.go +/** + * UpdateChatUsageLimitGroupOverrideRequest is kept as a compatibility alias. + */ +export interface UpdateChatUsageLimitGroupOverrideRequest { + readonly spend_limit_micros: number; // Must be greater than 0. +} + +// From codersdk/chats.go +/** + * UpdateChatUsageLimitOverrideRequest is kept as a compatibility alias. + */ +export interface UpdateChatUsageLimitOverrideRequest { + readonly spend_limit_micros: number; // Must be greater than 0. +} + // From codersdk/updatecheck.go /** * UpdateCheckResponse contains information on the latest release of Coder. @@ -6832,6 +6943,24 @@ export interface UploadResponse { readonly hash: string; } +// From codersdk/chats.go +/** + * UpsertChatUsageLimitGroupOverrideRequest is the request to create or update + * a group-level spend limit override. + */ +export interface UpsertChatUsageLimitGroupOverrideRequest { + readonly spend_limit_micros: number; // Must be greater than 0. +} + +// From codersdk/chats.go +/** + * UpsertChatUsageLimitOverrideRequest is the body for creating/updating a + * per-user usage limit override. + */ +export interface UpsertChatUsageLimitOverrideRequest { + readonly spend_limit_micros: number; // Must be greater than 0. +} + // From codersdk/workspaceagentportshare.go export interface UpsertWorkspaceAgentPortShareRequest { readonly agent_name: string;