Files
sub2api/backend/internal/service/grok_quota_service_test.go
T
李建琦 6d655c9903
Release / update-version (push) Has been cancelled
Release / build-frontend (push) Has been cancelled
Release / release (push) Has been cancelled
Release / sync-version-file (push) Has been cancelled
CI / shell (push) Canceled after 0s
CI / test (push) Canceled after 0s
CI / frontend (push) Canceled after 0s
CI / golangci-lint (push) Canceled after 0s
Security Scan / backend-security (push) Canceled after 0s
Security Scan / frontend-security (push) Canceled after 0s
Sub2API v1.0 - AI API 网关(二开初始版本,基于上游 Wei-Shaw/sub2api)
2026-08-21 18:30:13 +08:00

1271 lines
44 KiB
Go

//go:build unit
package service
import (
"bytes"
"context"
"errors"
"io"
"log/slog"
"net/http"
"strconv"
"strings"
"sync"
"testing"
"time"
"github.com/Wei-Shaw/sub2api/internal/config"
infraerrors "github.com/Wei-Shaw/sub2api/internal/pkg/errors"
"github.com/Wei-Shaw/sub2api/internal/pkg/usagestats"
"github.com/Wei-Shaw/sub2api/internal/pkg/xai"
"github.com/stretchr/testify/require"
"github.com/tidwall/gjson"
)
type grokQuotaAccountRepo struct {
*mockAccountRepoForPlatform
updates map[int64]map[string]any
updateCalls int
rateLimitedCalls int
lastRateLimitedID int64
lastRateLimitResetAt time.Time
tempUnschedCalls int
lastTempUnschedID int64
lastTempUnschedUntil time.Time
lastTempUnschedReason string
recoveryClearCalls int
recoveryObservedAt time.Time
recoveryObservedReset time.Time
recoveryClearResult bool
}
func (r *grokQuotaAccountRepo) UpdateExtra(_ context.Context, id int64, updates map[string]any) error {
r.updateCalls++
if r.updates == nil {
r.updates = make(map[int64]map[string]any)
}
r.updates[id] = updates
if r.mockAccountRepoForPlatform != nil {
account := r.accountsByID[id]
if account == nil {
return nil
}
if account.Extra == nil {
account.Extra = make(map[string]any)
}
for key, value := range updates {
account.Extra[key] = value
}
}
return nil
}
func (r *grokQuotaAccountRepo) SetRateLimited(_ context.Context, id int64, resetAt time.Time) error {
r.rateLimitedCalls++
r.lastRateLimitedID = id
r.lastRateLimitResetAt = resetAt
return nil
}
func (r *grokQuotaAccountRepo) SetRateLimitedIfLater(ctx context.Context, id int64, resetAt time.Time) error {
return r.SetRateLimited(ctx, id, resetAt)
}
func (r *grokQuotaAccountRepo) ClearRateLimitIfObserved(_ context.Context, _ int64, observedLimitedAt, observedResetAt time.Time) (bool, error) {
r.recoveryClearCalls++
r.recoveryObservedAt = observedLimitedAt
r.recoveryObservedReset = observedResetAt
return r.recoveryClearResult, nil
}
func (r *grokQuotaAccountRepo) SetTempUnschedulable(_ context.Context, id int64, until time.Time, reason string) error {
r.tempUnschedCalls++
r.lastTempUnschedID = id
r.lastTempUnschedUntil = until
r.lastTempUnschedReason = reason
return nil
}
func TestSyncGrokObservedModelsRejectsOAuthCustomURLOutsideOperatorPolicy(t *testing.T) {
account := &Account{
ID: 901,
Platform: PlatformGrok,
Type: AccountTypeOAuth,
Credentials: map[string]any{
"access_token": "secret-token",
"base_url": "https://blocked.example.test/v1",
},
}
repo := &grokQuotaAccountRepo{mockAccountRepoForPlatform: &mockAccountRepoForPlatform{
accountsByID: map[int64]*Account{account.ID: account},
}}
upstream := &httpUpstreamRecorder{}
cfg := &config.Config{}
cfg.Security.URLAllowlist.Enabled = true
cfg.Security.URLAllowlist.UpstreamHosts = []string{"allowed.example.test"}
svc := &GrokQuotaService{accountRepo: repo, httpUpstream: upstream, cfg: cfg}
err := svc.syncGrokObservedModels(context.Background(), account)
require.ErrorContains(t, err, "base URL rejected by URL security policy")
require.Nil(t, upstream.lastReq)
}
func TestSyncGrokObservedModelsUsesCLIIdentityAndAccountHeaders(t *testing.T) {
account := &Account{
ID: 902,
Platform: PlatformGrok,
Type: AccountTypeOAuth,
Credentials: map[string]any{
"access_token": "secret-token",
"sub": "user-902",
"email": "user902@example.test",
},
}
repo := &grokQuotaAccountRepo{mockAccountRepoForPlatform: &mockAccountRepoForPlatform{
accountsByID: map[int64]*Account{account.ID: account},
}}
upstream := &httpUpstreamRecorder{resp: &http.Response{
StatusCode: http.StatusOK,
Body: io.NopCloser(strings.NewReader(`{"data":[{"id":"grok-4.5"}]}`)),
}}
svc := &GrokQuotaService{accountRepo: repo, httpUpstream: upstream, cfg: &config.Config{}}
require.NoError(t, svc.syncGrokObservedModels(context.Background(), account))
require.Equal(t, xai.DefaultCLIBaseURL+"/models", upstream.lastReq.URL.String())
require.NotEmpty(t, upstream.lastReq.Header.Get("x-grok-client-version"))
require.Equal(t, xai.CLIClientIdentifier, upstream.lastReq.Header.Get("x-grok-client-identifier"))
require.Equal(t, "interactive", upstream.lastReq.Header.Get("X-Grok-Client-Mode"))
require.Equal(t, "user-902", upstream.lastReq.Header.Get("X-UserID"))
require.Equal(t, "user902@example.test", upstream.lastReq.Header.Get("X-Email"))
require.Contains(t, repo.updates[account.ID], grokObservedModelsExtraKey)
}
type grokQuotaProxyRepo struct {
proxyRepoStub
proxies map[int64]*Proxy
calls int
}
type grokQuotaUsageLogRepo struct {
UsageLogRepository
stats *usagestats.AccountStats
err error
calls int
startTimes []time.Time
}
func (r *grokQuotaUsageLogRepo) GetAccountWindowStats(_ context.Context, _ int64, start time.Time) (*usagestats.AccountStats, error) {
r.calls++
r.startTimes = append(r.startTimes, start)
return r.stats, r.err
}
func (r *grokQuotaUsageLogRepo) GetAccountTodayStats(context.Context, int64) (*usagestats.AccountStats, error) {
return nil, nil
}
type grokHybridUpstream struct {
httpUpstreamRecorder
mu sync.Mutex
requests []*http.Request
bodies [][]byte
weeklyUsagePercent *float64
monthlyLimitCents *float64
activeStatus int
activeHeaders http.Header
billingStarted chan struct{}
billingRelease <-chan struct{}
billingStartOnce sync.Once
billingStatus int
weeklyBillingStatus int
monthlyBillingStatus int
billingHeaders http.Header
}
type grokQuotaUpstreamStep struct {
status int
body string
err error
}
type grokQuotaSequenceUpstream struct {
httpUpstreamRecorder
mu sync.Mutex
steps []grokQuotaUpstreamStep
requests []*http.Request
}
func (u *grokQuotaSequenceUpstream) Do(req *http.Request, _ string, _ int64, _ int) (*http.Response, error) {
u.mu.Lock()
defer u.mu.Unlock()
u.requests = append(u.requests, req)
index := len(u.requests) - 1
if index >= len(u.steps) {
return nil, errors.New("unexpected upstream request")
}
step := u.steps[index]
if step.err != nil {
return nil, step.err
}
return &http.Response{
StatusCode: step.status,
Header: make(http.Header),
Body: io.NopCloser(strings.NewReader(step.body)),
}, nil
}
func (u *grokQuotaSequenceUpstream) snapshotRequests() []*http.Request {
u.mu.Lock()
defer u.mu.Unlock()
return append([]*http.Request(nil), u.requests...)
}
func (u *grokHybridUpstream) Do(req *http.Request, _ string, _ int64, _ int) (*http.Response, error) {
var body []byte
if req != nil && req.Body != nil {
body, _ = io.ReadAll(req.Body)
}
u.mu.Lock()
u.requests = append(u.requests, req)
u.bodies = append(u.bodies, body)
u.mu.Unlock()
if req.URL.Path == "/v1/responses" {
status := u.activeStatus
if status == 0 {
status = http.StatusOK
}
headers := u.activeHeaders
if headers == nil {
headers = http.Header{
"X-Ratelimit-Limit-Tokens": []string{"2000000"},
"X-Ratelimit-Remaining-Tokens": []string{"1500000"},
}
}
return &http.Response{StatusCode: status, Header: headers, Body: io.NopCloser(strings.NewReader(`{"id":"resp_probe"}`))}, nil
}
if u.billingStarted != nil {
u.billingStartOnce.Do(func() { close(u.billingStarted) })
}
if u.billingRelease != nil {
select {
case <-u.billingRelease:
case <-req.Context().Done():
return nil, req.Context().Err()
}
}
billingStatus := u.billingStatus
if req.URL.RawQuery == "format=credits" && u.weeklyBillingStatus != 0 {
billingStatus = u.weeklyBillingStatus
}
if req.URL.RawQuery != "format=credits" && u.monthlyBillingStatus != 0 {
billingStatus = u.monthlyBillingStatus
}
if billingStatus != 0 && billingStatus != http.StatusOK {
return &http.Response{
StatusCode: billingStatus,
Header: u.billingHeaders,
Body: io.NopCloser(strings.NewReader(`{"error":{"message":"billing limited"}}`)),
}, nil
}
if req.URL.RawQuery == "format=credits" {
usage := ""
if u.weeklyUsagePercent != nil {
usage = `,"creditUsagePercent":` + strconv.FormatFloat(*u.weeklyUsagePercent, 'f', -1, 64)
}
payload := `{"config":{"currentPeriod":{"type":"WEEKLY","start":"2026-07-09T03:25:00Z","end":"2026-07-16T03:25:00Z"}` + usage + `}}`
return &http.Response{StatusCode: http.StatusOK, Header: make(http.Header), Body: io.NopCloser(strings.NewReader(payload))}, nil
}
monthlyLimit := ""
if u.monthlyLimitCents != nil {
monthlyLimit = `,"monthlyLimit":{"val":` + strconv.FormatFloat(*u.monthlyLimitCents, 'f', -1, 64) + `}`
}
monthlyPayload := `{"config":{"billingPeriodStart":"2026-07-01T00:00:00Z","billingPeriodEnd":"2026-08-01T00:00:00Z"` + monthlyLimit + `}}`
return &http.Response{
StatusCode: http.StatusOK,
Header: make(http.Header),
Body: io.NopCloser(strings.NewReader(monthlyPayload)),
}, nil
}
func (u *grokHybridUpstream) snapshot() ([]*http.Request, [][]byte) {
u.mu.Lock()
defer u.mu.Unlock()
requests := append([]*http.Request(nil), u.requests...)
bodies := make([][]byte, len(u.bodies))
for i := range u.bodies {
bodies[i] = append([]byte(nil), u.bodies[i]...)
}
return requests, bodies
}
func (r *grokQuotaProxyRepo) GetByID(_ context.Context, id int64) (*Proxy, error) {
r.calls++
return r.proxies[id], nil
}
func healthyGrokQuotaOAuthAccount(id int64) *Account {
return &Account{
ID: id,
Platform: PlatformGrok,
Type: AccountTypeOAuth,
Status: StatusActive,
Schedulable: true,
Concurrency: 1,
Credentials: map[string]any{
"access_token": "access-token",
"refresh_token": "refresh-token",
"expires_at": time.Now().Add(2 * grokTokenRefreshSkew).UTC().Format(time.RFC3339),
},
}
}
func TestGrokQuotaServiceFetchBillingRetries502ThenSucceeds(t *testing.T) {
account := healthyGrokQuotaOAuthAccount(401)
upstream := &grokQuotaSequenceUpstream{steps: []grokQuotaUpstreamStep{
{status: http.StatusBadGateway, body: `The origin web server returned an invalid or incomplete response to Cloudflare.`},
{status: http.StatusOK, body: `{"config":{"currentPeriod":{"type":"WEEKLY","start":"2026-07-09T03:25:00Z","end":"2026-07-16T03:25:00Z"},"creditUsagePercent":12}}`},
}}
svc := &GrokQuotaService{httpUpstream: upstream}
summary, status, err := svc.fetchBilling(context.Background(), account, "access-token", "", true)
require.NoError(t, err)
require.Equal(t, http.StatusOK, status)
require.NotNil(t, summary)
require.NotNil(t, summary.UsagePercent)
require.Equal(t, 12.0, *summary.UsagePercent)
requests := upstream.snapshotRequests()
require.Len(t, requests, 2)
require.Equal(t, http.MethodGet, requests[0].Method)
require.Equal(t, "/v1/billing", requests[0].URL.Path)
require.Equal(t, "format=credits", requests[0].URL.RawQuery)
}
func TestGrokQuotaServiceFetchBillingRetriesTransportErrorThenSucceeds(t *testing.T) {
account := healthyGrokQuotaOAuthAccount(402)
upstream := &grokQuotaSequenceUpstream{steps: []grokQuotaUpstreamStep{
{err: errors.New("temporary transport failure")},
{status: http.StatusOK, body: `{"config":{"currentPeriod":{"type":"WEEKLY"},"creditUsagePercent":8}}`},
}}
svc := &GrokQuotaService{httpUpstream: upstream}
summary, status, err := svc.fetchBilling(context.Background(), account, "access-token", "", true)
require.NoError(t, err)
require.Equal(t, http.StatusOK, status)
require.NotNil(t, summary)
require.Len(t, upstream.snapshotRequests(), 2)
}
func TestGrokQuotaServiceFetchBillingStopsAfterSingleTransientRetry(t *testing.T) {
account := healthyGrokQuotaOAuthAccount(403)
upstream := &grokQuotaSequenceUpstream{steps: []grokQuotaUpstreamStep{
{status: http.StatusBadGateway, body: `cloudflare failure`},
{status: http.StatusBadGateway, body: `cloudflare failure`},
{status: http.StatusOK, body: `{"config":{"currentPeriod":{"type":"WEEKLY"}}}`},
}}
svc := &GrokQuotaService{httpUpstream: upstream}
summary, status, err := svc.fetchBilling(context.Background(), account, "access-token", "", true)
require.Error(t, err)
require.Nil(t, summary)
require.Equal(t, http.StatusBadGateway, status)
require.Equal(t, "GROK_QUOTA_PROBE_UPSTREAM_ERROR", infraerrors.Reason(err))
require.Contains(t, infraerrors.Message(err), "billing returned 502: cloudflare failure")
require.Len(t, upstream.snapshotRequests(), 2)
}
func TestGrokQuotaServiceFetchBillingDoesNotRetryNonTransientStatuses(t *testing.T) {
tests := []struct {
name string
status int
wantErr bool
}{
{name: "unauthorized", status: http.StatusUnauthorized, wantErr: true},
{name: "forbidden", status: http.StatusForbidden, wantErr: true},
{name: "rate limited", status: http.StatusTooManyRequests, wantErr: false},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
account := healthyGrokQuotaOAuthAccount(404)
upstream := &grokQuotaSequenceUpstream{steps: []grokQuotaUpstreamStep{
{status: tt.status, body: `{"error":{"message":"rejected"}}`},
{status: http.StatusOK, body: `{"config":{"currentPeriod":{"type":"WEEKLY"}}}`},
}}
svc := &GrokQuotaService{httpUpstream: upstream}
summary, status, err := svc.fetchBilling(context.Background(), account, "access-token", "", true)
if tt.wantErr {
require.Error(t, err)
} else {
require.NoError(t, err)
}
require.Nil(t, summary)
require.Equal(t, tt.status, status)
require.Len(t, upstream.snapshotRequests(), 1)
})
}
}
func TestIsRetryableGrokBillingStatus(t *testing.T) {
tests := []struct {
status int
want bool
}{
{status: http.StatusBadGateway, want: true},
{status: http.StatusServiceUnavailable, want: true},
{status: http.StatusGatewayTimeout, want: true},
{status: http.StatusUnauthorized, want: false},
{status: http.StatusForbidden, want: false},
{status: http.StatusTooManyRequests, want: false},
}
for _, tt := range tests {
t.Run(strconv.Itoa(tt.status), func(t *testing.T) {
require.Equal(t, tt.want, isRetryableGrokBillingStatus(tt.status))
})
}
}
func TestGrokQuotaServiceProbeUsageDoesNotRetryResponsesPost(t *testing.T) {
account := healthyGrokQuotaOAuthAccount(405)
repo := &grokQuotaAccountRepo{mockAccountRepoForPlatform: &mockAccountRepoForPlatform{
accountsByID: map[int64]*Account{account.ID: account},
}}
upstream := &grokQuotaSequenceUpstream{steps: []grokQuotaUpstreamStep{
{status: http.StatusBadGateway, body: `cloudflare failure`},
{status: http.StatusOK, body: `{"id":"unexpected_retry"}`},
}}
svc := NewGrokQuotaService(repo, nil, NewGrokTokenProvider(repo, nil), upstream, nil)
result, err := svc.ProbeUsage(context.Background(), account.ID)
require.Error(t, err)
require.Nil(t, result)
require.Equal(t, "GROK_QUOTA_PROBE_UPSTREAM_ERROR", infraerrors.Reason(err))
requests := upstream.snapshotRequests()
require.Len(t, requests, 1)
require.Equal(t, http.MethodPost, requests[0].Method)
require.Equal(t, "/v1/responses", requests[0].URL.Path)
}
func TestGrokQuotaServiceProbeUsageStoresHeaders(t *testing.T) {
t.Parallel()
account := healthyGrokQuotaOAuthAccount(42)
repo := &grokQuotaAccountRepo{
mockAccountRepoForPlatform: &mockAccountRepoForPlatform{
accountsByID: map[int64]*Account{42: account},
},
}
upstream := &httpUpstreamRecorder{resp: &http.Response{
StatusCode: http.StatusOK,
Header: http.Header{
"X-Ratelimit-Limit-Requests": []string{"10"},
"X-Ratelimit-Remaining-Requests": []string{"7"},
"X-Ratelimit-Reset-Requests": []string{"2000000000"},
"X-Ratelimit-Limit-Tokens": []string{"1000"},
"X-Ratelimit-Remaining-Tokens": []string{"900"},
},
Body: io.NopCloser(strings.NewReader(`{"id":"resp_probe"}`)),
}}
svc := NewGrokQuotaService(repo, nil, NewGrokTokenProvider(repo, nil), upstream, nil)
result, err := svc.ProbeUsage(context.Background(), 42)
require.NoError(t, err)
require.Equal(t, http.StatusOK, result.StatusCode)
require.Equal(t, "grok-4.5", result.Model)
require.True(t, result.HeadersObserved)
require.NotNil(t, result.Snapshot)
require.True(t, result.Snapshot.HeadersObserved)
require.Equal(t, "active_probe", result.Snapshot.ObservationSource)
require.NotEmpty(t, result.Snapshot.LastProbeAt)
require.NotEmpty(t, result.Snapshot.LastHeadersSeenAt)
require.NotNil(t, result.Snapshot.Requests)
require.EqualValues(t, 10, *result.Snapshot.Requests.Limit)
require.EqualValues(t, 7, *result.Snapshot.Requests.Remaining)
require.Equal(t, "https://cli-chat-proxy.grok.com/v1/responses", upstream.lastReq.URL.String())
require.Equal(t, "Bearer access-token", upstream.lastReq.Header.Get("Authorization"))
require.Equal(t, grokCLIVersion, upstream.lastReq.Header.Get("X-Grok-Client-Version"))
require.Equal(t, "application/json, text/event-stream", upstream.lastReq.Header.Get("Accept"))
require.Equal(t, "grok-4.5", gjson.GetBytes(upstream.lastBody, "model").String())
require.Equal(t, grokQuotaProbeInput, gjson.GetBytes(upstream.lastBody, "input").String())
require.True(t, gjson.GetBytes(upstream.lastBody, "stream").Bool())
require.False(t, gjson.GetBytes(upstream.lastBody, "max_output_tokens").Exists())
require.False(t, gjson.GetBytes(upstream.lastBody, "store").Exists())
require.NotNil(t, repo.updates[42][grokQuotaSnapshotExtraKey])
}
func TestGrokQuotaServiceProbeUsageIgnoresAccountGrokMapping(t *testing.T) {
t.Parallel()
account := healthyGrokQuotaOAuthAccount(47)
account.Credentials["model_mapping"] = map[string]any{
"grok": "grok-composer",
"grok-composer": "grok-composer-2.5-fast",
}
repo := &grokQuotaAccountRepo{
mockAccountRepoForPlatform: &mockAccountRepoForPlatform{
accountsByID: map[int64]*Account{47: account},
},
}
upstream := &httpUpstreamRecorder{resp: &http.Response{
StatusCode: http.StatusOK,
Header: http.Header{},
Body: io.NopCloser(strings.NewReader(`{"id":"resp_probe"}`)),
}}
svc := NewGrokQuotaService(repo, nil, NewGrokTokenProvider(repo, nil), upstream, nil)
result, err := svc.ProbeUsage(context.Background(), 47)
require.NoError(t, err)
require.Equal(t, "grok-4.5", result.Model)
require.Equal(t, "grok-4.5", gjson.GetBytes(upstream.lastBody, "model").String())
require.NotContains(t, string(upstream.lastBody), "grok-composer")
}
func TestGrokQuotaServiceProbeUsageReportsProbeModelOnUpstreamError(t *testing.T) {
t.Parallel()
account := healthyGrokQuotaOAuthAccount(48)
repo := &grokQuotaAccountRepo{
mockAccountRepoForPlatform: &mockAccountRepoForPlatform{
accountsByID: map[int64]*Account{48: account},
},
}
upstream := &httpUpstreamRecorder{resp: &http.Response{
StatusCode: http.StatusBadRequest,
Header: http.Header{},
Body: io.NopCloser(strings.NewReader(`{"code":"invalid-argument","error":"Model not found"}`)),
}}
svc := NewGrokQuotaService(repo, nil, NewGrokTokenProvider(repo, nil), upstream, nil)
_, err := svc.ProbeUsage(context.Background(), 48)
require.Error(t, err)
require.Equal(t, "GROK_QUOTA_PROBE_UPSTREAM_ERROR", infraerrors.Reason(err))
require.Contains(t, infraerrors.Message(err), `probe model "grok-4.5"`)
}
func TestGrokQuotaServiceProbeUsageRedactsUpstreamErrorBodyFromErrorAndLogs(t *testing.T) {
const upstreamSecret = "upstream-secret-refresh-token"
account := healthyGrokQuotaOAuthAccount(49)
repo := &grokQuotaAccountRepo{
mockAccountRepoForPlatform: &mockAccountRepoForPlatform{
accountsByID: map[int64]*Account{49: account},
},
}
upstream := &httpUpstreamRecorder{resp: &http.Response{
StatusCode: http.StatusBadRequest,
Header: http.Header{},
Body: io.NopCloser(strings.NewReader(
`{"error":"` + upstreamSecret + `","detail":"credential rejected"}`,
)),
}}
svc := NewGrokQuotaService(
repo,
nil,
NewGrokTokenProvider(repo, nil),
upstream,
nil,
)
var logs bytes.Buffer
previousLogger := slog.Default()
slog.SetDefault(slog.New(slog.NewTextHandler(&logs, nil)))
defer slog.SetDefault(previousLogger)
_, err := svc.ProbeUsage(context.Background(), account.ID)
require.Error(t, err)
require.Equal(t, "GROK_QUOTA_PROBE_UPSTREAM_ERROR", infraerrors.Reason(err))
require.Contains(t, infraerrors.Message(err), `probe model "grok-4.5"`)
require.NotContains(t, err.Error(), upstreamSecret)
require.NotContains(t, infraerrors.Message(err), upstreamSecret)
require.Contains(t, logs.String(), "GROK_QUOTA_PROBE_UPSTREAM_ERROR")
require.NotContains(t, logs.String(), upstreamSecret)
require.NotContains(t, logs.String(), "credential rejected")
require.Equal(t, xai.DefaultCLIBaseURL+"/responses", upstream.lastReq.URL.String())
}
func TestGrokQuotaServiceProbeUsageLoadsProxyWhenAccountEdgeMissing(t *testing.T) {
t.Parallel()
proxyID := int64(7)
account := healthyGrokQuotaOAuthAccount(46)
account.ProxyID = &proxyID
repo := &grokQuotaAccountRepo{
mockAccountRepoForPlatform: &mockAccountRepoForPlatform{
accountsByID: map[int64]*Account{46: account},
},
}
proxyRepo := &grokQuotaProxyRepo{
proxies: map[int64]*Proxy{
proxyID: {
ID: proxyID,
Protocol: "http",
Host: "proxy.test",
Port: 3128,
},
},
}
upstream := &httpUpstreamRecorder{resp: &http.Response{
StatusCode: http.StatusOK,
Header: http.Header{},
Body: io.NopCloser(strings.NewReader(`{"id":"resp_probe"}`)),
}}
svc := NewGrokQuotaService(repo, proxyRepo, NewGrokTokenProvider(repo, nil), upstream, nil)
_, err := svc.ProbeUsage(context.Background(), 46)
require.NoError(t, err)
require.Equal(t, 1, proxyRepo.calls)
require.Equal(t, "http://proxy.test:3128", upstream.lastProxyURL)
}
func TestGrokQuotaServiceProbeUsageStoresNoHeadersState(t *testing.T) {
t.Parallel()
account := healthyGrokQuotaOAuthAccount(45)
observedResetAt := time.Now().Add(-time.Second).UTC().Truncate(time.Second)
observedLimitedAt := observedResetAt.Add(-grokRateLimitRepeatCooldown)
account.RateLimitedAt = &observedLimitedAt
account.RateLimitResetAt = &observedResetAt
repo := &grokQuotaAccountRepo{
mockAccountRepoForPlatform: &mockAccountRepoForPlatform{
accountsByID: map[int64]*Account{45: account},
},
recoveryClearResult: true,
}
upstream := &httpUpstreamRecorder{resp: &http.Response{
StatusCode: http.StatusOK,
Header: http.Header{},
Body: io.NopCloser(strings.NewReader(`{"id":"resp_probe"}`)),
}}
svc := NewGrokQuotaService(repo, nil, NewGrokTokenProvider(repo, nil), upstream, nil)
result, err := svc.ProbeUsage(context.Background(), 45)
require.NoError(t, err)
require.Equal(t, http.StatusOK, result.StatusCode)
require.False(t, result.HeadersObserved)
require.NotNil(t, result.Snapshot)
require.False(t, result.Snapshot.HeadersObserved)
require.Equal(t, "active_probe", result.Snapshot.ObservationSource)
require.NotEmpty(t, result.Snapshot.LastProbeAt)
require.Empty(t, result.Snapshot.LastHeadersSeenAt)
stored, ok := repo.updates[45][grokQuotaSnapshotExtraKey].(*xai.QuotaSnapshot)
require.True(t, ok)
require.False(t, stored.HeadersObserved)
require.Equal(t, http.StatusOK, stored.StatusCode)
require.Equal(t, 1, repo.recoveryClearCalls)
require.Equal(t, observedLimitedAt, repo.recoveryObservedAt)
require.Equal(t, observedResetAt, repo.recoveryObservedReset)
}
func TestGrokQuotaServiceProbeUsageDoesNotOverwriteSnapshotOnUnauthorized(t *testing.T) {
t.Parallel()
account := healthyGrokQuotaOAuthAccount(44)
previous := &xai.QuotaSnapshot{StatusCode: http.StatusOK, HeadersObserved: true}
account.Extra = map[string]any{grokQuotaSnapshotExtraKey: previous}
repo := &grokQuotaAccountRepo{mockAccountRepoForPlatform: &mockAccountRepoForPlatform{
accountsByID: map[int64]*Account{account.ID: account},
}}
upstream := &httpUpstreamRecorder{resp: &http.Response{
StatusCode: http.StatusUnauthorized,
Header: http.Header{},
Body: io.NopCloser(strings.NewReader(`{"error":"unauthorized"}`)),
}}
svc := NewGrokQuotaService(repo, nil, NewGrokTokenProvider(repo, nil), upstream, nil)
_, err := svc.ProbeUsage(context.Background(), account.ID)
require.Error(t, err)
require.Equal(t, 0, repo.updateCalls)
require.Same(t, previous, account.Extra[grokQuotaSnapshotExtraKey])
}
func TestGrokQuotaServiceProbeUsageReturnsRateLimitedSnapshot(t *testing.T) {
t.Parallel()
account := healthyGrokQuotaOAuthAccount(43)
repo := &grokQuotaAccountRepo{
mockAccountRepoForPlatform: &mockAccountRepoForPlatform{
accountsByID: map[int64]*Account{43: account},
},
}
upstream := &httpUpstreamRecorder{resp: &http.Response{
StatusCode: http.StatusTooManyRequests,
Header: http.Header{"Retry-After": []string{"45"}},
Body: io.NopCloser(strings.NewReader(`{"error":{"message":"rate limited"}}`)),
}}
svc := NewGrokQuotaService(repo, nil, NewGrokTokenProvider(repo, nil), upstream, nil)
result, err := svc.ProbeUsage(context.Background(), 43)
require.NoError(t, err)
require.Equal(t, http.StatusTooManyRequests, result.StatusCode)
require.NotNil(t, result.Snapshot)
require.NotNil(t, result.Snapshot.RetryAfterSeconds)
require.Equal(t, 45, *result.Snapshot.RetryAfterSeconds)
require.Equal(t, 1, repo.rateLimitedCalls)
require.Equal(t, account.ID, repo.lastRateLimitedID)
require.WithinDuration(t, time.Now().Add(45*time.Second), repo.lastRateLimitResetAt, time.Second)
require.Zero(t, repo.tempUnschedCalls)
}
func TestGrokQuotaServiceQueryQuotaFreeFallsBackToGrok45(t *testing.T) {
t.Parallel()
account := healthyGrokQuotaOAuthAccount(51)
repo := &grokQuotaAccountRepo{mockAccountRepoForPlatform: &mockAccountRepoForPlatform{
accountsByID: map[int64]*Account{account.ID: account},
}}
upstream := &grokHybridUpstream{}
usageRepo := &grokQuotaUsageLogRepo{stats: &usagestats.AccountStats{Tokens: 1_000_000}}
svc := NewGrokQuotaService(repo, nil, NewGrokTokenProvider(repo, nil), upstream, nil, usageRepo)
result, err := svc.QueryQuota(context.Background(), account.ID)
require.NoError(t, err)
require.Equal(t, "hybrid_probe", result.Source)
require.Equal(t, "grok-4.5", result.Model)
require.NotNil(t, result.Billing)
require.Nil(t, result.Billing.UsagePercent)
require.NotNil(t, result.LocalUsage24h)
require.EqualValues(t, 1_000_000, result.LocalUsage24h.Tokens)
require.Equal(t, 1, usageRepo.calls)
require.WithinDuration(t, time.Now().UTC().Add(-24*time.Hour), usageRepo.startTimes[0], time.Second)
require.NotNil(t, result.Snapshot)
require.NotNil(t, result.Snapshot.Tokens)
require.EqualValues(t, 2_000_000, *result.Snapshot.Tokens.Limit)
require.True(t, result.HeadersObserved)
requests, bodies := upstream.snapshot()
require.Len(t, requests, 3)
responseCalls := 0
for i, req := range requests {
if req.URL.Path != "/v1/responses" {
continue
}
responseCalls++
require.Equal(t, http.MethodPost, req.Method)
require.Equal(t, "application/json, text/event-stream", req.Header.Get("Accept"))
require.Equal(t, "grok-4.5", gjson.GetBytes(bodies[i], "model").String())
require.Equal(t, grokQuotaProbeInput, gjson.GetBytes(bodies[i], "input").String())
require.True(t, gjson.GetBytes(bodies[i], "stream").Bool())
require.False(t, gjson.GetBytes(bodies[i], "max_output_tokens").Exists())
require.False(t, gjson.GetBytes(bodies[i], "store").Exists())
}
require.Equal(t, 1, responseCalls)
}
func TestGrokQuotaServiceQueryQuotaPaidBillingSkipsActiveProbe(t *testing.T) {
t.Parallel()
account := healthyGrokQuotaOAuthAccount(52)
repo := &grokQuotaAccountRepo{mockAccountRepoForPlatform: &mockAccountRepoForPlatform{
accountsByID: map[int64]*Account{account.ID: account},
}}
usagePercent := 25.0
upstream := &grokHybridUpstream{weeklyUsagePercent: &usagePercent}
usageRepo := &grokQuotaUsageLogRepo{stats: &usagestats.AccountStats{Tokens: 1_000_000}}
svc := NewGrokQuotaService(repo, nil, NewGrokTokenProvider(repo, nil), upstream, nil, usageRepo)
result, err := svc.QueryQuota(context.Background(), account.ID)
require.NoError(t, err)
require.Equal(t, "billing_probe", result.Source)
require.NotNil(t, result.Billing)
require.InDelta(t, usagePercent, *result.Billing.UsagePercent, 1e-9)
require.Nil(t, result.Snapshot)
require.Empty(t, result.Model)
require.Nil(t, result.LocalUsage24h)
requests, _ := upstream.snapshot()
require.Len(t, requests, 2)
for _, req := range requests {
require.Equal(t, "/v1/billing", req.URL.Path)
}
}
func TestGrokQuotaServiceQueryQuotaCustomPaidMonthlyLimitSkipsActiveProbe(t *testing.T) {
t.Parallel()
account := healthyGrokQuotaOAuthAccount(57)
repo := &grokQuotaAccountRepo{mockAccountRepoForPlatform: &mockAccountRepoForPlatform{
accountsByID: map[int64]*Account{account.ID: account},
}}
monthlyLimit := 25_000.0
upstream := &grokHybridUpstream{monthlyLimitCents: &monthlyLimit}
svc := NewGrokQuotaService(repo, nil, NewGrokTokenProvider(repo, nil), upstream, nil)
result, err := svc.QueryQuota(context.Background(), account.ID)
require.NoError(t, err)
require.Equal(t, "billing_probe", result.Source)
require.NotNil(t, result.Billing)
require.InDelta(t, monthlyLimit, *result.Billing.MonthlyLimitCents, 1e-9)
require.Nil(t, result.Snapshot)
requests, _ := upstream.snapshot()
require.Len(t, requests, 2)
for _, req := range requests {
require.Equal(t, "/v1/billing", req.URL.Path)
}
}
func TestGrokLocalUsage24hUsesRollingUTCWindow(t *testing.T) {
t.Parallel()
now := time.Date(2026, 7, 14, 20, 30, 0, 0, time.FixedZone("UTC+8", 8*60*60))
t.Run("returns usage from exact rolling window", func(t *testing.T) {
repo := &grokQuotaUsageLogRepo{stats: &usagestats.AccountStats{Tokens: 1_250_000}}
stats := grokLocalUsage24h(context.Background(), repo, 57, now)
require.NotNil(t, stats)
require.EqualValues(t, 1_250_000, stats.Tokens)
require.Equal(t, []time.Time{now.UTC().Add(-24 * time.Hour)}, repo.startTimes)
})
t.Run("query failure returns no stats", func(t *testing.T) {
repo := &grokQuotaUsageLogRepo{err: context.DeadlineExceeded}
stats := grokLocalUsage24h(context.Background(), repo, 57, now)
require.Nil(t, stats)
require.Equal(t, []time.Time{now.UTC().Add(-24 * time.Hour)}, repo.startTimes)
})
t.Run("missing repository returns no stats", func(t *testing.T) {
require.Nil(t, grokLocalUsage24h(context.Background(), nil, 57, now))
})
t.Run("invalid account returns no stats without query", func(t *testing.T) {
repo := &grokQuotaUsageLogRepo{}
require.Nil(t, grokLocalUsage24h(context.Background(), repo, 0, now))
require.Zero(t, repo.calls)
})
}
func TestGrokLocalUsageForQuotaSelectsFreeOrPaidWindows(t *testing.T) {
t.Parallel()
now := time.Date(2026, 7, 14, 12, 0, 0, 0, time.UTC)
billing := &xai.BillingSummary{
PeriodType: "weekly",
PeriodStart: now.Add(-4 * 24 * time.Hour).Format(time.RFC3339),
PeriodEnd: now.Add(3 * 24 * time.Hour).Format(time.RFC3339),
BillingPeriodStart: now.Add(-13 * 24 * time.Hour).Format(time.RFC3339),
BillingPeriodEnd: now.Add(17 * 24 * time.Hour).Format(time.RFC3339),
}
t.Run("free queries only rolling 24h", func(t *testing.T) {
repo := &grokQuotaUsageLogRepo{stats: &usagestats.AccountStats{Tokens: 500_000}}
rolling, weekly, monthly := grokLocalUsageForQuota(context.Background(), repo, 57, billing, now)
require.NotNil(t, rolling)
require.Nil(t, weekly)
require.Nil(t, monthly)
require.Equal(t, []time.Time{now.Add(-24 * time.Hour)}, repo.startTimes)
})
t.Run("paid queries only billing windows", func(t *testing.T) {
usagePercent := 25.0
paidBilling := *billing
paidBilling.UsagePercent = &usagePercent
repo := &grokQuotaUsageLogRepo{stats: &usagestats.AccountStats{Tokens: 500_000}}
rolling, weekly, monthly := grokLocalUsageForQuota(context.Background(), repo, 57, &paidBilling, now)
require.Nil(t, rolling)
require.NotNil(t, weekly)
require.NotNil(t, monthly)
require.Equal(t, []time.Time{
now.Add(-4 * 24 * time.Hour),
now.Add(-13 * 24 * time.Hour),
}, repo.startTimes)
})
}
func TestGrokLocalUsageForBillingOnlyReturnsAvailableWindows(t *testing.T) {
t.Parallel()
now := time.Date(2026, 7, 13, 12, 0, 0, 0, time.UTC)
billing := &xai.BillingSummary{
PeriodType: "weekly",
PeriodStart: now.Add(-4 * 24 * time.Hour).Format(time.RFC3339),
PeriodEnd: now.Add(3 * 24 * time.Hour).Format(time.RFC3339),
}
t.Run("valid weekly window", func(t *testing.T) {
repo := &grokQuotaUsageLogRepo{stats: &usagestats.AccountStats{Tokens: 1_500_000}}
weekly, monthly := grokLocalUsageForBilling(context.Background(), repo, 57, billing, now)
require.NotNil(t, weekly)
require.EqualValues(t, 1_500_000, weekly.Tokens)
require.Nil(t, monthly)
require.Equal(t, 1, repo.calls)
})
t.Run("query failure", func(t *testing.T) {
repo := &grokQuotaUsageLogRepo{err: context.DeadlineExceeded}
weekly, monthly := grokLocalUsageForBilling(context.Background(), repo, 57, billing, now)
require.Nil(t, weekly)
require.Nil(t, monthly)
require.Equal(t, 1, repo.calls)
})
t.Run("missing billing window", func(t *testing.T) {
repo := &grokQuotaUsageLogRepo{}
weekly, monthly := grokLocalUsageForBilling(context.Background(), repo, 57, nil, now)
require.Nil(t, weekly)
require.Nil(t, monthly)
require.Zero(t, repo.calls)
})
}
func TestAccountUsageServiceGrokRefreshUsesBillingOnly(t *testing.T) {
t.Parallel()
account := healthyGrokQuotaOAuthAccount(54)
repo := &grokQuotaAccountRepo{mockAccountRepoForPlatform: &mockAccountRepoForPlatform{
accountsByID: map[int64]*Account{account.ID: account},
}}
upstream := &grokHybridUpstream{}
usageRepo := &grokQuotaUsageLogRepo{stats: &usagestats.AccountStats{Tokens: 750_000}}
quotaService := NewGrokQuotaService(repo, nil, NewGrokTokenProvider(repo, nil), upstream, nil, usageRepo)
usageService := &AccountUsageService{
grokQuotaFetcher: NewGrokQuotaFetcher(),
grokQuotaService: quotaService,
usageLogRepo: usageRepo,
cache: NewUsageCache(),
}
usage, err := usageService.getGrokUsage(context.Background(), account, false)
require.NoError(t, err)
require.NotNil(t, usage.GrokBilling)
require.Nil(t, usage.GrokBilling.UsagePercent)
require.NotNil(t, usage.GrokLocalUsage24h)
require.EqualValues(t, 750_000, usage.GrokLocalUsage24h.Tokens)
require.Equal(t, 1, usageRepo.calls)
require.Len(t, usageRepo.startTimes, 1)
require.WithinDuration(t, time.Now().UTC().Add(-24*time.Hour), usageRepo.startTimes[0], time.Second)
requests, _ := upstream.snapshot()
require.Len(t, requests, 2)
for _, req := range requests {
require.Equal(t, http.MethodGet, req.Method)
require.Equal(t, "/v1/billing", req.URL.Path)
}
}
func TestGrokQuotaServiceProbeFlightsDeduplicateBillingAndSeparateActive(t *testing.T) {
t.Parallel()
account := healthyGrokQuotaOAuthAccount(55)
repo := &grokQuotaAccountRepo{mockAccountRepoForPlatform: &mockAccountRepoForPlatform{
accountsByID: map[int64]*Account{account.ID: account},
}}
billingStarted := make(chan struct{})
billingRelease := make(chan struct{})
upstream := &grokHybridUpstream{billingStarted: billingStarted, billingRelease: billingRelease}
svc := NewGrokQuotaService(repo, nil, NewGrokTokenProvider(repo, nil), upstream, nil)
type probeOutcome struct {
result *GrokQuotaProbeResult
err error
}
billingOutcomes := make(chan probeOutcome, 2)
go func() {
result, err := svc.ProbeBilling(context.Background(), account.ID)
billingOutcomes <- probeOutcome{result: result, err: err}
}()
<-billingStarted
secondStarted := make(chan struct{})
go func() {
close(secondStarted)
result, err := svc.ProbeBilling(context.Background(), account.ID)
billingOutcomes <- probeOutcome{result: result, err: err}
}()
<-secondStarted
time.Sleep(25 * time.Millisecond)
activeResult, err := svc.ProbeUsage(context.Background(), account.ID)
require.NoError(t, err)
require.NotNil(t, activeResult.Snapshot)
close(billingRelease)
for range 2 {
outcome := <-billingOutcomes
require.NoError(t, outcome.err)
require.NotNil(t, outcome.result.Billing)
}
requests, _ := upstream.snapshot()
billingCalls := 0
activeCalls := 0
for _, req := range requests {
switch req.URL.Path {
case "/v1/billing":
billingCalls++
case "/v1/responses":
activeCalls++
}
}
require.Equal(t, 2, billingCalls)
require.Equal(t, 1, activeCalls)
}
func TestGrokQuotaServiceBilling429DoesNotPauseModelScheduling(t *testing.T) {
t.Parallel()
account := healthyGrokQuotaOAuthAccount(56)
repo := &grokQuotaAccountRepo{mockAccountRepoForPlatform: &mockAccountRepoForPlatform{
accountsByID: map[int64]*Account{account.ID: account},
}}
upstream := &grokHybridUpstream{
billingStatus: http.StatusTooManyRequests,
billingHeaders: http.Header{"Retry-After": []string{"45"}},
}
svc := NewGrokQuotaService(repo, nil, NewGrokTokenProvider(repo, nil), upstream, nil)
result, err := svc.ProbeBilling(context.Background(), account.ID)
require.Error(t, err)
require.Nil(t, result)
require.Zero(t, repo.rateLimitedCalls)
}
func TestGrokQuotaServiceBilling403PersistsMediaEligibilitySignal(t *testing.T) {
t.Parallel()
account := healthyGrokQuotaOAuthAccount(58)
repo := &grokQuotaAccountRepo{mockAccountRepoForPlatform: &mockAccountRepoForPlatform{
accountsByID: map[int64]*Account{account.ID: account},
}}
upstream := &grokHybridUpstream{billingStatus: http.StatusForbidden}
svc := NewGrokQuotaService(repo, nil, NewGrokTokenProvider(repo, nil), upstream, nil)
result, err := svc.ProbeBilling(context.Background(), account.ID)
require.Error(t, err)
require.Nil(t, result)
require.Equal(t, 1, repo.updateCalls)
raw := repo.updates[account.ID][grokBillingExtraKey]
billing, ok := raw.(*xai.BillingSummary)
require.True(t, ok)
require.Equal(t, http.StatusForbidden, billing.StatusCode)
require.Equal(t, http.StatusForbidden, billing.WeeklyStatusCode)
require.Equal(t, http.StatusForbidden, billing.MonthlyStatusCode)
require.True(t, billing.Partial)
account.Extra = map[string]any{grokBillingExtraKey: billing}
eligible, reason := account.GrokMediaGenerationEligibility()
require.False(t, eligible)
require.Equal(t, "billing_forbidden", reason)
}
func TestGrokQuotaServicePartialBilling403PersistsMediaEligibilitySignal(t *testing.T) {
t.Parallel()
account := healthyGrokQuotaOAuthAccount(59)
repo := &grokQuotaAccountRepo{mockAccountRepoForPlatform: &mockAccountRepoForPlatform{
accountsByID: map[int64]*Account{account.ID: account},
}}
upstream := &grokHybridUpstream{
weeklyBillingStatus: http.StatusForbidden,
monthlyBillingStatus: http.StatusOK,
}
svc := NewGrokQuotaService(repo, nil, NewGrokTokenProvider(repo, nil), upstream, nil)
result, err := svc.ProbeBilling(context.Background(), account.ID)
require.NoError(t, err)
require.NotNil(t, result)
require.NotNil(t, result.Billing)
require.Equal(t, http.StatusOK, result.StatusCode)
require.Equal(t, http.StatusForbidden, result.Billing.WeeklyStatusCode)
require.Equal(t, http.StatusOK, result.Billing.MonthlyStatusCode)
require.True(t, result.Billing.Partial)
require.Contains(t, result.Billing.FailedWindows, "weekly")
require.Equal(t, 1, repo.updateCalls)
account.Extra = map[string]any{grokBillingExtraKey: result.Billing}
eligible, reason := account.GrokMediaGenerationEligibility()
require.False(t, eligible)
require.Equal(t, "billing_forbidden", reason)
}
func TestGrokQuotaServiceProbeMediaEligibility(t *testing.T) {
t.Run("positive paid evidence enables media", func(t *testing.T) {
usagePercent := 10.0
monthlyLimit := 15_000.0
account := healthyGrokQuotaOAuthAccount(60)
repo := &grokQuotaAccountRepo{mockAccountRepoForPlatform: &mockAccountRepoForPlatform{
accountsByID: map[int64]*Account{account.ID: account},
}}
upstream := &grokHybridUpstream{weeklyUsagePercent: &usagePercent, monthlyLimitCents: &monthlyLimit}
svc := NewGrokQuotaService(repo, nil, NewGrokTokenProvider(repo, nil), upstream, nil)
eligible, reason, err := svc.ProbeMediaEligibility(context.Background(), account.ID)
require.NoError(t, err)
require.True(t, eligible)
require.Equal(t, "eligible", reason)
})
t.Run("successful empty billing identifies free account", func(t *testing.T) {
account := healthyGrokQuotaOAuthAccount(61)
repo := &grokQuotaAccountRepo{mockAccountRepoForPlatform: &mockAccountRepoForPlatform{
accountsByID: map[int64]*Account{account.ID: account},
}}
svc := NewGrokQuotaService(repo, nil, NewGrokTokenProvider(repo, nil), &grokHybridUpstream{}, nil)
eligible, reason, err := svc.ProbeMediaEligibility(context.Background(), account.ID)
require.NoError(t, err)
require.False(t, eligible)
require.Equal(t, "billing_free_tier", reason)
})
t.Run("forbidden billing is deterministic ineligibility", func(t *testing.T) {
account := healthyGrokQuotaOAuthAccount(62)
repo := &grokQuotaAccountRepo{mockAccountRepoForPlatform: &mockAccountRepoForPlatform{
accountsByID: map[int64]*Account{account.ID: account},
}}
svc := NewGrokQuotaService(repo, nil, NewGrokTokenProvider(repo, nil), &grokHybridUpstream{billingStatus: http.StatusForbidden}, nil)
eligible, reason, err := svc.ProbeMediaEligibility(context.Background(), account.ID)
require.NoError(t, err)
require.False(t, eligible)
require.Equal(t, "billing_forbidden", reason)
})
}
func TestPreferBillingObservationStatus(t *testing.T) {
t.Parallel()
tests := []struct {
name string
weeklyStatus int
monthlyStatus int
want int
}{
{name: "weekly forbidden wins", weeklyStatus: http.StatusForbidden, monthlyStatus: http.StatusBadGateway, want: http.StatusForbidden},
{name: "monthly forbidden wins", weeklyStatus: http.StatusBadGateway, monthlyStatus: http.StatusForbidden, want: http.StatusForbidden},
{name: "weekly observation otherwise wins", weeklyStatus: http.StatusTooManyRequests, monthlyStatus: http.StatusBadGateway, want: http.StatusTooManyRequests},
{name: "monthly observation is fallback", monthlyStatus: http.StatusBadGateway, want: http.StatusBadGateway},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
require.Equal(t, tt.want, preferBillingObservationStatus(tt.weeklyStatus, tt.monthlyStatus))
})
}
}
func TestGrokQuotaServiceQueryQuotaFree429PersistsLimitAndKeepsBilling(t *testing.T) {
t.Parallel()
account := healthyGrokQuotaOAuthAccount(53)
repo := &grokQuotaAccountRepo{mockAccountRepoForPlatform: &mockAccountRepoForPlatform{
accountsByID: map[int64]*Account{account.ID: account},
}}
upstream := &grokHybridUpstream{
activeStatus: http.StatusTooManyRequests,
activeHeaders: http.Header{"Retry-After": []string{"45"}},
}
svc := NewGrokQuotaService(repo, nil, NewGrokTokenProvider(repo, nil), upstream, nil)
result, err := svc.QueryQuota(context.Background(), account.ID)
require.NoError(t, err)
require.Equal(t, http.StatusTooManyRequests, result.StatusCode)
require.NotNil(t, result.Billing)
require.NotNil(t, result.Snapshot)
require.Equal(t, 45, *result.Snapshot.RetryAfterSeconds)
require.Equal(t, 1, repo.rateLimitedCalls)
require.Equal(t, account.ID, repo.lastRateLimitedID)
require.WithinDuration(t, time.Now().Add(45*time.Second), repo.lastRateLimitResetAt, time.Second)
}
func TestGrokQuotaServiceResetQuotaUnsupported(t *testing.T) {
t.Parallel()
account := &Account{
ID: 44,
Platform: PlatformGrok,
Type: AccountTypeOAuth,
}
repo := &grokQuotaAccountRepo{
mockAccountRepoForPlatform: &mockAccountRepoForPlatform{
accountsByID: map[int64]*Account{44: account},
},
}
svc := NewGrokQuotaService(repo, nil, nil, nil, nil)
_, err := svc.ResetQuota(context.Background(), 44)
require.Error(t, err)
require.Equal(t, http.StatusNotImplemented, infraerrors.Code(err))
require.Equal(t, "GROK_QUOTA_RESET_UNSUPPORTED", infraerrors.Reason(err))
}
func TestShouldAutoPauseGrokAccountByQuota(t *testing.T) {
t.Parallel()
zero := int64(0)
limit := int64(10)
resetFuture := time.Now().Add(time.Minute).Unix()
retryAfter := 30
tests := []struct {
name string
snapshot xai.QuotaSnapshot
want bool
}{
{
name: "remaining requests exhausted",
snapshot: xai.QuotaSnapshot{
Requests: &xai.QuotaWindow{Limit: &limit, Remaining: &zero, ResetUnix: &resetFuture},
UpdatedAt: time.Now().UTC().Format(time.RFC3339),
},
want: true,
},
{
name: "retry after active",
snapshot: xai.QuotaSnapshot{
RetryAfterSeconds: &retryAfter,
UpdatedAt: time.Now().UTC().Format(time.RFC3339),
},
want: true,
},
{
name: "retry after expired",
snapshot: xai.QuotaSnapshot{
RetryAfterSeconds: &retryAfter,
UpdatedAt: time.Now().Add(-time.Duration(retryAfter+1) * time.Second).UTC().Format(time.RFC3339),
},
want: false,
},
{
name: "stale snapshot ignored",
snapshot: xai.QuotaSnapshot{
Requests: &xai.QuotaWindow{Limit: &limit, Remaining: &zero, ResetUnix: &resetFuture},
UpdatedAt: time.Now().Add(-3 * time.Hour).UTC().Format(time.RFC3339),
},
want: false,
},
}
for _, tt := range tests {
tt := tt
t.Run(tt.name, func(t *testing.T) {
t.Parallel()
account := &Account{
Platform: PlatformGrok,
Type: AccountTypeOAuth,
Extra: map[string]any{
grokQuotaSnapshotExtraKey: tt.snapshot,
},
}
got, _ := shouldAutoPauseGrokAccountByQuota(account)
require.Equal(t, tt.want, got)
})
}
}