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:
Susana Ferreira
2026-07-28 09:22:58 +01:00
committed by GitHub
parent bfcfb71860
commit c351280a37
20 changed files with 835 additions and 30 deletions
+1
View File
@@ -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
+29 -4
View File
@@ -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
}
+84 -26
View File
@@ -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)
}
})
}
}
+3
View File
@@ -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)
+136
View File
@@ -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)
})
}
+3
View File
@@ -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
+8
View File
@@ -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 {
+7
View File
@@ -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
View File
@@ -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)
+15
View File
@@ -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()
+5
View File
@@ -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)
+191
View File
@@ -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()
+91
View File
@@ -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
+62
View File
@@ -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;