mirror of
https://github.com/coder/coder.git
synced 2026-09-24 15:04:27 +08:00
feat: add Prometheus metrics for AI Governance cost control (#27490)
## Description Adds Prometheus metrics for AI budget cost control, emitted by the aibridged server under the `cost_control` subsystem (full names are prefixed `coder_ai_gateway_`). - `blocked_requests_total` (counter) — labels: `group_id` - `blocked_users` (gauge) — labels: `group_id` - `unpriced_requests_total` (counter) — labels: `provider`, `model` - `enforcement_duration_seconds` (histogram) — labels: `outcome` ## Changes - Add `GetOverBudgetUsersPerGroup` query (plus dbauthz/dbmetrics/dbmock wiring) to count over-budget users per effective group. - Add a background collector that refreshes the `blocked_users` gauge on an interval, started only when Prometheus is enabled. - Wire `Metrics` through the aibridged server, coderd API, `cli/server.go`, and the enterprise AI gateway handler; recording is nil-safe when metrics are unset. Closes https://linear.app/codercom/issue/AIGOV-296/add-prometheus-metrics-for-cost-control > [!NOTE] > Initially generated by Claude Opus 4.7, modified and reviewed by @ssncferreira
This commit is contained in:
@@ -76,6 +76,7 @@ func (api *API) CreateInMemoryAIBridgeServer(dialCtx context.Context) (client ai
|
||||
Experiments: api.Experiments,
|
||||
Logger: api.Logger.Named("aibridgedserver"),
|
||||
Clock: api.Clock,
|
||||
Metrics: api.AIGatewayServerMetrics,
|
||||
})
|
||||
if err != nil {
|
||||
return nil, err
|
||||
|
||||
@@ -124,6 +124,8 @@ type Server struct {
|
||||
budgetPeriod codersdk.AIBudgetPeriod
|
||||
clock quartz.Clock
|
||||
notifEnqueuer notifications.Enqueuer
|
||||
// metrics records cost-control metrics. May be nil.
|
||||
metrics *Metrics
|
||||
}
|
||||
|
||||
// Options carries the dependencies required to construct an aibridged Server.
|
||||
@@ -140,8 +142,9 @@ type Options struct {
|
||||
ExternalAuthConfigs []*externalauth.Config
|
||||
Experiments codersdk.Experiments
|
||||
|
||||
Logger slog.Logger
|
||||
Clock quartz.Clock
|
||||
Logger slog.Logger
|
||||
Clock quartz.Clock
|
||||
Metrics *Metrics
|
||||
}
|
||||
|
||||
func NewServer(lifecycleCtx context.Context, opts Options) (*Server, error) {
|
||||
@@ -173,6 +176,7 @@ func NewServer(lifecycleCtx context.Context, opts Options) (*Server, error) {
|
||||
budgetPeriod: codersdk.NewAIBudgetPeriodFromString(opts.GatewayCfg.BudgetPeriod),
|
||||
clock: opts.Clock,
|
||||
notifEnqueuer: enqueuer,
|
||||
metrics: opts.Metrics,
|
||||
}
|
||||
|
||||
if opts.GatewayCfg.InjectCoderMCPTools {
|
||||
@@ -811,10 +815,25 @@ func (s *Server) IsAuthorized(ctx context.Context, in *proto.IsAuthorizedRequest
|
||||
// IsBudgetExceeded reports whether the user's AI spend has reached their
|
||||
// effective limit over [periodStart, now], where periodStart is the start of
|
||||
// the current deployment-configured budget period.
|
||||
func (s *Server) IsBudgetExceeded(ctx context.Context, in *proto.IsBudgetExceededRequest) (*proto.IsBudgetExceededResponse, error) {
|
||||
func (s *Server) IsBudgetExceeded(ctx context.Context, in *proto.IsBudgetExceededRequest) (resp *proto.IsBudgetExceededResponse, retErr error) {
|
||||
//nolint:gocritic // AIBridged has specific authz rules.
|
||||
ctx = dbauthz.AsAIBridged(ctx)
|
||||
|
||||
start := s.clock.Now()
|
||||
defer func() {
|
||||
if s.metrics == nil {
|
||||
return
|
||||
}
|
||||
outcome := "allowed"
|
||||
switch {
|
||||
case retErr != nil:
|
||||
outcome = "error"
|
||||
case resp != nil && resp.Exceeded:
|
||||
outcome = "blocked"
|
||||
}
|
||||
s.metrics.EnforcementDuration.WithLabelValues(outcome).Observe(s.clock.Since(start).Seconds())
|
||||
}()
|
||||
|
||||
userID, err := uuid.Parse(in.GetUserId())
|
||||
if err != nil {
|
||||
return nil, xerrors.Errorf("invalid user_id %q: %w", in.GetUserId(), err)
|
||||
@@ -829,6 +848,9 @@ func (s *Server) IsBudgetExceeded(ctx context.Context, in *proto.IsBudgetExceede
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if userBudget.Exceeded && s.metrics != nil {
|
||||
s.metrics.BlockedRequests.WithLabelValues(userBudget.GroupID.String()).Inc()
|
||||
}
|
||||
return &proto.IsBudgetExceededResponse{
|
||||
Exceeded: userBudget.Exceeded,
|
||||
SpendLimitMicros: userBudget.SpendLimitMicros,
|
||||
@@ -836,10 +858,12 @@ func (s *Server) IsBudgetExceeded(ctx context.Context, in *proto.IsBudgetExceede
|
||||
}
|
||||
|
||||
// userAIBudget is a snapshot of a user's AI budget status. SpendLimitMicros
|
||||
// is nil when no budget is configured for the user (unlimited).
|
||||
// is nil when no budget is configured for the user (unlimited). GroupID is the
|
||||
// effective group the limit resolved to, set only when a limit applies.
|
||||
type userAIBudget struct {
|
||||
Exceeded bool
|
||||
SpendLimitMicros *int64
|
||||
GroupID uuid.UUID
|
||||
}
|
||||
|
||||
// checkUserAIBudget evaluates the user's AI budget status aggregated over
|
||||
@@ -894,6 +918,7 @@ func (s *Server) checkUserAIBudget(ctx context.Context, userID uuid.UUID, period
|
||||
return userAIBudget{
|
||||
Exceeded: exceeded,
|
||||
SpendLimitMicros: ptr.Ref(effectiveGroup.Limit.SpendLimitMicros),
|
||||
GroupID: effectiveGroup.GroupID,
|
||||
}, nil
|
||||
}
|
||||
|
||||
|
||||
@@ -16,6 +16,7 @@ import (
|
||||
|
||||
"github.com/google/uuid"
|
||||
"github.com/prometheus/client_golang/prometheus"
|
||||
promtest "github.com/prometheus/client_golang/prometheus/testutil"
|
||||
"github.com/sqlc-dev/pqtype"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
@@ -36,6 +37,7 @@ import (
|
||||
agplaiseats "github.com/coder/coder/v2/coderd/aiseats"
|
||||
"github.com/coder/coder/v2/coderd/apikey"
|
||||
"github.com/coder/coder/v2/coderd/coderdtest"
|
||||
"github.com/coder/coder/v2/coderd/coderdtest/promhelp"
|
||||
"github.com/coder/coder/v2/coderd/database"
|
||||
"github.com/coder/coder/v2/coderd/database/dbauthz"
|
||||
"github.com/coder/coder/v2/coderd/database/dbgen"
|
||||
@@ -419,21 +421,23 @@ func TestIsBudgetExceeded(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
cases := []struct {
|
||||
name string
|
||||
userIDStr string
|
||||
setupMocks func(db *dbmock.MockStore, userID uuid.UUID) *proto.IsBudgetExceededResponse
|
||||
wantErrContains string
|
||||
name string
|
||||
userIDStr string
|
||||
setupMocks func(db *dbmock.MockStore, userID uuid.UUID) (resp *proto.IsBudgetExceededResponse, blockedGroupID uuid.UUID)
|
||||
wantErrContains string
|
||||
wantMetricOutcome string
|
||||
}{
|
||||
{
|
||||
// Invalid UUID short-circuits before any store call.
|
||||
name: "invalid user_id",
|
||||
userIDStr: "not-a-uuid",
|
||||
wantErrContains: "invalid user_id",
|
||||
name: "invalid user_id",
|
||||
userIDStr: "not-a-uuid",
|
||||
wantErrContains: "invalid user_id",
|
||||
wantMetricOutcome: "error",
|
||||
},
|
||||
{
|
||||
// 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 {
|
||||
setupMocks: func(db *dbmock.MockStore, userID uuid.UUID) (*proto.IsBudgetExceededResponse, uuid.UUID) {
|
||||
db.EXPECT().GetUserAIBudgetOverride(gomock.Any(), userID).
|
||||
Return(database.UserAIBudgetOverride{}, sql.ErrNoRows)
|
||||
db.EXPECT().GetHighestGroupAIBudgetByUser(gomock.Any(), userID).
|
||||
@@ -441,13 +445,14 @@ func TestIsBudgetExceeded(t *testing.T) {
|
||||
return &proto.IsBudgetExceededResponse{
|
||||
Exceeded: false,
|
||||
SpendLimitMicros: nil,
|
||||
}
|
||||
}, uuid.Nil
|
||||
},
|
||||
wantMetricOutcome: "allowed",
|
||||
},
|
||||
{
|
||||
// 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 {
|
||||
setupMocks: func(db *dbmock.MockStore, userID uuid.UUID) (*proto.IsBudgetExceededResponse, uuid.UUID) {
|
||||
groupID := uuid.New()
|
||||
db.EXPECT().GetUserAIBudgetOverride(gomock.Any(), userID).
|
||||
Return(database.UserAIBudgetOverride{}, sql.ErrNoRows)
|
||||
@@ -458,13 +463,14 @@ func TestIsBudgetExceeded(t *testing.T) {
|
||||
return &proto.IsBudgetExceededResponse{
|
||||
Exceeded: false,
|
||||
SpendLimitMicros: ptr.Ref(int64(1_000)),
|
||||
}
|
||||
}, uuid.Nil
|
||||
},
|
||||
wantMetricOutcome: "allowed",
|
||||
},
|
||||
{
|
||||
// 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 {
|
||||
setupMocks: func(db *dbmock.MockStore, userID uuid.UUID) (*proto.IsBudgetExceededResponse, uuid.UUID) {
|
||||
groupID := uuid.New()
|
||||
db.EXPECT().GetUserAIBudgetOverride(gomock.Any(), userID).
|
||||
Return(database.UserAIBudgetOverride{}, sql.ErrNoRows)
|
||||
@@ -475,14 +481,15 @@ func TestIsBudgetExceeded(t *testing.T) {
|
||||
return &proto.IsBudgetExceededResponse{
|
||||
Exceeded: true,
|
||||
SpendLimitMicros: ptr.Ref(int64(1_000)),
|
||||
}
|
||||
}, groupID
|
||||
},
|
||||
wantMetricOutcome: "blocked",
|
||||
},
|
||||
{
|
||||
// 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 {
|
||||
setupMocks: func(db *dbmock.MockStore, userID uuid.UUID) (*proto.IsBudgetExceededResponse, uuid.UUID) {
|
||||
groupID := uuid.New()
|
||||
db.EXPECT().GetUserAIBudgetOverride(gomock.Any(), userID).
|
||||
Return(database.UserAIBudgetOverride{}, sql.ErrNoRows)
|
||||
@@ -493,13 +500,14 @@ func TestIsBudgetExceeded(t *testing.T) {
|
||||
return &proto.IsBudgetExceededResponse{
|
||||
Exceeded: true,
|
||||
SpendLimitMicros: ptr.Ref(int64(0)),
|
||||
}
|
||||
}, groupID
|
||||
},
|
||||
wantMetricOutcome: "blocked",
|
||||
},
|
||||
{
|
||||
// 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 {
|
||||
setupMocks: func(db *dbmock.MockStore, userID uuid.UUID) (*proto.IsBudgetExceededResponse, uuid.UUID) {
|
||||
groupID := uuid.New()
|
||||
db.EXPECT().GetUserAIBudgetOverride(gomock.Any(), userID).
|
||||
Return(database.UserAIBudgetOverride{}, sql.ErrNoRows)
|
||||
@@ -510,14 +518,15 @@ func TestIsBudgetExceeded(t *testing.T) {
|
||||
return &proto.IsBudgetExceededResponse{
|
||||
Exceeded: true,
|
||||
SpendLimitMicros: ptr.Ref(int64(1_000)),
|
||||
}
|
||||
}, groupID
|
||||
},
|
||||
wantMetricOutcome: "blocked",
|
||||
},
|
||||
{
|
||||
// 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 {
|
||||
setupMocks: func(db *dbmock.MockStore, userID uuid.UUID) (*proto.IsBudgetExceededResponse, uuid.UUID) {
|
||||
overrideGroupID := uuid.New()
|
||||
db.EXPECT().GetUserAIBudgetOverride(gomock.Any(), userID).
|
||||
Return(database.UserAIBudgetOverride{
|
||||
@@ -531,32 +540,35 @@ func TestIsBudgetExceeded(t *testing.T) {
|
||||
return &proto.IsBudgetExceededResponse{
|
||||
Exceeded: true,
|
||||
SpendLimitMicros: ptr.Ref(int64(500)),
|
||||
}
|
||||
}, overrideGroupID
|
||||
},
|
||||
wantMetricOutcome: "blocked",
|
||||
},
|
||||
{
|
||||
// Unexpected error from budget override lookup propagates.
|
||||
name: "budget resolution error propagates",
|
||||
setupMocks: func(db *dbmock.MockStore, userID uuid.UUID) *proto.IsBudgetExceededResponse {
|
||||
setupMocks: func(db *dbmock.MockStore, userID uuid.UUID) (*proto.IsBudgetExceededResponse, uuid.UUID) {
|
||||
db.EXPECT().GetUserAIBudgetOverride(gomock.Any(), userID).
|
||||
Return(database.UserAIBudgetOverride{}, sql.ErrConnDone)
|
||||
return nil
|
||||
return nil, uuid.Nil
|
||||
},
|
||||
wantErrContains: "resolve effective AI budget",
|
||||
wantErrContains: "resolve effective AI budget",
|
||||
wantMetricOutcome: "error",
|
||||
},
|
||||
{
|
||||
// Error from spend aggregation propagates (fail-closed).
|
||||
name: "spend aggregation error propagates",
|
||||
setupMocks: func(db *dbmock.MockStore, userID uuid.UUID) *proto.IsBudgetExceededResponse {
|
||||
setupMocks: func(db *dbmock.MockStore, userID uuid.UUID) (*proto.IsBudgetExceededResponse, uuid.UUID) {
|
||||
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
|
||||
return nil, uuid.Nil
|
||||
},
|
||||
wantErrContains: "get user AI spend",
|
||||
wantErrContains: "get user AI spend",
|
||||
wantMetricOutcome: "error",
|
||||
},
|
||||
}
|
||||
|
||||
@@ -575,10 +587,13 @@ func TestIsBudgetExceeded(t *testing.T) {
|
||||
}
|
||||
|
||||
var wantResp *proto.IsBudgetExceededResponse
|
||||
var blockedGroupID uuid.UUID
|
||||
if tc.setupMocks != nil {
|
||||
wantResp = tc.setupMocks(db, userID)
|
||||
wantResp, blockedGroupID = tc.setupMocks(db, userID)
|
||||
}
|
||||
|
||||
reg := prometheus.NewRegistry()
|
||||
metrics := aibridgedserver.NewMetrics(reg)
|
||||
srv, err := aibridgedserver.NewServer(t.Context(), aibridgedserver.Options{
|
||||
Store: db,
|
||||
AISeatTracker: agplaiseats.Noop{},
|
||||
@@ -587,11 +602,30 @@ func TestIsBudgetExceeded(t *testing.T) {
|
||||
Experiments: requiredExperiments,
|
||||
Logger: logger,
|
||||
Clock: quartz.NewReal(),
|
||||
Metrics: metrics,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
req := &proto.IsBudgetExceededRequest{UserId: userIDStr}
|
||||
resp, err := srv.IsBudgetExceeded(t.Context(), req)
|
||||
|
||||
// The enforcement duration is always observed once, labeled by the
|
||||
// outcome, even when the check errors.
|
||||
require.Equal(t, 1, promtest.CollectAndCount(metrics.EnforcementDuration))
|
||||
require.EqualValues(t, 1, promhelp.HistogramValue(t, reg,
|
||||
"cost_control_enforcement_duration_seconds",
|
||||
prometheus.Labels{"outcome": tc.wantMetricOutcome}).GetSampleCount())
|
||||
wantBlocked := 0
|
||||
if tc.wantMetricOutcome == "blocked" {
|
||||
wantBlocked = 1
|
||||
}
|
||||
require.Equal(t, wantBlocked, promtest.CollectAndCount(metrics.BlockedRequests))
|
||||
if wantBlocked == 1 {
|
||||
require.Equal(t, 1, promhelp.CounterValue(t, reg,
|
||||
"cost_control_blocked_requests_total",
|
||||
prometheus.Labels{"group_id": blockedGroupID.String()}))
|
||||
}
|
||||
|
||||
if tc.wantErrContains != "" {
|
||||
require.Error(t, err)
|
||||
require.Nil(t, resp)
|
||||
@@ -1714,6 +1748,11 @@ func TestRecordTokenUsage(t *testing.T) {
|
||||
db.EXPECT().GetUserAISpendSince(gomock.Any(), gomock.Any()).
|
||||
Return(database.GetUserAISpendSinceRow{SpendMicros: wantCost}, nil)
|
||||
},
|
||||
// A priced model does not increment unpriced_token_usage_records_total.
|
||||
assertMetrics: func(t *testing.T, reg *prometheus.Registry) {
|
||||
require.Nil(t, promhelp.MetricValue(t, reg, "cost_control_unpriced_token_usage_records_total",
|
||||
prometheus.Labels{"provider": "anthropic", "model": "claude-sonnet-4-6"}))
|
||||
},
|
||||
},
|
||||
{
|
||||
// Budget resolves via user override, model is priced.
|
||||
@@ -1856,6 +1895,11 @@ func TestRecordTokenUsage(t *testing.T) {
|
||||
// Spend update is skipped because cost is NULL.
|
||||
db.EXPECT().IncrementUserAIDailySpend(gomock.Any(), gomock.Any()).Times(0)
|
||||
},
|
||||
// A missing price row increments unpriced_token_usage_records_total.
|
||||
assertMetrics: func(t *testing.T, reg *prometheus.Registry) {
|
||||
require.Equal(t, 1, promhelp.CounterValue(t, reg, "cost_control_unpriced_token_usage_records_total",
|
||||
prometheus.Labels{"provider": "anthropic", "model": "claude-sonnet-4-6"}))
|
||||
},
|
||||
},
|
||||
{
|
||||
// Price row exists with NULL columns, so cost is 0 (Valid).
|
||||
@@ -2031,6 +2075,11 @@ func TestRecordTokenUsage(t *testing.T) {
|
||||
// Spend update is skipped because cost is NULL.
|
||||
db.EXPECT().IncrementUserAIDailySpend(gomock.Any(), gomock.Any()).Times(0)
|
||||
},
|
||||
// A missing price row increments unpriced_token_usage_records_total.
|
||||
assertMetrics: func(t *testing.T, reg *prometheus.Registry) {
|
||||
require.Equal(t, 1, promhelp.CounterValue(t, reg, "cost_control_unpriced_token_usage_records_total",
|
||||
prometheus.Labels{"provider": "anthropic", "model": "claude-sonnet-4-6"}))
|
||||
},
|
||||
},
|
||||
{
|
||||
// A user with no organization has no effective group. Spend is
|
||||
@@ -3168,6 +3217,9 @@ type testRecordMethodCase[Req any] struct {
|
||||
// setupMocks is called with the mock store and the above request.
|
||||
setupMocks func(t *testing.T, db *dbmock.MockStore, req Req)
|
||||
expectedErr string
|
||||
// assertMetrics, when set, is called after the method returns to assert
|
||||
// the metrics recorded on the server's registry.
|
||||
assertMetrics func(t *testing.T, reg *prometheus.Registry)
|
||||
}
|
||||
|
||||
// testRecordMethod is a helper that abstracts the common testing pattern for all Record* methods.
|
||||
@@ -3191,6 +3243,8 @@ func testRecordMethod[Req any, Resp any](
|
||||
}
|
||||
|
||||
ctx := testutil.Context(t, testutil.WaitLong)
|
||||
reg := prometheus.NewRegistry()
|
||||
metrics := aibridgedserver.NewMetrics(reg)
|
||||
srv, err := aibridgedserver.NewServer(ctx, aibridgedserver.Options{
|
||||
Store: db,
|
||||
AISeatTracker: agplaiseats.Noop{},
|
||||
@@ -3199,6 +3253,7 @@ func testRecordMethod[Req any, Resp any](
|
||||
Experiments: requiredExperiments,
|
||||
Logger: logger,
|
||||
Clock: quartz.NewReal(),
|
||||
Metrics: metrics,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
@@ -3210,6 +3265,9 @@ func testRecordMethod[Req any, Resp any](
|
||||
require.NoError(t, err, "Unexpected error for test case: %s", tc.name)
|
||||
require.NotNil(t, resp)
|
||||
}
|
||||
if tc.assertMetrics != nil {
|
||||
tc.assertMetrics(t, reg)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
@@ -73,6 +73,9 @@ func (s *Server) resolveTokenUsageCost(ctx context.Context, intc database.AIBrid
|
||||
// Model not in the price table: record tokens but leave cost NULL.
|
||||
s.logger.Debug(ctx, "no price found for model, recording token usage with NULL cost",
|
||||
slog.F("provider", intc.Provider), slog.F("model", intc.Model))
|
||||
if s.metrics != nil {
|
||||
s.metrics.UnpricedTokenUsageRecords.WithLabelValues(intc.Provider, intc.Model).Inc()
|
||||
}
|
||||
return result, nil
|
||||
case err != nil:
|
||||
return tokenUsageCost{}, xerrors.Errorf("look up model price for %s/%s: %w", intc.Provider, intc.Model, err)
|
||||
|
||||
@@ -0,0 +1,136 @@
|
||||
package aibridgedserver
|
||||
|
||||
import (
|
||||
"context"
|
||||
"time"
|
||||
|
||||
"github.com/prometheus/client_golang/prometheus"
|
||||
"github.com/prometheus/client_golang/prometheus/promauto"
|
||||
"golang.org/x/xerrors"
|
||||
|
||||
"cdr.dev/slog/v3"
|
||||
"github.com/coder/coder/v2/coderd/aibridge/budget"
|
||||
"github.com/coder/coder/v2/coderd/database"
|
||||
"github.com/coder/coder/v2/coderd/database/dbauthz"
|
||||
"github.com/coder/coder/v2/codersdk"
|
||||
"github.com/coder/quartz"
|
||||
)
|
||||
|
||||
// blockedUsersRefreshInterval is the default cadence for recomputing the
|
||||
// blocked_users gauge.
|
||||
const blockedUsersRefreshInterval = 5 * time.Minute
|
||||
|
||||
// Metrics holds the AI budget cost-control metrics emitted by the aibridged
|
||||
// server.
|
||||
type Metrics struct {
|
||||
// Requests blocked because the initiator's AI budget was exceeded.
|
||||
BlockedRequests *prometheus.CounterVec
|
||||
// Users currently over their AI budget. Updated periodically.
|
||||
BlockedUsers *prometheus.GaugeVec
|
||||
// Recorded token-usage records for which no model price was found.
|
||||
UnpricedTokenUsageRecords *prometheus.CounterVec
|
||||
// Duration of budget enforcement checks.
|
||||
EnforcementDuration *prometheus.HistogramVec
|
||||
}
|
||||
|
||||
// NewMetrics creates and registers metrics. It will panic if a collector has
|
||||
// already been registered. The provided registerer may specify a namespace
|
||||
// prefix using [prometheus.WrapRegistererWithPrefix].
|
||||
func NewMetrics(reg prometheus.Registerer) *Metrics {
|
||||
return &Metrics{
|
||||
// Pessimistic cardinality: one series per group, bounded per deployment.
|
||||
BlockedRequests: promauto.With(reg).NewCounterVec(prometheus.CounterOpts{
|
||||
Subsystem: "cost_control",
|
||||
Name: "blocked_requests_total",
|
||||
Help: "The number of AI requests blocked because the initiator's budget was exceeded.",
|
||||
}, []string{"group_id"}),
|
||||
// Pessimistic cardinality: one series per group with an over-budget user.
|
||||
BlockedUsers: promauto.With(reg).NewGaugeVec(prometheus.GaugeOpts{
|
||||
Subsystem: "cost_control",
|
||||
Name: "blocked_users",
|
||||
Help: "The number of users currently over their AI budget.",
|
||||
}, []string{"group_id"}),
|
||||
// Pessimistic cardinality: 3 providers, 5 models = up to 15.
|
||||
UnpricedTokenUsageRecords: promauto.With(reg).NewCounterVec(prometheus.CounterOpts{
|
||||
Subsystem: "cost_control",
|
||||
Name: "unpriced_token_usage_records_total",
|
||||
Help: "The number of recorded AI token-usage records for which no model price was found " +
|
||||
"(provider: anthropic, openai, copilot).",
|
||||
}, []string{"provider", "model"}),
|
||||
// Pessimistic cardinality: 3 outcomes, 8 buckets + 3 extra series
|
||||
// (count, sum, +Inf) = up to 33.
|
||||
EnforcementDuration: promauto.With(reg).NewHistogramVec(prometheus.HistogramOpts{
|
||||
Subsystem: "cost_control",
|
||||
Name: "enforcement_duration_seconds",
|
||||
Help: "The duration of AI budget enforcement checks, in seconds " +
|
||||
"(outcome: allowed, blocked, error).",
|
||||
Buckets: []float64{0.005, 0.01, 0.025, 0.05, 0.1, 0.25, 0.5, 1},
|
||||
NativeHistogramBucketFactor: 1.1,
|
||||
NativeHistogramMaxBucketNumber: 100,
|
||||
NativeHistogramMinResetDuration: time.Hour,
|
||||
NativeHistogramZeroThreshold: 0,
|
||||
NativeHistogramMaxZeroThreshold: 0,
|
||||
}, []string{"outcome"}),
|
||||
}
|
||||
}
|
||||
|
||||
// StartBlockedUsersCollector periodically updates the blocked_users gauge from
|
||||
// the database until ctx is canceled. A non-positive interval uses the 5m
|
||||
// default. The returned function stops the collector and waits for it to exit.
|
||||
// It is a no-op returning a no-op closer when m is nil, so callers need not
|
||||
// nil-check.
|
||||
func (m *Metrics) StartBlockedUsersCollector(ctx context.Context, logger slog.Logger, clk quartz.Clock, db database.Store, budgetPeriod codersdk.AIBudgetPeriod, interval time.Duration) func() {
|
||||
if m == nil {
|
||||
return func() {}
|
||||
}
|
||||
if interval <= 0 {
|
||||
interval = blockedUsersRefreshInterval
|
||||
}
|
||||
|
||||
ctx, cancel := context.WithCancel(ctx)
|
||||
done := make(chan struct{})
|
||||
ticker := clk.NewTicker(interval)
|
||||
go func() {
|
||||
defer close(done)
|
||||
defer ticker.Stop()
|
||||
// Update immediately so the gauge is populated at startup rather than
|
||||
// absent for a full interval.
|
||||
for {
|
||||
if err := m.updateBlockedUsers(ctx, clk, db, budgetPeriod); err != nil {
|
||||
logger.Error(ctx, "update blocked_users gauge", slog.Error(err))
|
||||
}
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return
|
||||
case <-ticker.C:
|
||||
}
|
||||
}
|
||||
}()
|
||||
return func() {
|
||||
cancel()
|
||||
<-done
|
||||
}
|
||||
}
|
||||
|
||||
// updateBlockedUsers sets the blocked_users gauge to the current per-group
|
||||
// count of users at or over their AI budget for the active period.
|
||||
func (m *Metrics) updateBlockedUsers(ctx context.Context, clk quartz.Clock, db database.Store, budgetPeriod codersdk.AIBudgetPeriod) error {
|
||||
period, err := budget.CurrentPeriod(clk.Now(), budgetPeriod)
|
||||
if err != nil {
|
||||
return xerrors.Errorf("compute AI budget period: %w", err)
|
||||
}
|
||||
//nolint:gocritic // Cost-control metrics need deployment-wide access to
|
||||
// group budgets and user spend.
|
||||
rows, err := db.GetOverBudgetUsersPerGroup(dbauthz.AsSystemRestricted(ctx), period.Start)
|
||||
if err != nil {
|
||||
return xerrors.Errorf("get over-budget users per group: %w", err)
|
||||
}
|
||||
|
||||
// Reset clears groups that dropped to zero since the last cycle so their
|
||||
// stale series do not linger.
|
||||
m.BlockedUsers.Reset()
|
||||
for _, row := range rows {
|
||||
m.BlockedUsers.WithLabelValues(row.GroupID.String()).Set(float64(row.OverBudgetUsers))
|
||||
}
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,153 @@
|
||||
package aibridgedserver
|
||||
|
||||
import (
|
||||
"database/sql"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/google/uuid"
|
||||
"github.com/prometheus/client_golang/prometheus"
|
||||
promtest "github.com/prometheus/client_golang/prometheus/testutil"
|
||||
"github.com/stretchr/testify/require"
|
||||
"go.uber.org/mock/gomock"
|
||||
|
||||
"github.com/coder/coder/v2/coderd/database"
|
||||
"github.com/coder/coder/v2/coderd/database/dbmock"
|
||||
"github.com/coder/coder/v2/codersdk"
|
||||
"github.com/coder/coder/v2/testutil"
|
||||
"github.com/coder/quartz"
|
||||
)
|
||||
|
||||
func TestUpdateBlockedUsers(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
groupA := uuid.New()
|
||||
groupB := uuid.New()
|
||||
|
||||
// Each round is one updateBlockedUsers call and the rows its query returns.
|
||||
// The final gauge state is asserted after the last round, so sequential
|
||||
// rounds cover Reset clearing series that dropped to zero.
|
||||
type round struct {
|
||||
rows []database.GetOverBudgetUsersPerGroupRow
|
||||
err error
|
||||
}
|
||||
tests := []struct {
|
||||
name string
|
||||
rounds []round
|
||||
wantErr error
|
||||
wantCount int
|
||||
wantValues map[uuid.UUID]float64
|
||||
}{
|
||||
{
|
||||
name: "SetsPerGroupGauges",
|
||||
rounds: []round{{rows: []database.GetOverBudgetUsersPerGroupRow{
|
||||
{GroupID: groupA, OverBudgetUsers: 3},
|
||||
{GroupID: groupB, OverBudgetUsers: 1},
|
||||
}}},
|
||||
wantCount: 2,
|
||||
wantValues: map[uuid.UUID]float64{groupA: 3, groupB: 1},
|
||||
},
|
||||
{
|
||||
// groupB drops to zero in the second round, so its stale series
|
||||
// is cleared by Reset.
|
||||
name: "ResetClearsStaleSeries",
|
||||
rounds: []round{
|
||||
{rows: []database.GetOverBudgetUsersPerGroupRow{
|
||||
{GroupID: groupA, OverBudgetUsers: 3},
|
||||
{GroupID: groupB, OverBudgetUsers: 1},
|
||||
}},
|
||||
{rows: []database.GetOverBudgetUsersPerGroupRow{
|
||||
{GroupID: groupA, OverBudgetUsers: 2},
|
||||
}},
|
||||
},
|
||||
wantCount: 1,
|
||||
wantValues: map[uuid.UUID]float64{groupA: 2},
|
||||
},
|
||||
{
|
||||
name: "PropagatesDBError",
|
||||
rounds: []round{{err: sql.ErrConnDone}},
|
||||
wantErr: sql.ErrConnDone,
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
// Given: a database returning the over-budget rows for each round.
|
||||
db := dbmock.NewMockStore(gomock.NewController(t))
|
||||
calls := make([]any, 0, len(tt.rounds))
|
||||
for _, r := range tt.rounds {
|
||||
calls = append(calls, db.EXPECT().GetOverBudgetUsersPerGroup(gomock.Any(), gomock.Any()).
|
||||
Return(r.rows, r.err))
|
||||
}
|
||||
gomock.InOrder(calls...)
|
||||
|
||||
m := NewMetrics(prometheus.NewRegistry())
|
||||
clk := quartz.NewMock(t)
|
||||
|
||||
// When: the gauge is updated once per round.
|
||||
var err error
|
||||
for range tt.rounds {
|
||||
err = m.updateBlockedUsers(t.Context(), clk, db, codersdk.AIBudgetPeriodMonth)
|
||||
}
|
||||
|
||||
// Then: the query error propagates, or the gauge holds the final
|
||||
// per-group counts.
|
||||
if tt.wantErr != nil {
|
||||
require.ErrorIs(t, err, tt.wantErr)
|
||||
return
|
||||
}
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, tt.wantCount, promtest.CollectAndCount(m.BlockedUsers))
|
||||
for group, want := range tt.wantValues {
|
||||
require.Equal(t, want, promtest.ToFloat64(m.BlockedUsers.WithLabelValues(group.String())))
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestStartBlockedUsersCollector(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
t.Run("NilReceiverNoop", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
// Given: a nil Metrics.
|
||||
var m *Metrics
|
||||
|
||||
// When: the collector is started.
|
||||
closeFn := m.StartBlockedUsersCollector(t.Context(), testutil.Logger(t), quartz.NewMock(t), nil, codersdk.AIBudgetPeriodMonth, time.Minute)
|
||||
|
||||
// Then: the returned closer is a no-op and does not panic.
|
||||
require.NotPanics(t, closeFn)
|
||||
})
|
||||
|
||||
t.Run("TicksAndStops", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
ctx := testutil.Context(t, testutil.WaitShort)
|
||||
groupA := uuid.New()
|
||||
|
||||
// Given: a running collector on a mock clock.
|
||||
db := dbmock.NewMockStore(gomock.NewController(t))
|
||||
m := NewMetrics(prometheus.NewRegistry())
|
||||
clk := quartz.NewMock(t)
|
||||
db.EXPECT().GetOverBudgetUsersPerGroup(gomock.Any(), gomock.Any()).
|
||||
Return([]database.GetOverBudgetUsersPerGroupRow{
|
||||
{GroupID: groupA, OverBudgetUsers: 4},
|
||||
}, nil).AnyTimes()
|
||||
|
||||
closeFn := m.StartBlockedUsersCollector(ctx, testutil.Logger(t), clk, db, codersdk.AIBudgetPeriodMonth, time.Minute)
|
||||
defer closeFn()
|
||||
|
||||
// When: the ticker fires.
|
||||
_, w := clk.AdvanceNext()
|
||||
w.MustWait(ctx)
|
||||
|
||||
// Then: the gauge reflects the queried per-group count.
|
||||
require.Eventually(t, func() bool {
|
||||
return promtest.ToFloat64(m.BlockedUsers.WithLabelValues(groupA.String())) == 4.0
|
||||
}, testutil.WaitShort, testutil.IntervalFast)
|
||||
})
|
||||
}
|
||||
@@ -47,6 +47,7 @@ import (
|
||||
"github.com/coder/coder/v2/coderd/agentapi/metadatabatcher"
|
||||
"github.com/coder/coder/v2/coderd/aibridge"
|
||||
"github.com/coder/coder/v2/coderd/aibridge/prices"
|
||||
"github.com/coder/coder/v2/coderd/aibridgedserver"
|
||||
"github.com/coder/coder/v2/coderd/aiseats"
|
||||
_ "github.com/coder/coder/v2/coderd/apidoc" // Used for swagger docs.
|
||||
"github.com/coder/coder/v2/coderd/appearance"
|
||||
@@ -2318,6 +2319,8 @@ type API struct {
|
||||
// routes (license-gated) which apply their own StripPrefix, and by
|
||||
// the in-memory transport (used by chatd, license-exempt).
|
||||
aiGatewayHandler http.Handler
|
||||
// AIGatewayServerMetrics records AI budget cost-control metrics. May be nil.
|
||||
AIGatewayServerMetrics *aibridgedserver.Metrics
|
||||
|
||||
UpdatesProvider tailnet.WorkspaceUpdatesProvider
|
||||
|
||||
|
||||
@@ -4321,6 +4321,14 @@ func (q *querier) GetOrganizationsWithPrebuildStatus(ctx context.Context, arg da
|
||||
return q.db.GetOrganizationsWithPrebuildStatus(ctx, arg)
|
||||
}
|
||||
|
||||
func (q *querier) GetOverBudgetUsersPerGroup(ctx context.Context, periodStart time.Time) ([]database.GetOverBudgetUsersPerGroupRow, error) {
|
||||
// Aggregates over-budget user counts per group for cost-control metrics.
|
||||
if err := q.authorizeContext(ctx, policy.ActionRead, rbac.ResourceGroup.All()); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return q.db.GetOverBudgetUsersPerGroup(ctx, periodStart)
|
||||
}
|
||||
|
||||
func (q *querier) GetParameterSchemasByJobID(ctx context.Context, jobID uuid.UUID) ([]database.ParameterSchema, error) {
|
||||
version, err := q.db.GetTemplateVersionByJobID(ctx, jobID)
|
||||
if err != nil {
|
||||
|
||||
@@ -7073,6 +7073,13 @@ func (s *MethodTestSuite) TestAIBridge() {
|
||||
check.Args(arg).Asserts(user, policy.ActionRead).Returns(row)
|
||||
}))
|
||||
|
||||
s.Run("GetOverBudgetUsersPerGroup", s.Mocked(func(dbm *dbmock.MockStore, _ *gofakeit.Faker, check *expects) {
|
||||
periodStart := time.Now().UTC().Truncate(24 * time.Hour)
|
||||
dbm.EXPECT().GetOverBudgetUsersPerGroup(gomock.Any(), periodStart).
|
||||
Return([]database.GetOverBudgetUsersPerGroupRow{}, nil).AnyTimes()
|
||||
check.Args(periodStart).Asserts(rbac.ResourceGroup.All(), policy.ActionRead)
|
||||
}))
|
||||
|
||||
s.Run("IncrementUserAIDailySpend", s.Mocked(func(dbm *dbmock.MockStore, faker *gofakeit.Faker, check *expects) {
|
||||
arg := database.IncrementUserAIDailySpendParams{
|
||||
UserID: uuid.New(),
|
||||
|
||||
+8
@@ -2609,6 +2609,14 @@ func (m queryMetricsStore) GetOrganizationsWithPrebuildStatus(ctx context.Contex
|
||||
return r0, r1
|
||||
}
|
||||
|
||||
func (m queryMetricsStore) GetOverBudgetUsersPerGroup(ctx context.Context, periodStart time.Time) ([]database.GetOverBudgetUsersPerGroupRow, error) {
|
||||
start := time.Now()
|
||||
r0, r1 := m.s.GetOverBudgetUsersPerGroup(ctx, periodStart)
|
||||
m.queryLatencies.WithLabelValues("GetOverBudgetUsersPerGroup").Observe(time.Since(start).Seconds())
|
||||
m.queryCounts.WithLabelValues(httpmw.ExtractHTTPRoute(ctx), httpmw.ExtractHTTPMethod(ctx), "GetOverBudgetUsersPerGroup").Inc()
|
||||
return r0, r1
|
||||
}
|
||||
|
||||
func (m queryMetricsStore) GetParameterSchemasByJobID(ctx context.Context, jobID uuid.UUID) ([]database.ParameterSchema, error) {
|
||||
start := time.Now()
|
||||
r0, r1 := m.s.GetParameterSchemasByJobID(ctx, jobID)
|
||||
|
||||
Generated
+15
@@ -4843,6 +4843,21 @@ func (mr *MockStoreMockRecorder) GetOrganizationsWithPrebuildStatus(ctx, arg any
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetOrganizationsWithPrebuildStatus", reflect.TypeOf((*MockStore)(nil).GetOrganizationsWithPrebuildStatus), ctx, arg)
|
||||
}
|
||||
|
||||
// GetOverBudgetUsersPerGroup mocks base method.
|
||||
func (m *MockStore) GetOverBudgetUsersPerGroup(ctx context.Context, periodStart time.Time) ([]database.GetOverBudgetUsersPerGroupRow, error) {
|
||||
m.ctrl.T.Helper()
|
||||
ret := m.ctrl.Call(m, "GetOverBudgetUsersPerGroup", ctx, periodStart)
|
||||
ret0, _ := ret[0].([]database.GetOverBudgetUsersPerGroupRow)
|
||||
ret1, _ := ret[1].(error)
|
||||
return ret0, ret1
|
||||
}
|
||||
|
||||
// GetOverBudgetUsersPerGroup indicates an expected call of GetOverBudgetUsersPerGroup.
|
||||
func (mr *MockStoreMockRecorder) GetOverBudgetUsersPerGroup(ctx, periodStart any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetOverBudgetUsersPerGroup", reflect.TypeOf((*MockStore)(nil).GetOverBudgetUsersPerGroup), ctx, periodStart)
|
||||
}
|
||||
|
||||
// GetParameterSchemasByJobID mocks base method.
|
||||
func (m *MockStore) GetParameterSchemasByJobID(ctx context.Context, jobID uuid.UUID) ([]database.ParameterSchema, error) {
|
||||
m.ctrl.T.Helper()
|
||||
|
||||
Generated
+5
@@ -690,6 +690,11 @@ type sqlcQuerier interface {
|
||||
// GetOrganizationsWithPrebuildStatus returns organizations with prebuilds configured and their
|
||||
// membership status for the prebuilds system user (org membership, group existence, group membership).
|
||||
GetOrganizationsWithPrebuildStatus(ctx context.Context, arg GetOrganizationsWithPrebuildStatusParams) ([]GetOrganizationsWithPrebuildStatusRow, error)
|
||||
// Returns, per effective group, the number of users at or over their spend
|
||||
// limit since period_start. Only users with an enforceable limit (override or
|
||||
// budgeted group) count, and the unlimited Everyone fallback does not.
|
||||
// TODO(AIGOV-527): unify effective group resolution in a single place.
|
||||
GetOverBudgetUsersPerGroup(ctx context.Context, periodStart time.Time) ([]GetOverBudgetUsersPerGroupRow, error)
|
||||
GetParameterSchemasByJobID(ctx context.Context, jobID uuid.UUID) ([]ParameterSchema, error)
|
||||
GetPrebuildMetrics(ctx context.Context) ([]GetPrebuildMetricsRow, error)
|
||||
GetPrebuildsSettings(ctx context.Context) (string, error)
|
||||
|
||||
@@ -13901,6 +13901,197 @@ func TestGetHighestGroupAIBudgetByUser(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestGetOverBudgetUsersPerGroup(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
periodStart := dbtime.Now().UTC().Truncate(24 * time.Hour)
|
||||
|
||||
// seedSpendOnDay attributes micros of spend to (user, effectiveGroup) on a
|
||||
// specific day.
|
||||
seedSpendOnDay := func(t *testing.T, ctx context.Context, db database.Store, userID, effectiveGroupID uuid.UUID, day time.Time, micros int64) {
|
||||
t.Helper()
|
||||
_, err := db.IncrementUserAIDailySpend(ctx, database.IncrementUserAIDailySpendParams{
|
||||
UserID: userID,
|
||||
EffectiveGroupID: effectiveGroupID,
|
||||
Day: day,
|
||||
CostMicros: micros,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
}
|
||||
|
||||
// seedSpend attributes micros of spend to (user, effectiveGroup) within the
|
||||
// current period.
|
||||
seedSpend := func(t *testing.T, ctx context.Context, db database.Store, userID, effectiveGroupID uuid.UUID, micros int64) {
|
||||
t.Helper()
|
||||
seedSpendOnDay(t, ctx, db, userID, effectiveGroupID, periodStart, micros)
|
||||
}
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
setup func(t *testing.T, ctx context.Context, db database.Store) []database.GetOverBudgetUsersPerGroupRow
|
||||
}{
|
||||
{
|
||||
// A user whose spend exceeds their group budget is counted.
|
||||
name: "OverBudgetCounted",
|
||||
setup: func(t *testing.T, ctx context.Context, db database.Store) []database.GetOverBudgetUsersPerGroupRow {
|
||||
user := dbgen.User(t, db, database.User{})
|
||||
org := dbgen.Organization(t, db, database.Organization{})
|
||||
group := dbgen.Group(t, db, database.Group{OrganizationID: org.ID})
|
||||
dbgen.OrganizationMember(t, db, database.OrganizationMember{OrganizationID: org.ID, UserID: user.ID})
|
||||
dbgen.GroupMember(t, db, database.GroupMemberTable{GroupID: group.ID, UserID: user.ID})
|
||||
_, err := db.UpsertGroupAIBudget(ctx, database.UpsertGroupAIBudgetParams{GroupID: group.ID, SpendLimitMicros: 1_000_000})
|
||||
require.NoError(t, err)
|
||||
seedSpend(t, ctx, db, user.ID, group.ID, 1_500_000)
|
||||
return []database.GetOverBudgetUsersPerGroupRow{{GroupID: group.ID, OverBudgetUsers: 1}}
|
||||
},
|
||||
},
|
||||
{
|
||||
// A user under their group budget is not counted.
|
||||
name: "UnderBudgetNotCounted",
|
||||
setup: func(t *testing.T, ctx context.Context, db database.Store) []database.GetOverBudgetUsersPerGroupRow {
|
||||
user := dbgen.User(t, db, database.User{})
|
||||
org := dbgen.Organization(t, db, database.Organization{})
|
||||
group := dbgen.Group(t, db, database.Group{OrganizationID: org.ID})
|
||||
dbgen.OrganizationMember(t, db, database.OrganizationMember{OrganizationID: org.ID, UserID: user.ID})
|
||||
dbgen.GroupMember(t, db, database.GroupMemberTable{GroupID: group.ID, UserID: user.ID})
|
||||
_, err := db.UpsertGroupAIBudget(ctx, database.UpsertGroupAIBudgetParams{GroupID: group.ID, SpendLimitMicros: 1_000_000})
|
||||
require.NoError(t, err)
|
||||
seedSpend(t, ctx, db, user.ID, group.ID, 500_000)
|
||||
return nil
|
||||
},
|
||||
},
|
||||
{
|
||||
// Spend exactly at the limit counts, since the check is inclusive.
|
||||
name: "AtLimitCounted",
|
||||
setup: func(t *testing.T, ctx context.Context, db database.Store) []database.GetOverBudgetUsersPerGroupRow {
|
||||
user := dbgen.User(t, db, database.User{})
|
||||
org := dbgen.Organization(t, db, database.Organization{})
|
||||
group := dbgen.Group(t, db, database.Group{OrganizationID: org.ID})
|
||||
dbgen.OrganizationMember(t, db, database.OrganizationMember{OrganizationID: org.ID, UserID: user.ID})
|
||||
dbgen.GroupMember(t, db, database.GroupMemberTable{GroupID: group.ID, UserID: user.ID})
|
||||
_, err := db.UpsertGroupAIBudget(ctx, database.UpsertGroupAIBudgetParams{GroupID: group.ID, SpendLimitMicros: 1_000_000})
|
||||
require.NoError(t, err)
|
||||
seedSpend(t, ctx, db, user.ID, group.ID, 1_000_000)
|
||||
return []database.GetOverBudgetUsersPerGroupRow{{GroupID: group.ID, OverBudgetUsers: 1}}
|
||||
},
|
||||
},
|
||||
{
|
||||
// A zero limit blocks a user with no spend, since zero spend is at the
|
||||
// limit.
|
||||
name: "ZeroLimitCounted",
|
||||
setup: func(t *testing.T, ctx context.Context, db database.Store) []database.GetOverBudgetUsersPerGroupRow {
|
||||
user := dbgen.User(t, db, database.User{})
|
||||
org := dbgen.Organization(t, db, database.Organization{})
|
||||
group := dbgen.Group(t, db, database.Group{OrganizationID: org.ID})
|
||||
dbgen.OrganizationMember(t, db, database.OrganizationMember{OrganizationID: org.ID, UserID: user.ID})
|
||||
dbgen.GroupMember(t, db, database.GroupMemberTable{GroupID: group.ID, UserID: user.ID})
|
||||
_, err := db.UpsertGroupAIBudget(ctx, database.UpsertGroupAIBudgetParams{GroupID: group.ID, SpendLimitMicros: 0})
|
||||
require.NoError(t, err)
|
||||
return []database.GetOverBudgetUsersPerGroupRow{{GroupID: group.ID, OverBudgetUsers: 1}}
|
||||
},
|
||||
},
|
||||
{
|
||||
// A per-user override overrides the group budget, both for the limit
|
||||
// and the group the spend is attributed to.
|
||||
name: "OverrideWins",
|
||||
setup: func(t *testing.T, ctx context.Context, db database.Store) []database.GetOverBudgetUsersPerGroupRow {
|
||||
user := dbgen.User(t, db, database.User{})
|
||||
org := dbgen.Organization(t, db, database.Organization{})
|
||||
group := dbgen.Group(t, db, database.Group{OrganizationID: org.ID})
|
||||
overrideGroup := dbgen.Group(t, db, database.Group{OrganizationID: org.ID})
|
||||
dbgen.OrganizationMember(t, db, database.OrganizationMember{OrganizationID: org.ID, UserID: user.ID})
|
||||
dbgen.GroupMember(t, db, database.GroupMemberTable{GroupID: group.ID, UserID: user.ID})
|
||||
dbgen.GroupMember(t, db, database.GroupMemberTable{GroupID: overrideGroup.ID, UserID: user.ID})
|
||||
_, err := db.UpsertGroupAIBudget(ctx, database.UpsertGroupAIBudgetParams{GroupID: group.ID, SpendLimitMicros: 5_000_000})
|
||||
require.NoError(t, err)
|
||||
_, err = db.UpsertUserAIBudgetOverride(ctx, database.UpsertUserAIBudgetOverrideParams{UserID: user.ID, GroupID: overrideGroup.ID, SpendLimitMicros: 1_000_000})
|
||||
require.NoError(t, err)
|
||||
// Over the override limit but under the group limit.
|
||||
seedSpend(t, ctx, db, user.ID, overrideGroup.ID, 1_500_000)
|
||||
return []database.GetOverBudgetUsersPerGroupRow{{GroupID: overrideGroup.ID, OverBudgetUsers: 1}}
|
||||
},
|
||||
},
|
||||
{
|
||||
// A user in multiple budgeted groups is attributed to their
|
||||
// highest-limit group.
|
||||
name: "HighestGroupWins",
|
||||
setup: func(t *testing.T, ctx context.Context, db database.Store) []database.GetOverBudgetUsersPerGroupRow {
|
||||
user := dbgen.User(t, db, database.User{})
|
||||
org := dbgen.Organization(t, db, database.Organization{})
|
||||
lower := dbgen.Group(t, db, database.Group{OrganizationID: org.ID})
|
||||
higher := dbgen.Group(t, db, database.Group{OrganizationID: org.ID})
|
||||
dbgen.OrganizationMember(t, db, database.OrganizationMember{OrganizationID: org.ID, UserID: user.ID})
|
||||
dbgen.GroupMember(t, db, database.GroupMemberTable{GroupID: lower.ID, UserID: user.ID})
|
||||
dbgen.GroupMember(t, db, database.GroupMemberTable{GroupID: higher.ID, UserID: user.ID})
|
||||
_, err := db.UpsertGroupAIBudget(ctx, database.UpsertGroupAIBudgetParams{GroupID: lower.ID, SpendLimitMicros: 1_000_000})
|
||||
require.NoError(t, err)
|
||||
_, err = db.UpsertGroupAIBudget(ctx, database.UpsertGroupAIBudgetParams{GroupID: higher.ID, SpendLimitMicros: 2_000_000})
|
||||
require.NoError(t, err)
|
||||
seedSpend(t, ctx, db, user.ID, higher.ID, 2_000_000)
|
||||
return []database.GetOverBudgetUsersPerGroupRow{{GroupID: higher.ID, OverBudgetUsers: 1}}
|
||||
},
|
||||
},
|
||||
{
|
||||
// A user with only the unlimited Everyone fallback is never counted.
|
||||
name: "EveryoneFallbackNotCounted",
|
||||
setup: func(t *testing.T, ctx context.Context, db database.Store) []database.GetOverBudgetUsersPerGroupRow {
|
||||
user := dbgen.User(t, db, database.User{})
|
||||
org := dbgen.Organization(t, db, database.Organization{})
|
||||
dbgen.OrganizationMember(t, db, database.OrganizationMember{OrganizationID: org.ID, UserID: user.ID})
|
||||
// Spend attributed to the Everyone group (id == organization_id).
|
||||
seedSpend(t, ctx, db, user.ID, org.ID, 9_000_000)
|
||||
return nil
|
||||
},
|
||||
},
|
||||
{
|
||||
// Multiple over-budget users in the same group are summed.
|
||||
name: "AggregatesUsersPerGroup",
|
||||
setup: func(t *testing.T, ctx context.Context, db database.Store) []database.GetOverBudgetUsersPerGroupRow {
|
||||
org := dbgen.Organization(t, db, database.Organization{})
|
||||
group := dbgen.Group(t, db, database.Group{OrganizationID: org.ID})
|
||||
_, err := db.UpsertGroupAIBudget(ctx, database.UpsertGroupAIBudgetParams{GroupID: group.ID, SpendLimitMicros: 1_000_000})
|
||||
require.NoError(t, err)
|
||||
for range 2 {
|
||||
user := dbgen.User(t, db, database.User{})
|
||||
dbgen.OrganizationMember(t, db, database.OrganizationMember{OrganizationID: org.ID, UserID: user.ID})
|
||||
dbgen.GroupMember(t, db, database.GroupMemberTable{GroupID: group.ID, UserID: user.ID})
|
||||
seedSpend(t, ctx, db, user.ID, group.ID, 2_000_000)
|
||||
}
|
||||
return []database.GetOverBudgetUsersPerGroupRow{{GroupID: group.ID, OverBudgetUsers: 2}}
|
||||
},
|
||||
},
|
||||
{
|
||||
// Spend on days before the period start is excluded, so a user whose
|
||||
// only over-limit spend predates the period is not counted.
|
||||
name: "SpendBeforePeriodNotCounted",
|
||||
setup: func(t *testing.T, ctx context.Context, db database.Store) []database.GetOverBudgetUsersPerGroupRow {
|
||||
user := dbgen.User(t, db, database.User{})
|
||||
org := dbgen.Organization(t, db, database.Organization{})
|
||||
group := dbgen.Group(t, db, database.Group{OrganizationID: org.ID})
|
||||
dbgen.OrganizationMember(t, db, database.OrganizationMember{OrganizationID: org.ID, UserID: user.ID})
|
||||
dbgen.GroupMember(t, db, database.GroupMemberTable{GroupID: group.ID, UserID: user.ID})
|
||||
_, err := db.UpsertGroupAIBudget(ctx, database.UpsertGroupAIBudgetParams{GroupID: group.ID, SpendLimitMicros: 1_000_000})
|
||||
require.NoError(t, err)
|
||||
seedSpendOnDay(t, ctx, db, user.ID, group.ID, periodStart.AddDate(0, 0, -1), 1_500_000)
|
||||
return nil
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
db, _ := dbtestutil.NewDB(t)
|
||||
ctx := testutil.Context(t, testutil.WaitShort)
|
||||
|
||||
want := tt.setup(t, ctx, db)
|
||||
got, err := db.GetOverBudgetUsersPerGroup(ctx, periodStart)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, want, got)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestGetUserEveryoneFallbackGroup(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
|
||||
Generated
+91
@@ -2819,6 +2819,97 @@ func (q *sqlQuerier) GetOrganizationGroupsAISpend(ctx context.Context, arg GetOr
|
||||
return items, nil
|
||||
}
|
||||
|
||||
const getOverBudgetUsersPerGroup = `-- name: GetOverBudgetUsersPerGroup :many
|
||||
WITH budgeted_users AS (
|
||||
-- Users with an override or membership in a budgeted group.
|
||||
SELECT user_id FROM user_ai_budget_overrides
|
||||
UNION
|
||||
SELECT DISTINCT member.user_id
|
||||
FROM group_ai_budgets budget
|
||||
JOIN group_members_expanded member ON member.group_id = budget.group_id
|
||||
),
|
||||
user_highest_group AS (
|
||||
-- Per user, their highest-limit group ("highest" budget policy).
|
||||
SELECT DISTINCT ON (member.user_id)
|
||||
member.user_id,
|
||||
budget.group_id,
|
||||
budget.spend_limit_micros
|
||||
FROM group_ai_budgets budget
|
||||
JOIN group_members_expanded member ON member.group_id = budget.group_id
|
||||
JOIN organizations ON organizations.id = member.organization_id
|
||||
JOIN organization_members
|
||||
ON organization_members.user_id = member.user_id
|
||||
AND organization_members.organization_id = member.organization_id
|
||||
WHERE member.user_id IN (SELECT user_id FROM budgeted_users)
|
||||
AND organizations.deleted = false
|
||||
ORDER BY member.user_id, budget.spend_limit_micros DESC, organization_members.created_at ASC, budget.group_id ASC
|
||||
),
|
||||
effective AS (
|
||||
-- An override wins over the highest-limit group, and users with neither drop.
|
||||
SELECT
|
||||
budgeted_users.user_id,
|
||||
COALESCE(override.group_id, user_highest_group.group_id) AS effective_group_id,
|
||||
COALESCE(override.spend_limit_micros, user_highest_group.spend_limit_micros) AS spend_limit_micros
|
||||
FROM budgeted_users
|
||||
LEFT JOIN user_ai_budget_overrides override ON override.user_id = budgeted_users.user_id
|
||||
LEFT JOIN user_highest_group ON user_highest_group.user_id = budgeted_users.user_id
|
||||
WHERE COALESCE(override.group_id, user_highest_group.group_id) IS NOT NULL
|
||||
),
|
||||
user_spend AS (
|
||||
-- Each user's spend against their effective group since period_start.
|
||||
SELECT
|
||||
effective.user_id,
|
||||
effective.effective_group_id,
|
||||
effective.spend_limit_micros,
|
||||
COALESCE(SUM(spend.spend_micros), 0)::BIGINT AS current_spend_micros
|
||||
FROM effective
|
||||
LEFT JOIN ai_user_daily_spend spend
|
||||
ON spend.user_id = effective.user_id
|
||||
AND spend.effective_group_id = effective.effective_group_id
|
||||
AND spend.day >= (($1::timestamptz) AT TIME ZONE 'UTC')::date
|
||||
GROUP BY effective.user_id, effective.effective_group_id, effective.spend_limit_micros
|
||||
)
|
||||
SELECT
|
||||
effective_group_id AS group_id,
|
||||
COUNT(*)::BIGINT AS over_budget_users
|
||||
FROM user_spend
|
||||
WHERE current_spend_micros >= spend_limit_micros
|
||||
GROUP BY effective_group_id
|
||||
ORDER BY effective_group_id
|
||||
`
|
||||
|
||||
type GetOverBudgetUsersPerGroupRow struct {
|
||||
GroupID uuid.UUID `db:"group_id" json:"group_id"`
|
||||
OverBudgetUsers int64 `db:"over_budget_users" json:"over_budget_users"`
|
||||
}
|
||||
|
||||
// Returns, per effective group, the number of users at or over their spend
|
||||
// limit since period_start. Only users with an enforceable limit (override or
|
||||
// budgeted group) count, and the unlimited Everyone fallback does not.
|
||||
// TODO(AIGOV-527): unify effective group resolution in a single place.
|
||||
func (q *sqlQuerier) GetOverBudgetUsersPerGroup(ctx context.Context, periodStart time.Time) ([]GetOverBudgetUsersPerGroupRow, error) {
|
||||
rows, err := q.db.QueryContext(ctx, getOverBudgetUsersPerGroup, periodStart)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer rows.Close()
|
||||
var items []GetOverBudgetUsersPerGroupRow
|
||||
for rows.Next() {
|
||||
var i GetOverBudgetUsersPerGroupRow
|
||||
if err := rows.Scan(&i.GroupID, &i.OverBudgetUsers); 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 getUserAIBudgetOverride = `-- name: GetUserAIBudgetOverride :one
|
||||
SELECT user_id, group_id, spend_limit_micros, created_at, updated_at
|
||||
FROM user_ai_budget_overrides
|
||||
|
||||
@@ -243,3 +243,65 @@ GROUP BY
|
||||
applied_budget.spend_limit_micros,
|
||||
applied_budget.limit_source
|
||||
ORDER BY effective.user_id;
|
||||
|
||||
-- name: GetOverBudgetUsersPerGroup :many
|
||||
-- Returns, per effective group, the number of users at or over their spend
|
||||
-- limit since period_start. Only users with an enforceable limit (override or
|
||||
-- budgeted group) count, and the unlimited Everyone fallback does not.
|
||||
-- TODO(AIGOV-527): unify effective group resolution in a single place.
|
||||
WITH budgeted_users AS (
|
||||
-- Users with an override or membership in a budgeted group.
|
||||
SELECT user_id FROM user_ai_budget_overrides
|
||||
UNION
|
||||
SELECT DISTINCT member.user_id
|
||||
FROM group_ai_budgets budget
|
||||
JOIN group_members_expanded member ON member.group_id = budget.group_id
|
||||
),
|
||||
user_highest_group AS (
|
||||
-- Per user, their highest-limit group ("highest" budget policy).
|
||||
SELECT DISTINCT ON (member.user_id)
|
||||
member.user_id,
|
||||
budget.group_id,
|
||||
budget.spend_limit_micros
|
||||
FROM group_ai_budgets budget
|
||||
JOIN group_members_expanded member ON member.group_id = budget.group_id
|
||||
JOIN organizations ON organizations.id = member.organization_id
|
||||
JOIN organization_members
|
||||
ON organization_members.user_id = member.user_id
|
||||
AND organization_members.organization_id = member.organization_id
|
||||
WHERE member.user_id IN (SELECT user_id FROM budgeted_users)
|
||||
AND organizations.deleted = false
|
||||
ORDER BY member.user_id, budget.spend_limit_micros DESC, organization_members.created_at ASC, budget.group_id ASC
|
||||
),
|
||||
effective AS (
|
||||
-- An override wins over the highest-limit group, and users with neither drop.
|
||||
SELECT
|
||||
budgeted_users.user_id,
|
||||
COALESCE(override.group_id, user_highest_group.group_id) AS effective_group_id,
|
||||
COALESCE(override.spend_limit_micros, user_highest_group.spend_limit_micros) AS spend_limit_micros
|
||||
FROM budgeted_users
|
||||
LEFT JOIN user_ai_budget_overrides override ON override.user_id = budgeted_users.user_id
|
||||
LEFT JOIN user_highest_group ON user_highest_group.user_id = budgeted_users.user_id
|
||||
WHERE COALESCE(override.group_id, user_highest_group.group_id) IS NOT NULL
|
||||
),
|
||||
user_spend AS (
|
||||
-- Each user's spend against their effective group since period_start.
|
||||
SELECT
|
||||
effective.user_id,
|
||||
effective.effective_group_id,
|
||||
effective.spend_limit_micros,
|
||||
COALESCE(SUM(spend.spend_micros), 0)::BIGINT AS current_spend_micros
|
||||
FROM effective
|
||||
LEFT JOIN ai_user_daily_spend spend
|
||||
ON spend.user_id = effective.user_id
|
||||
AND spend.effective_group_id = effective.effective_group_id
|
||||
AND spend.day >= ((@period_start::timestamptz) AT TIME ZONE 'UTC')::date
|
||||
GROUP BY effective.user_id, effective.effective_group_id, effective.spend_limit_micros
|
||||
)
|
||||
SELECT
|
||||
effective_group_id AS group_id,
|
||||
COUNT(*)::BIGINT AS over_budget_users
|
||||
FROM user_spend
|
||||
WHERE current_spend_micros >= spend_limit_micros
|
||||
GROUP BY effective_group_id
|
||||
ORDER BY effective_group_id;
|
||||
|
||||
Reference in New Issue
Block a user