feat(coderd): enforce ai budget on pre-request path (#26915)

## Description

Adds pre-request AI budget enforcement to `aibridged`. Requests are rejected with HTTP 403 when the user's aggregated spend for the current period has reached their effective limit.

## Changes

- Add `IsBudgetExceeded` RPC to `aibridgedserver`. Resolves the user's effective budget, aggregates spend over the caller-supplied `[period_start, now]` window, and returns whether the limit has been reached along with the effective limit.
- Wire the check into `aibridged`'s HTTP handler. The caller computes the period start (monthly for now) and passes it in the request.
- Reject exceeded requests with HTTP 403 Forbidden and a message directing the user to contact an administrator.
- Add `dbtime.StartOfMonth` alongside `StartOfDay` for period computation.
- Add real-DB tests covering the enforcement path: month-boundary excludes prior-period spend, and a new user override unblocks a previously-exceeded user.

Closes https://linear.app/codercom/issue/AIGOV-428/add-pre-request-budget-enforcement

> [!NOTE]
> Initially generated by Claude Opus 4.7, modified and reviewed by @ssncferreira
This commit is contained in:
Susana Ferreira
2026-07-02 16:53:36 +01:00
committed by GitHub
parent be9c95c8f5
commit 1989db0e2b
9 changed files with 885 additions and 181 deletions
+89
View File
@@ -8,6 +8,7 @@ import (
"slices"
"strings"
"sync"
"time"
"github.com/google/uuid"
"github.com/hashicorp/go-multierror"
@@ -16,6 +17,7 @@ import (
"google.golang.org/protobuf/types/known/structpb"
"cdr.dev/slog/v3"
"github.com/coder/coder/v2/coderd/aibridge/budget"
"github.com/coder/coder/v2/coderd/aibridged"
"github.com/coder/coder/v2/coderd/aibridged/proto"
"github.com/coder/coder/v2/coderd/aiseats"
@@ -75,6 +77,7 @@ type store interface {
GetAIModelPriceByProviderModel(ctx context.Context, arg database.GetAIModelPriceByProviderModelParams) (database.AIModelPrice, error)
GetUserAIBudgetOverride(ctx context.Context, userID uuid.UUID) (database.UserAIBudgetOverride, error)
GetHighestGroupAIBudgetByUser(ctx context.Context, userID uuid.UUID) (database.GetHighestGroupAIBudgetByUserRow, error)
GetUserAISpendSince(ctx context.Context, arg database.GetUserAISpendSinceParams) (database.GetUserAISpendSinceRow, error)
// MCPConfigurator-related queries.
GetExternalAuthLinksByUserID(ctx context.Context, userID uuid.UUID) ([]database.ExternalAuthLink, error)
@@ -727,6 +730,92 @@ func (s *Server) IsAuthorized(ctx context.Context, in *proto.IsAuthorizedRequest
}, nil
}
// IsBudgetExceeded reports whether the user's AI spend has reached their
// effective limit over [PeriodStart, now].
func (s *Server) IsBudgetExceeded(ctx context.Context, in *proto.IsBudgetExceededRequest) (*proto.IsBudgetExceededResponse, error) {
//nolint:gocritic // AIBridged has specific authz rules.
ctx = dbauthz.AsAIBridged(ctx)
userID, err := uuid.Parse(in.GetUserId())
if err != nil {
return nil, xerrors.Errorf("invalid user_id %q: %w", in.GetUserId(), err)
}
// An unset PeriodStart deserializes to time.Unix(0, 0), which would
// incorrectly aggregate the user's lifetime spend against a period budget.
if in.PeriodStart == nil {
return nil, xerrors.New("period_start is required")
}
periodStart := in.GetPeriodStart().AsTime()
userBudget, err := s.checkUserAIBudget(ctx, userID, periodStart)
if err != nil {
return nil, err
}
return &proto.IsBudgetExceededResponse{
Exceeded: userBudget.Exceeded,
SpendLimitMicros: userBudget.SpendLimitMicros,
}, nil
}
// userAIBudget is a snapshot of a user's AI budget status. SpendLimitMicros
// is nil when no budget is configured for the user (unlimited).
type userAIBudget struct {
Exceeded bool
SpendLimitMicros *int64
}
// checkUserAIBudget evaluates the user's AI budget status aggregated over
// [periodStart, now].
//
// Note: there is a potential race condition where two concurrent requests
// from the same user can both pass the 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 request cost × concurrency; once the accumulated
// spend crosses the limit, subsequent requests are blocked.
// - Cost accounting is advisory, not strict. The goal is to prevent
// overages, not build an accounting system.
// - Fail-open is acceptable for this case.
func (s *Server) checkUserAIBudget(ctx context.Context, userID uuid.UUID, periodStart time.Time) (userAIBudget, error) {
effectiveBudget, ok, err := budget.ResolveUserAIBudget(ctx, s.store, userID, s.budgetPolicy)
if err != nil {
return userAIBudget{}, xerrors.Errorf("resolve effective AI budget for user %q with budget policy %q: %w", userID, s.budgetPolicy, err)
}
if !ok {
// No budget configured for the user; return zero-valued status.
return userAIBudget{}, nil
}
spend, err := s.store.GetUserAISpendSince(ctx, database.GetUserAISpendSinceParams{
UserID: userID,
EffectiveGroupID: effectiveBudget.GroupID,
PeriodStart: periodStart,
})
if err != nil {
return userAIBudget{}, xerrors.Errorf("get user AI spend for user %q in group %q: %w", userID, effectiveBudget.GroupID, err)
}
exceeded := spend.SpendMicros >= effectiveBudget.SpendLimitMicros
logger := s.logger.With(
slog.F("user_id", userID),
slog.F("effective_group_id", effectiveBudget.GroupID),
slog.F("period_start", periodStart),
slog.F("current_spend_micros", spend.SpendMicros),
slog.F("spend_limit_micros", effectiveBudget.SpendLimitMicros),
slog.F("exceeded", exceeded),
)
logger.Debug(ctx, "user AI spend status")
if exceeded {
logger.Warn(ctx, "user AI budget exceeded")
}
return userAIBudget{
Exceeded: exceeded,
SpendLimitMicros: ptr.Ref(effectiveBudget.SpendLimitMicros),
}, nil
}
// GetAIProviders returns the full AI provider set (enabled and disabled) from
// the database, which is the single source of truth seeded from coderd's
// environment. Embedded and standalone AI Gateway daemons call this over DRPC
@@ -391,6 +391,314 @@ func TestAuthorization_Delegated(t *testing.T) {
}
}
func TestIsBudgetExceeded(t *testing.T) {
t.Parallel()
cases := []struct {
name string
userIDStr string
omitPeriodStart bool
setupMocks func(db *dbmock.MockStore, userID uuid.UUID) *proto.IsBudgetExceededResponse
wantErrContains string
}{
{
// Invalid UUID short-circuits before any store call.
name: "invalid user_id",
userIDStr: "not-a-uuid",
wantErrContains: "invalid user_id",
},
{
// Missing period_start is rejected before any store call.
name: "missing period_start",
omitPeriodStart: true,
wantErrContains: "period_start is required",
},
{
// No override and no group budget resolves: pass-through.
name: "no budget configured returns not exceeded",
setupMocks: func(db *dbmock.MockStore, userID uuid.UUID) *proto.IsBudgetExceededResponse {
db.EXPECT().GetUserAIBudgetOverride(gomock.Any(), userID).
Return(database.UserAIBudgetOverride{}, sql.ErrNoRows)
db.EXPECT().GetHighestGroupAIBudgetByUser(gomock.Any(), userID).
Return(database.GetHighestGroupAIBudgetByUserRow{}, sql.ErrNoRows)
return &proto.IsBudgetExceededResponse{
Exceeded: false,
SpendLimitMicros: nil,
}
},
},
{
// Group budget resolves, spend below limit (spend 500 < limit 1000): pass-through.
name: "under limit returns not exceeded",
setupMocks: func(db *dbmock.MockStore, userID uuid.UUID) *proto.IsBudgetExceededResponse {
groupID := uuid.New()
db.EXPECT().GetUserAIBudgetOverride(gomock.Any(), userID).
Return(database.UserAIBudgetOverride{}, sql.ErrNoRows)
db.EXPECT().GetHighestGroupAIBudgetByUser(gomock.Any(), userID).
Return(database.GetHighestGroupAIBudgetByUserRow{GroupID: groupID, SpendLimitMicros: 1_000}, nil)
db.EXPECT().GetUserAISpendSince(gomock.Any(), gomock.Any()).
Return(database.GetUserAISpendSinceRow{SpendMicros: 500}, nil)
return &proto.IsBudgetExceededResponse{
Exceeded: false,
SpendLimitMicros: ptr.Ref(int64(1_000)),
}
},
},
{
// Group budget resolves, spend at limit (spend 1000 == limit 1000): blocked.
name: "at limit returns exceeded",
setupMocks: func(db *dbmock.MockStore, userID uuid.UUID) *proto.IsBudgetExceededResponse {
groupID := uuid.New()
db.EXPECT().GetUserAIBudgetOverride(gomock.Any(), userID).
Return(database.UserAIBudgetOverride{}, sql.ErrNoRows)
db.EXPECT().GetHighestGroupAIBudgetByUser(gomock.Any(), userID).
Return(database.GetHighestGroupAIBudgetByUserRow{GroupID: groupID, SpendLimitMicros: 1_000}, nil)
db.EXPECT().GetUserAISpendSince(gomock.Any(), gomock.Any()).
Return(database.GetUserAISpendSinceRow{SpendMicros: 1_000}, nil)
return &proto.IsBudgetExceededResponse{
Exceeded: true,
SpendLimitMicros: ptr.Ref(int64(1_000)),
}
},
},
{
// Limit of 0 is a valid "block-all" setting, distinct from
// "no budget configured": blocked.
name: "zero limit blocks all requests",
setupMocks: func(db *dbmock.MockStore, userID uuid.UUID) *proto.IsBudgetExceededResponse {
groupID := uuid.New()
db.EXPECT().GetUserAIBudgetOverride(gomock.Any(), userID).
Return(database.UserAIBudgetOverride{}, sql.ErrNoRows)
db.EXPECT().GetHighestGroupAIBudgetByUser(gomock.Any(), userID).
Return(database.GetHighestGroupAIBudgetByUserRow{GroupID: groupID, SpendLimitMicros: 0}, nil)
db.EXPECT().GetUserAISpendSince(gomock.Any(), gomock.Any()).
Return(database.GetUserAISpendSinceRow{SpendMicros: 0}, nil)
return &proto.IsBudgetExceededResponse{
Exceeded: true,
SpendLimitMicros: ptr.Ref(int64(0)),
}
},
},
{
// Group budget resolves, spend above limit (spend 1500 > limit 1000): blocked.
name: "over limit returns exceeded",
setupMocks: func(db *dbmock.MockStore, userID uuid.UUID) *proto.IsBudgetExceededResponse {
groupID := uuid.New()
db.EXPECT().GetUserAIBudgetOverride(gomock.Any(), userID).
Return(database.UserAIBudgetOverride{}, sql.ErrNoRows)
db.EXPECT().GetHighestGroupAIBudgetByUser(gomock.Any(), userID).
Return(database.GetHighestGroupAIBudgetByUserRow{GroupID: groupID, SpendLimitMicros: 1_000}, nil)
db.EXPECT().GetUserAISpendSince(gomock.Any(), gomock.Any()).
Return(database.GetUserAISpendSinceRow{SpendMicros: 1_500}, nil)
return &proto.IsBudgetExceededResponse{
Exceeded: true,
SpendLimitMicros: ptr.Ref(int64(1_000)),
}
},
},
{
// User override wins, group lookup skipped, spend aggregated against
// the override's group (spend 600 > limit 500): blocked.
name: "user override wins over group budget",
setupMocks: func(db *dbmock.MockStore, userID uuid.UUID) *proto.IsBudgetExceededResponse {
overrideGroupID := uuid.New()
db.EXPECT().GetUserAIBudgetOverride(gomock.Any(), userID).
Return(database.UserAIBudgetOverride{
UserID: userID,
GroupID: overrideGroupID,
SpendLimitMicros: 500,
}, nil)
db.EXPECT().GetUserAISpendSince(gomock.Any(), gomock.Cond(func(p database.GetUserAISpendSinceParams) bool {
return assert.Equal(t, overrideGroupID, p.EffectiveGroupID, "spend aggregated against override group")
})).Return(database.GetUserAISpendSinceRow{SpendMicros: 600}, nil)
return &proto.IsBudgetExceededResponse{
Exceeded: true,
SpendLimitMicros: ptr.Ref(int64(500)),
}
},
},
{
// Unexpected error from budget override lookup propagates.
name: "budget resolution error propagates",
setupMocks: func(db *dbmock.MockStore, userID uuid.UUID) *proto.IsBudgetExceededResponse {
db.EXPECT().GetUserAIBudgetOverride(gomock.Any(), userID).
Return(database.UserAIBudgetOverride{}, sql.ErrConnDone)
return nil
},
wantErrContains: "resolve effective AI budget",
},
{
// Error from spend aggregation propagates (fail-closed).
name: "spend aggregation error propagates",
setupMocks: func(db *dbmock.MockStore, userID uuid.UUID) *proto.IsBudgetExceededResponse {
db.EXPECT().GetUserAIBudgetOverride(gomock.Any(), userID).
Return(database.UserAIBudgetOverride{}, sql.ErrNoRows)
db.EXPECT().GetHighestGroupAIBudgetByUser(gomock.Any(), userID).
Return(database.GetHighestGroupAIBudgetByUserRow{GroupID: uuid.New(), SpendLimitMicros: 1_000}, nil)
db.EXPECT().GetUserAISpendSince(gomock.Any(), gomock.Any()).
Return(database.GetUserAISpendSinceRow{}, sql.ErrConnDone)
return nil
},
wantErrContains: "get user AI spend",
},
}
for _, tc := range cases {
t.Run(tc.name, func(t *testing.T) {
t.Parallel()
ctrl := gomock.NewController(t)
db := dbmock.NewMockStore(ctrl)
logger := testutil.Logger(t)
userID := uuid.New()
userIDStr := tc.userIDStr
if userIDStr == "" {
userIDStr = userID.String()
}
var wantResp *proto.IsBudgetExceededResponse
if tc.setupMocks != nil {
wantResp = tc.setupMocks(db, userID)
}
srv, err := aibridgedserver.NewServer(t.Context(), db, logger, "/", codersdk.AIBridgeConfig{}, nil, requiredExperiments, agplaiseats.Noop{})
require.NoError(t, err)
req := &proto.IsBudgetExceededRequest{UserId: userIDStr}
if !tc.omitPeriodStart {
req.PeriodStart = timestamppb.New(dbtime.StartOfMonth(dbtime.Now().UTC()))
}
resp, err := srv.IsBudgetExceeded(t.Context(), req)
if tc.wantErrContains != "" {
require.Error(t, err)
require.Nil(t, resp)
assert.ErrorContains(t, err, tc.wantErrContains)
return
}
require.NoError(t, err)
require.NotNil(t, resp)
require.Equal(t, wantResp.GetExceeded(), resp.GetExceeded(), "exceeded")
require.Equal(t, wantResp.SpendLimitMicros, resp.SpendLimitMicros, "spend_limit_micros")
})
}
}
// TestIsBudgetExceeded_Enforcement exercises real-DB scenarios that drive
// enforcement decisions.
func TestIsBudgetExceeded_Enforcement(t *testing.T) {
t.Parallel()
const groupLimitMicros = 1_000_000
// setup provisions a user in an organization with a single budgeted group.
setup := func(t *testing.T) (context.Context, database.Store, *aibridgedserver.Server, database.User, database.Group) {
t.Helper()
ctx := testutil.Context(t, testutil.WaitLong)
logger := testutil.Logger(t)
rawDB, _ := dbtestutil.NewDB(t)
authzDB := dbauthz.New(rawDB, rbac.NewStrictAuthorizer(prometheus.NewRegistry()), logger, coderdtest.AccessControlStorePointer())
org := dbgen.Organization(t, rawDB, database.Organization{})
user := dbgen.User(t, rawDB, database.User{})
dbgen.OrganizationMember(t, rawDB, database.OrganizationMember{OrganizationID: org.ID, UserID: user.ID})
group := dbgen.Group(t, rawDB, database.Group{OrganizationID: org.ID})
dbgen.GroupMember(t, rawDB, database.GroupMemberTable{UserID: user.ID, GroupID: group.ID})
_, err := rawDB.UpsertGroupAIBudget(ctx, database.UpsertGroupAIBudgetParams{
GroupID: group.ID,
SpendLimitMicros: groupLimitMicros,
})
require.NoError(t, err)
srv, err := aibridgedserver.NewServer(ctx, authzDB, logger, "/", codersdk.AIBridgeConfig{}, nil, requiredExperiments, agplaiseats.Noop{})
require.NoError(t, err)
return ctx, rawDB, srv, user, group
}
t.Run("period boundary excludes prior period spend", func(t *testing.T) {
t.Parallel()
ctx, rawDB, srv, user, group := setup(t)
prevMonth := time.Date(2026, time.January, 1, 0, 0, 0, 0, time.UTC)
newMonth := time.Date(2026, time.February, 1, 0, 0, 0, 0, time.UTC)
// User spend on 2026-01-15.
_, err := rawDB.IncrementUserAIDailySpend(ctx, database.IncrementUserAIDailySpendParams{
UserID: user.ID,
EffectiveGroupID: group.ID,
Day: prevMonth.AddDate(0, 0, 14),
CostMicros: 1_500_000,
})
require.NoError(t, err)
// Query with period_start 2026-01-01: includes the 2026-01-15 spend, user exceeded.
prevMonthResp, err := srv.IsBudgetExceeded(ctx, &proto.IsBudgetExceededRequest{
UserId: user.ID.String(),
PeriodStart: timestamppb.New(prevMonth),
})
require.NoError(t, err)
require.True(t, prevMonthResp.GetExceeded())
require.Equal(t, int64(groupLimitMicros), prevMonthResp.GetSpendLimitMicros())
// Query with period_start 2026-02-01: excludes the 2026-01-15 spend, user not exceeded.
newMonthResp, err := srv.IsBudgetExceeded(ctx, &proto.IsBudgetExceededRequest{
UserId: user.ID.String(),
PeriodStart: timestamppb.New(newMonth),
})
require.NoError(t, err)
require.False(t, newMonthResp.GetExceeded())
require.Equal(t, int64(groupLimitMicros), newMonthResp.GetSpendLimitMicros())
})
t.Run("new user override unblocks user", func(t *testing.T) {
t.Parallel()
ctx, rawDB, srv, user, group := setup(t)
// Use fixed dates to keep the test deterministic.
periodStart := time.Date(2026, time.March, 1, 0, 0, 0, 0, time.UTC)
// User spend on 2026-03-15.
_, err := rawDB.IncrementUserAIDailySpend(ctx, database.IncrementUserAIDailySpendParams{
UserID: user.ID,
EffectiveGroupID: group.ID,
Day: periodStart.AddDate(0, 0, 14),
CostMicros: 1_500_000,
})
require.NoError(t, err)
// User's spend exceeds the group limit.
beforeResp, err := srv.IsBudgetExceeded(ctx, &proto.IsBudgetExceededRequest{
UserId: user.ID.String(),
PeriodStart: timestamppb.New(periodStart),
})
require.NoError(t, err)
require.True(t, beforeResp.GetExceeded())
require.Equal(t, int64(groupLimitMicros), beforeResp.GetSpendLimitMicros())
// Add user override with a higher limit on the same group. The override
// wins, so the user's spend is now under the effective limit.
const overrideLimitMicros = 2_000_000
_, err = rawDB.UpsertUserAIBudgetOverride(ctx, database.UpsertUserAIBudgetOverrideParams{
UserID: user.ID,
GroupID: group.ID,
SpendLimitMicros: overrideLimitMicros,
})
require.NoError(t, err)
afterResp, err := srv.IsBudgetExceeded(ctx, &proto.IsBudgetExceededRequest{
UserId: user.ID.String(),
PeriodStart: timestamppb.New(periodStart),
})
require.NoError(t, err)
require.False(t, afterResp.GetExceeded())
require.Equal(t, int64(overrideLimitMicros), afterResp.GetSpendLimitMicros())
})
}
func TestGetMCPServerConfigs(t *testing.T) {
t.Parallel()