test: harden grok quota readiness

This commit is contained in:
Heatherm Huang
2026-06-26 10:42:21 +08:00
parent 2a80495880
commit 720db8983f
22 changed files with 607 additions and 37 deletions
@@ -6,6 +6,7 @@ import (
"github.com/Wei-Shaw/sub2api/internal/handler/dto"
"github.com/Wei-Shaw/sub2api/internal/pkg/response"
"github.com/Wei-Shaw/sub2api/internal/pkg/xai"
"github.com/Wei-Shaw/sub2api/internal/service"
"github.com/gin-gonic/gin"
)
@@ -239,3 +240,7 @@ func (h *GrokOAuthHandler) ResetQuota(c *gin.Context) {
}
response.Success(c, result)
}
func (h *GrokOAuthHandler) RuntimeSanity(c *gin.Context) {
response.Success(c, xai.RuntimeSanity())
}
@@ -125,3 +125,23 @@ func TestGrokOAuthHandlerResetQuotaReturnsUnsupported(t *testing.T) {
require.Contains(t, rec.Body.String(), `"reason":"GROK_QUOTA_RESET_UNSUPPORTED"`)
require.NotContains(t, rec.Body.String(), "access-token")
}
func TestGrokOAuthHandlerRuntimeSanityDoesNotExposeSecrets(t *testing.T) {
gin.SetMode(gin.TestMode)
t.Setenv(xai.EnvBaseURL, "http://127.0.0.1:8080/v1?access_token=secret")
t.Setenv(xai.EnvClientID, "client-secret-like-value")
handler := NewGrokOAuthHandler(nil, nil, nil)
router := gin.New()
router.GET("/api/v1/admin/grok/runtime-sanity", handler.RuntimeSanity)
rec := httptest.NewRecorder()
req := httptest.NewRequest(http.MethodGet, "/api/v1/admin/grok/runtime-sanity", nil)
router.ServeHTTP(rec, req)
require.Equal(t, http.StatusOK, rec.Code)
require.Contains(t, rec.Body.String(), `"public_gateway_scope":"responses_only"`)
require.Contains(t, rec.Body.String(), `"valid":false`)
require.NotContains(t, rec.Body.String(), "access_token")
require.NotContains(t, rec.Body.String(), "secret")
require.NotContains(t, rec.Body.String(), "client-secret-like-value")
}
+179 -10
View File
@@ -11,6 +11,9 @@ import (
"strings"
"sync"
"time"
"github.com/Wei-Shaw/sub2api/internal/util/logredact"
"github.com/Wei-Shaw/sub2api/internal/util/urlvalidator"
)
const (
@@ -19,17 +22,25 @@ const (
DefaultAuthorizeURL = OAuthIssuer + "/oauth2/authorize"
DefaultTokenURL = OAuthIssuer + "/oauth2/token"
DefaultBaseURL = "https://api.x.ai/v1"
DefaultCLIBaseURL = "https://cli-chat-proxy.grok.com/v1"
DefaultClientID = "b1a00492-073a-47ea-816f-4c329264a828"
DefaultScope = "openid profile email offline_access grok-cli:access api:access"
DefaultRedirectURI = "http://127.0.0.1:56121/callback"
SessionTTL = 30 * time.Minute
EnvAuthorizeURL = "XAI_OAUTH_AUTHORIZE_URL"
EnvTokenURL = "XAI_OAUTH_TOKEN_URL"
EnvClientID = "XAI_OAUTH_CLIENT_ID"
EnvScope = "XAI_OAUTH_SCOPE"
EnvRedirectURI = "XAI_OAUTH_REDIRECT_URI"
EnvBaseURL = "XAI_BASE_URL"
EnvAuthorizeURL = "XAI_OAUTH_AUTHORIZE_URL"
EnvTokenURL = "XAI_OAUTH_TOKEN_URL"
EnvClientID = "XAI_OAUTH_CLIENT_ID"
EnvScope = "XAI_OAUTH_SCOPE"
EnvRedirectURI = "XAI_OAUTH_REDIRECT_URI"
EnvBaseURL = "XAI_BASE_URL"
EnvAllowUnsafeURLOverrides = "XAI_ALLOW_UNSAFE_URL_OVERRIDES"
EnvUnsafeAllowHighConcurrency = "XAI_GROK_UNSAFE_ALLOW_CONCURRENCY_GT_ONE"
)
var (
oauthEndpointAllowedHosts = []string{"x.ai", "*.x.ai"}
baseURLAllowedHosts = []string{"api.x.ai", "cli-chat-proxy.grok.com"}
)
// OAuthSession stores one PKCE OAuth flow.
@@ -115,10 +126,18 @@ func EffectiveAuthorizeURL() string {
return envOrDefault(EnvAuthorizeURL, DefaultAuthorizeURL)
}
func ValidatedAuthorizeURL() (string, error) {
return ValidateOAuthEndpointURL(EffectiveAuthorizeURL())
}
func EffectiveTokenURL() string {
return envOrDefault(EnvTokenURL, DefaultTokenURL)
}
func ValidatedTokenURL() (string, error) {
return ValidateOAuthEndpointURL(EffectiveTokenURL())
}
func EffectiveClientID() string {
return envOrDefault(EnvClientID, DefaultClientID)
}
@@ -141,6 +160,139 @@ func EffectiveBaseURL(override string) string {
return strings.TrimRight(envOrDefault(EnvBaseURL, DefaultBaseURL), "/")
}
func ValidatedBaseURL(override string) (string, error) {
return ValidateBaseURL(EffectiveBaseURL(override))
}
type RuntimeSanityCheck struct {
Value string `json:"value"`
Valid bool `json:"valid"`
Error string `json:"error,omitempty"`
IsDefault bool `json:"is_default,omitempty"`
}
type RuntimeSanityReport struct {
BaseURL RuntimeSanityCheck `json:"base_url"`
OAuthAuthorizeURL RuntimeSanityCheck `json:"oauth_authorize_url"`
OAuthTokenURL RuntimeSanityCheck `json:"oauth_token_url"`
OAuthRedirectURI RuntimeSanityCheck `json:"oauth_redirect_uri"`
UnsafeURLOverrides bool `json:"unsafe_url_overrides"`
UnsafeHighConcurrency bool `json:"unsafe_high_concurrency"`
PublicGatewayScope string `json:"public_gateway_scope"`
ProxyPolicy string `json:"proxy_policy"`
}
func RuntimeSanity() RuntimeSanityReport {
return RuntimeSanityReport{
BaseURL: runtimeSanityCheck(EffectiveBaseURL(""), EnvBaseURL, ValidatedBaseURL),
OAuthAuthorizeURL: runtimeSanityCheck(EffectiveAuthorizeURL(), EnvAuthorizeURL, func(string) (string, error) { return ValidatedAuthorizeURL() }),
OAuthTokenURL: runtimeSanityCheck(EffectiveTokenURL(), EnvTokenURL, func(string) (string, error) { return ValidatedTokenURL() }),
OAuthRedirectURI: runtimeSanityCheck(EffectiveRedirectURI(""), EnvRedirectURI, validateRedirectURI),
UnsafeURLOverrides: AllowUnsafeURLOverrides(),
UnsafeHighConcurrency: AllowUnsafeHighConcurrency(),
PublicGatewayScope: "responses_only",
ProxyPolicy: "account_proxy_optional; upstream URL allowlists enforced unless unsafe overrides are enabled",
}
}
func runtimeSanityCheck(value string, envKey string, validate func(string) (string, error)) RuntimeSanityCheck {
normalized, err := validate(value)
check := RuntimeSanityCheck{
Value: sanitizeRuntimeURLValue(normalized),
Valid: err == nil,
IsDefault: strings.TrimSpace(os.Getenv(envKey)) == "",
}
if err != nil {
check.Value = sanitizeRuntimeURLValue(value)
check.Error = sanitizeRuntimeError(err.Error(), value)
}
return check
}
func validateRedirectURI(raw string) (string, error) {
return urlvalidator.ValidateURLFormat(raw, true)
}
func sanitizeRuntimeURLValue(raw string) string {
trimmed := strings.TrimSpace(raw)
if trimmed == "" {
return ""
}
parsed, err := url.Parse(trimmed)
if err != nil || parsed.Scheme == "" || parsed.Host == "" {
return trimmed
}
parsed.User = nil
parsed.RawQuery = ""
parsed.Fragment = ""
return strings.TrimRight(parsed.String(), "/")
}
func sanitizeRuntimeError(rawErr string, rawValue string) string {
redacted := logredact.RedactText(rawErr)
trimmedValue := strings.TrimSpace(rawValue)
if trimmedValue == "" {
return redacted
}
sanitizedValue := sanitizeRuntimeURLValue(trimmedValue)
redacted = strings.ReplaceAll(redacted, trimmedValue, sanitizedValue)
redacted = strings.ReplaceAll(redacted, logredact.RedactText(trimmedValue), sanitizedValue)
return redacted
}
func ValidateOAuthEndpointURL(raw string) (string, error) {
if AllowUnsafeURLOverrides() {
return urlvalidator.ValidateURLFormat(raw, true)
}
return urlvalidator.ValidateHTTPSURL(raw, urlvalidator.ValidationOptions{
AllowedHosts: oauthEndpointAllowedHosts,
RequireAllowlist: true,
AllowPrivate: false,
})
}
func ValidateBaseURL(raw string) (string, error) {
if AllowUnsafeURLOverrides() {
return urlvalidator.ValidateURLFormat(raw, true)
}
normalized, err := urlvalidator.ValidateHTTPSURL(raw, urlvalidator.ValidationOptions{
AllowedHosts: baseURLAllowedHosts,
RequireAllowlist: true,
AllowPrivate: false,
})
if err != nil {
return "", err
}
return normalizeKnownBaseURLPath(normalized)
}
func normalizeKnownBaseURLPath(raw string) (string, error) {
parsed, err := url.Parse(raw)
if err != nil || parsed.Scheme == "" || parsed.Host == "" {
return "", fmt.Errorf("invalid url: %s", raw)
}
path := strings.TrimRight(parsed.Path, "/")
if path == "" {
parsed.Path = "/v1"
parsed.RawPath = ""
return strings.TrimRight(parsed.String(), "/"), nil
}
if path != "/v1" {
return "", fmt.Errorf("base URL path must be /v1")
}
parsed.Path = path
parsed.RawPath = ""
return strings.TrimRight(parsed.String(), "/"), nil
}
func AllowUnsafeURLOverrides() bool {
return envBool(EnvAllowUnsafeURLOverrides)
}
func AllowUnsafeHighConcurrency() bool {
return envBool(EnvUnsafeAllowHighConcurrency)
}
func envOrDefault(key, fallback string) string {
if value := strings.TrimSpace(os.Getenv(key)); value != "" {
return value
@@ -148,6 +300,15 @@ func envOrDefault(key, fallback string) string {
return fallback
}
func envBool(key string) bool {
switch strings.ToLower(strings.TrimSpace(os.Getenv(key))) {
case "1", "true", "yes", "y", "on":
return true
default:
return false
}
}
func GenerateRandomBytes(n int) ([]byte, error) {
b := make([]byte, n)
if _, err := rand.Read(b); err != nil {
@@ -197,8 +358,12 @@ func base64URLEncode(data []byte) string {
return strings.TrimRight(base64.URLEncoding.EncodeToString(data), "=")
}
func BuildAuthorizationURL(state, codeChallenge, redirectURI, nonce string) string {
func BuildAuthorizationURL(state, codeChallenge, redirectURI, nonce string) (string, error) {
redirectURI = EffectiveRedirectURI(redirectURI)
authorizeURL, err := ValidatedAuthorizeURL()
if err != nil {
return "", fmt.Errorf("invalid authorize url: %w", err)
}
params := url.Values{}
params.Set("response_type", "code")
@@ -212,7 +377,7 @@ func BuildAuthorizationURL(state, codeChallenge, redirectURI, nonce string) stri
params.Set("plan", "generic")
params.Set("referrer", "sub2api")
return fmt.Sprintf("%s?%s", EffectiveAuthorizeURL(), params.Encode())
return fmt.Sprintf("%s?%s", authorizeURL, params.Encode()), nil
}
// AuthorizationInput is a parsed manual OAuth callback input.
@@ -256,8 +421,12 @@ func ParseAuthorizationInput(raw string) AuthorizationInput {
return AuthorizationInput{Code: trimmed}
}
func BuildResponsesURL(baseURL string) string {
return EffectiveBaseURL(baseURL) + "/responses"
func BuildResponsesURL(baseURL string) (string, error) {
validatedBaseURL, err := ValidatedBaseURL(baseURL)
if err != nil {
return "", fmt.Errorf("invalid base url: %w", err)
}
return validatedBaseURL + "/responses", nil
}
func BuildChatCompletionsURL(baseURL string) (string, error) {
+43 -1
View File
@@ -67,8 +67,10 @@ func TestBuildAuthorizationURLIncludesHermesCompatibleParameters(t *testing.T) {
t.Setenv(EnvAuthorizeURL, "https://auth.example.test/oauth2/authorize")
t.Setenv(EnvClientID, "client-id")
t.Setenv(EnvScope, "openid profile offline_access api:access")
t.Setenv(EnvAllowUnsafeURLOverrides, "true")
authURL := BuildAuthorizationURL("state", "challenge", "http://127.0.0.1:56121/callback", "nonce")
authURL, err := BuildAuthorizationURL("state", "challenge", "http://127.0.0.1:56121/callback", "nonce")
require.NoError(t, err)
parsed, err := url.Parse(authURL)
require.NoError(t, err)
@@ -140,6 +142,46 @@ func TestValidateXAIURLsAllowUnsafeDevOverride(t *testing.T) {
require.Equal(t, "http://127.0.0.1:8080/v1", baseURL)
}
func TestRuntimeSanityReportsSafeDefaults(t *testing.T) {
t.Setenv(EnvBaseURL, "")
t.Setenv(EnvAuthorizeURL, "")
t.Setenv(EnvTokenURL, "")
t.Setenv(EnvRedirectURI, "")
t.Setenv(EnvAllowUnsafeURLOverrides, "")
t.Setenv(EnvUnsafeAllowHighConcurrency, "")
report := RuntimeSanity()
require.True(t, report.BaseURL.Valid)
require.Equal(t, DefaultBaseURL, report.BaseURL.Value)
require.True(t, report.BaseURL.IsDefault)
require.True(t, report.OAuthAuthorizeURL.Valid)
require.True(t, report.OAuthTokenURL.Valid)
require.True(t, report.OAuthRedirectURI.Valid)
require.False(t, report.UnsafeURLOverrides)
require.False(t, report.UnsafeHighConcurrency)
require.Equal(t, "responses_only", report.PublicGatewayScope)
require.Contains(t, report.ProxyPolicy, "account_proxy_optional")
}
func TestRuntimeSanityReportsInvalidOverridesWithoutSecrets(t *testing.T) {
t.Setenv(EnvBaseURL, "http://127.0.0.1:8080/v1?access_token=secret")
t.Setenv(EnvAuthorizeURL, "https://auth.example.test/oauth2/authorize")
t.Setenv(EnvTokenURL, "https://auth.example.test/oauth2/token")
t.Setenv(EnvRedirectURI, "not a url")
t.Setenv(EnvClientID, "client-secret-like-value")
t.Setenv(EnvAllowUnsafeURLOverrides, "")
report := RuntimeSanity()
require.False(t, report.BaseURL.Valid)
require.False(t, report.BaseURL.IsDefault)
require.Contains(t, report.BaseURL.Error, "invalid url")
require.NotContains(t, report.BaseURL.Value, "secret")
require.False(t, report.OAuthAuthorizeURL.Valid)
require.False(t, report.OAuthTokenURL.Valid)
require.False(t, report.OAuthRedirectURI.Valid)
require.NotContains(t, report.ProxyPolicy, "client-secret-like-value")
}
func TestDefaultModelMappingIncludesGrokAliases(t *testing.T) {
t.Parallel()
+41 -7
View File
@@ -22,9 +22,26 @@ type QuotaSnapshot struct {
EntitlementStatus string `json:"entitlement_status,omitempty"`
StatusCode int `json:"status_code,omitempty"`
Headers map[string]string `json:"headers,omitempty"`
HeadersObserved bool `json:"headers_observed"`
ObservationSource string `json:"observation_source,omitempty"`
LastProbeAt string `json:"last_probe_at,omitempty"`
LastHeadersSeenAt string `json:"last_headers_seen_at,omitempty"`
UpdatedAt string `json:"updated_at"`
}
func (s *QuotaSnapshot) HasObservedHeaders() bool {
if s == nil {
return false
}
return s.HeadersObserved ||
s.Requests != nil ||
s.Tokens != nil ||
s.RetryAfterSeconds != nil ||
s.SubscriptionTier != "" ||
s.EntitlementStatus != "" ||
len(s.Headers) > 0
}
var quotaHeaderAllowlist = []string{
"x-ratelimit-limit-requests",
"x-ratelimit-remaining-requests",
@@ -40,16 +57,28 @@ var quotaHeaderAllowlist = []string{
}
func ParseQuotaHeaders(headers http.Header, statusCode int) *QuotaSnapshot {
if headers == nil {
return parseQuotaHeaders(headers, statusCode, "", false)
}
func ObserveQuotaHeaders(headers http.Header, statusCode int, source string) *QuotaSnapshot {
return parseQuotaHeaders(headers, statusCode, source, true)
}
func parseQuotaHeaders(headers http.Header, statusCode int, source string, keepEmpty bool) *QuotaSnapshot {
if headers == nil && !keepEmpty {
return nil
}
now := time.Now().UTC().Format(time.RFC3339)
snapshot := &QuotaSnapshot{
Requests: parseQuotaWindow(headers, "requests"),
Tokens: parseQuotaWindow(headers, "tokens"),
StatusCode: statusCode,
Headers: make(map[string]string),
UpdatedAt: time.Now().UTC().Format(time.RFC3339),
Requests: parseQuotaWindow(headers, "requests"),
Tokens: parseQuotaWindow(headers, "tokens"),
StatusCode: statusCode,
Headers: make(map[string]string),
ObservationSource: strings.TrimSpace(source),
UpdatedAt: now,
}
if snapshot.ObservationSource == "active_probe" {
snapshot.LastProbeAt = now
}
if retryAfter := parseRetryAfter(headers.Get("retry-after")); retryAfter != nil {
snapshot.RetryAfterSeconds = retryAfter
@@ -69,8 +98,13 @@ func ParseQuotaHeaders(headers http.Header, statusCode int) *QuotaSnapshot {
snapshot.SubscriptionTier == "" &&
snapshot.EntitlementStatus == "" &&
len(snapshot.Headers) == 0 {
if keepEmpty {
return snapshot
}
return nil
}
snapshot.HeadersObserved = true
snapshot.LastHeadersSeenAt = now
return snapshot
}
+17
View File
@@ -26,6 +26,8 @@ func TestParseQuotaHeaders(t *testing.T) {
snapshot := ParseQuotaHeaders(headers, http.StatusTooManyRequests)
require.NotNil(t, snapshot)
require.Equal(t, http.StatusTooManyRequests, snapshot.StatusCode)
require.True(t, snapshot.HeadersObserved)
require.NotEmpty(t, snapshot.LastHeadersSeenAt)
require.Equal(t, int64(100), *snapshot.Requests.Limit)
require.Equal(t, int64(25), *snapshot.Requests.Remaining)
require.Equal(t, int64(1893456000), *snapshot.Requests.ResetUnix)
@@ -44,3 +46,18 @@ func TestParseQuotaHeadersReturnsNilForMissingHeaders(t *testing.T) {
require.Nil(t, ParseQuotaHeaders(http.Header{}, http.StatusOK))
}
func TestObserveQuotaHeadersRecordsNoHeaderProbe(t *testing.T) {
t.Parallel()
snapshot := ObserveQuotaHeaders(http.Header{}, http.StatusOK, "active_probe")
require.NotNil(t, snapshot)
require.False(t, snapshot.HeadersObserved)
require.Equal(t, http.StatusOK, snapshot.StatusCode)
require.Equal(t, "active_probe", snapshot.ObservationSource)
require.NotEmpty(t, snapshot.LastProbeAt)
require.Empty(t, snapshot.LastHeadersSeenAt)
require.Empty(t, snapshot.Headers)
require.Nil(t, snapshot.Requests)
require.Nil(t, snapshot.Tokens)
}
+1
View File
@@ -398,6 +398,7 @@ func registerGrokOAuthRoutes(admin *gin.RouterGroup, h *handler.Handlers) {
grok.POST("/accounts/:id/refresh", h.Admin.GrokOAuth.RefreshAccountToken)
grok.GET("/accounts/:id/quota", h.Admin.GrokOAuth.QueryQuota)
grok.POST("/accounts/:id/reset-quota", h.Admin.GrokOAuth.ResetQuota)
grok.GET("/runtime-sanity", h.Admin.GrokOAuth.RuntimeSanity)
}
}
@@ -200,6 +200,9 @@ type UsageInfo struct {
GrokRetryAfterSeconds *int `json:"grok_retry_after_seconds,omitempty"`
GrokEntitlementStatus string `json:"grok_entitlement_status,omitempty"`
GrokQuotaSnapshotState string `json:"grok_quota_snapshot_state,omitempty"`
GrokLastQuotaProbeAt string `json:"grok_last_quota_probe_at,omitempty"`
GrokLastHeadersSeenAt string `json:"grok_last_headers_seen_at,omitempty"`
GrokLastStatusCode int `json:"grok_last_status_code,omitempty"`
GrokLocalUsage *WindowStats `json:"grok_local_usage,omitempty"`
// Antigravity 账号级信息
@@ -864,10 +867,12 @@ func (s *AccountUsageService) getGrokUsage(ctx context.Context, account *Account
s.grokQuotaFetcher = NewGrokQuotaFetcher()
}
usage := s.grokQuotaFetcher.BuildUsageInfo(account)
if usage.ErrorCode == "quota_unknown" {
usage.GrokQuotaSnapshotState = "unknown_until_first_response"
} else {
usage.GrokQuotaSnapshotState = "observed"
if usage.GrokQuotaSnapshotState == "" {
if usage.ErrorCode == "quota_unknown" {
usage.GrokQuotaSnapshotState = "unknown_until_first_response"
} else {
usage.GrokQuotaSnapshotState = "observed"
}
}
if s.usageLogRepo != nil && account != nil {
@@ -60,6 +60,11 @@ func (s *GrokOAuthService) GenerateAuthURL(ctx context.Context, proxyID *int64,
redirectURI = xai.EffectiveRedirectURI(redirectURI)
codeChallenge := xai.GenerateCodeChallenge(codeVerifier)
authURL, err := xai.BuildAuthorizationURL(state, codeChallenge, redirectURI, nonce)
if err != nil {
return nil, infraerrors.Newf(http.StatusBadRequest, "GROK_OAUTH_INVALID_AUTHORIZE_URL", "%v", err)
}
s.sessionStore.Set(sessionID, &xai.OAuthSession{
State: state,
CodeVerifier: codeVerifier,
@@ -72,7 +77,7 @@ func (s *GrokOAuthService) GenerateAuthURL(ctx context.Context, proxyID *int64,
})
return &GrokAuthURLResult{
AuthURL: xai.BuildAuthorizationURL(state, codeChallenge, redirectURI, nonce),
AuthURL: authURL,
SessionID: sessionID,
State: state,
}, nil
@@ -44,6 +44,16 @@ func (f *GrokQuotaFetcher) BuildUsageInfo(account *Account) *UsageInfo {
usage.SubscriptionTier = snapshot.SubscriptionTier
usage.SubscriptionTierRaw = snapshot.SubscriptionTier
usage.GrokEntitlementStatus = snapshot.EntitlementStatus
usage.GrokLastQuotaProbeAt = snapshot.LastProbeAt
usage.GrokLastHeadersSeenAt = snapshot.LastHeadersSeenAt
usage.GrokLastStatusCode = snapshot.StatusCode
if snapshot.HasObservedHeaders() {
usage.GrokQuotaSnapshotState = "observed"
} else {
usage.GrokQuotaSnapshotState = "no_headers"
usage.ErrorCode = "quota_unknown"
usage.Error = "No xAI quota headers observed on the latest Grok probe"
}
switch snapshot.StatusCode {
case 401:
@@ -45,6 +45,8 @@ func TestGrokQuotaFetcherBuildUsageInfoFromSnapshot(t *testing.T) {
SubscriptionTier: "supergrok",
EntitlementStatus: "active",
StatusCode: http.StatusTooManyRequests,
LastProbeAt: updatedAt,
LastHeadersSeenAt: updatedAt,
UpdatedAt: updatedAt,
},
},
@@ -53,15 +55,48 @@ func TestGrokQuotaFetcherBuildUsageInfoFromSnapshot(t *testing.T) {
usage := NewGrokQuotaFetcher().BuildUsageInfo(account)
require.Equal(t, "passive", usage.Source)
require.Equal(t, "rate_limited", usage.ErrorCode)
require.Equal(t, "observed", usage.GrokQuotaSnapshotState)
require.Equal(t, "supergrok", usage.SubscriptionTier)
require.Equal(t, "active", usage.GrokEntitlementStatus)
require.Equal(t, int64(100), *usage.GrokRequestQuota.Limit)
require.Equal(t, int64(12), *usage.GrokRequestQuota.Remaining)
require.Equal(t, 30, *usage.GrokRetryAfterSeconds)
require.NotNil(t, usage.UpdatedAt)
require.Equal(t, updatedAt, usage.GrokLastQuotaProbeAt)
require.Equal(t, updatedAt, usage.GrokLastHeadersSeenAt)
require.Equal(t, http.StatusTooManyRequests, usage.GrokLastStatusCode)
require.True(t, usage.UpdatedAt.Equal(time.Date(2030, 1, 1, 0, 0, 0, 0, time.UTC)))
}
func TestGrokQuotaFetcherBuildUsageInfoFromNoHeadersProbe(t *testing.T) {
t.Parallel()
probedAt := "2030-01-01T00:00:00Z"
account := &Account{
Platform: PlatformGrok,
Type: AccountTypeOAuth,
Extra: map[string]any{
grokQuotaSnapshotExtraKey: xai.QuotaSnapshot{
StatusCode: http.StatusOK,
HeadersObserved: false,
ObservationSource: "active_probe",
LastProbeAt: probedAt,
UpdatedAt: probedAt,
},
},
}
usage := NewGrokQuotaFetcher().BuildUsageInfo(account)
require.Equal(t, "quota_unknown", usage.ErrorCode)
require.Equal(t, "no_headers", usage.GrokQuotaSnapshotState)
require.Contains(t, usage.Error, "No xAI quota headers observed")
require.Equal(t, probedAt, usage.GrokLastQuotaProbeAt)
require.Empty(t, usage.GrokLastHeadersSeenAt)
require.Equal(t, http.StatusOK, usage.GrokLastStatusCode)
require.Nil(t, usage.GrokRequestQuota)
require.Nil(t, usage.GrokTokenQuota)
}
func TestGrokQuotaFetcherClassifiesForbiddenAndReauth(t *testing.T) {
t.Parallel()
@@ -85,8 +120,9 @@ func TestGrokQuotaFetcherClassifiesForbiddenAndReauth(t *testing.T) {
Type: AccountTypeOAuth,
Extra: map[string]any{
grokQuotaSnapshotExtraKey: xai.QuotaSnapshot{
StatusCode: tt.statusCode,
UpdatedAt: time.Now().UTC().Format(time.RFC3339),
StatusCode: tt.statusCode,
HeadersObserved: true,
UpdatedAt: time.Now().UTC().Format(time.RFC3339),
},
},
}
@@ -85,18 +85,16 @@ func (s *GrokQuotaService) ProbeUsage(ctx context.Context, accountID int64) (*Gr
}
defer func() { _ = resp.Body.Close() }()
snapshot := xai.ParseQuotaHeaders(resp.Header, resp.StatusCode)
if snapshot != nil {
_ = s.accountRepo.UpdateExtra(ctx, account.ID, map[string]any{
grokQuotaSnapshotExtraKey: snapshot,
})
}
snapshot := xai.ObserveQuotaHeaders(resp.Header, resp.StatusCode, "active_probe")
_ = s.accountRepo.UpdateExtra(ctx, account.ID, map[string]any{
grokQuotaSnapshotExtraKey: snapshot,
})
result := &GrokQuotaProbeResult{
Source: "active_probe",
Snapshot: snapshot,
StatusCode: resp.StatusCode,
HeadersObserved: snapshot != nil,
HeadersObserved: snapshot.HeadersObserved,
ResetSupported: false,
FetchedAt: time.Now().Unix(),
}
@@ -17,7 +17,11 @@ import (
type grokQuotaAccountRepo struct {
*mockAccountRepoForPlatform
updates map[int64]map[string]any
updates map[int64]map[string]any
tempUnschedCalls int
lastTempUnschedID int64
lastTempUnschedUntil time.Time
lastTempUnschedReason string
}
func (r *grokQuotaAccountRepo) UpdateExtra(_ context.Context, id int64, updates map[string]any) error {
@@ -28,6 +32,14 @@ func (r *grokQuotaAccountRepo) UpdateExtra(_ context.Context, id int64, updates
return 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 TestGrokQuotaServiceProbeUsageStoresHeaders(t *testing.T) {
t.Parallel()
@@ -64,6 +76,10 @@ func TestGrokQuotaServiceProbeUsageStoresHeaders(t *testing.T) {
require.Equal(t, http.StatusOK, result.StatusCode)
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)
@@ -74,6 +90,47 @@ func TestGrokQuotaServiceProbeUsageStoresHeaders(t *testing.T) {
require.NotNil(t, repo.updates[42][grokQuotaSnapshotExtraKey])
}
func TestGrokQuotaServiceProbeUsageStoresNoHeadersState(t *testing.T) {
t.Parallel()
account := &Account{
ID: 45,
Platform: PlatformGrok,
Type: AccountTypeOAuth,
Concurrency: 1,
Credentials: map[string]any{
"access_token": "access-token",
"expires_at": time.Now().Add(time.Hour).UTC().Format(time.RFC3339),
},
}
repo := &grokQuotaAccountRepo{
mockAccountRepoForPlatform: &mockAccountRepoForPlatform{
accountsByID: map[int64]*Account{45: account},
},
}
upstream := &httpUpstreamRecorder{resp: &http.Response{
StatusCode: http.StatusOK,
Header: http.Header{},
Body: io.NopCloser(strings.NewReader(`{"id":"resp_probe"}`)),
}}
svc := NewGrokQuotaService(repo, NewGrokTokenProvider(repo, nil, nil), upstream)
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)
}
func TestGrokQuotaServiceProbeUsageReturnsRateLimitedSnapshot(t *testing.T) {
t.Parallel()
@@ -91,3 +91,41 @@ func TestGrokTokenProviderRefreshesExpiredTokenOnRequestPath(t *testing.T) {
require.Greater(t, cache.setTTL, time.Duration(0))
require.Equal(t, 1, cache.releaseCalls)
}
func TestGrokTokenProviderRefreshFailureUnschedulesWithRedactedReason(t *testing.T) {
expiredAt := time.Now().Add(-time.Minute).UTC().Format(time.RFC3339)
account := &Account{
ID: 55,
Platform: PlatformGrok,
Type: AccountTypeOAuth,
Credentials: map[string]any{
"access_token": "expired-access-token",
"refresh_token": "refresh-token",
"expires_at": expiredAt,
"base_url": xai.DefaultCLIBaseURL,
},
}
repo := &tokenRefreshAccountRepo{}
repo.accountsByID = map[int64]*Account{55: account}
cache := &grokTokenCacheForProviderTest{lockResult: true}
tempCache := &tempUnschedCacheStub{}
provider := NewGrokTokenProvider(repo, cache, nil)
provider.SetRefreshAPI(NewOAuthRefreshAPI(repo, cache), &tokenRefresherStub{
err: errors.New("temporary refresh failure access_token=leaked-access refresh_token=leaked-refresh"),
})
provider.SetTempUnschedCache(tempCache)
token, err := provider.GetAccessToken(context.Background(), account)
require.Error(t, err)
require.Empty(t, token)
require.Equal(t, 1, repo.setTempUnschedCalls)
require.Equal(t, 0, repo.setErrorCalls)
require.Contains(t, repo.lastTempUnschedReason, "access_token=***")
require.Contains(t, repo.lastTempUnschedReason, "refresh_token=***")
require.NotContains(t, repo.lastTempUnschedReason, "leaked-access")
require.NotContains(t, repo.lastTempUnschedReason, "leaked-refresh")
require.Equal(t, 1, tempCache.setCalls)
require.NotNil(t, tempCache.lastState)
require.NotContains(t, tempCache.lastState.ErrorMessage, "leaked-access")
require.NotContains(t, tempCache.lastState.ErrorMessage, "leaked-refresh")
}
@@ -151,7 +151,10 @@ func patchGrokResponsesBody(body []byte, upstreamModel string) ([]byte, error) {
}
func buildGrokResponsesRequest(ctx context.Context, c *gin.Context, account *Account, body []byte, token string) (*http.Request, error) {
targetURL := xai.BuildResponsesURL(account.GetGrokBaseURL())
targetURL, err := xai.BuildResponsesURL(account.GetGrokBaseURL())
if err != nil {
return nil, err
}
req, err := http.NewRequestWithContext(ctx, http.MethodPost, targetURL, bytes.NewReader(body))
if err != nil {
return nil, err
@@ -208,6 +211,9 @@ func (s *OpenAIGatewayService) tempUnscheduleGrok(ctx context.Context, account *
return
}
until := time.Now().Add(cooldown)
if account.TempUnschedulableUntil != nil && account.TempUnschedulableUntil.After(until) {
until = *account.TempUnschedulableUntil
}
s.BlockAccountScheduling(account, until, reason)
if s.accountRepo != nil {
stateCtx, cancel := openAIAccountStateContext(ctx)
@@ -40,7 +40,7 @@ func TestPatchGrokResponsesBodySetsMappedModelAndDropsUnsupportedFields(t *testi
}
func TestBuildGrokResponsesRequestUsesAccountBaseURLAndBearerToken(t *testing.T) {
t.Parallel()
t.Setenv(xai.EnvAllowUnsafeURLOverrides, "true")
account := &Account{
Platform: PlatformGrok,
@@ -270,3 +270,78 @@ func TestForwardAsChatCompletionsForGrokStreamingUsesRawXAIChatCompletions(t *te
require.Contains(t, recorder.Body.String(), "data: [DONE]")
require.NotNil(t, repo.updates[53][grokQuotaSnapshotExtraKey])
}
func TestHandleGrokAccountUpstreamErrorTempUnschedulesReadinessStates(t *testing.T) {
tests := []struct {
name string
status int
headers http.Header
wantReason string
wantMinCooldown time.Duration
wantMaxCooldown time.Duration
}{
{
name: "unauthorized reauth",
status: http.StatusUnauthorized,
wantReason: "grok oauth token unauthorized",
wantMinCooldown: 10*time.Minute - time.Second,
wantMaxCooldown: 10*time.Minute + time.Second,
},
{
name: "forbidden entitlement",
status: http.StatusForbidden,
wantReason: "grok entitlement or subscription tier denied",
wantMinCooldown: 30*time.Minute - time.Second,
wantMaxCooldown: 30*time.Minute + time.Second,
},
{
name: "rate limited retry after",
status: http.StatusTooManyRequests,
headers: http.Header{"Retry-After": []string{"45"}},
wantReason: "grok rate limited",
wantMinCooldown: 44 * time.Second,
wantMaxCooldown: 46 * time.Second,
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
account := &Account{ID: 61, Platform: PlatformGrok, Type: AccountTypeOAuth}
repo := &grokQuotaAccountRepo{}
svc := &OpenAIGatewayService{accountRepo: repo}
before := time.Now()
svc.handleGrokAccountUpstreamError(context.Background(), account, tt.status, tt.headers, nil)
require.True(t, svc.isOpenAIAccountRuntimeBlocked(account))
require.Equal(t, 1, repo.tempUnschedCalls)
require.Equal(t, account.ID, repo.lastTempUnschedID)
require.Equal(t, tt.wantReason, repo.lastTempUnschedReason)
require.True(t, repo.lastTempUnschedUntil.After(before.Add(tt.wantMinCooldown)))
require.True(t, repo.lastTempUnschedUntil.Before(before.Add(tt.wantMaxCooldown)))
})
}
}
func TestHandleGrokAccountUpstreamErrorDoesNotShortenExistingPause(t *testing.T) {
existingUntil := time.Now().Add(15 * time.Minute)
account := &Account{
ID: 62,
Platform: PlatformGrok,
Type: AccountTypeOAuth,
TempUnschedulableUntil: &existingUntil,
TempUnschedulableReason: "existing pause",
}
repo := &grokQuotaAccountRepo{}
svc := &OpenAIGatewayService{accountRepo: repo}
svc.handleGrokAccountUpstreamError(context.Background(), account, http.StatusTooManyRequests, http.Header{"Retry-After": []string{"45"}}, nil)
require.Equal(t, 1, repo.tempUnschedCalls)
require.WithinDuration(t, existingUntil, repo.lastTempUnschedUntil, time.Second)
value, ok := svc.openaiAccountRuntimeBlockUntil.Load(account.ID)
require.True(t, ok)
runtimeUntil, ok := value.(time.Time)
require.True(t, ok)
require.WithinDuration(t, existingUntil, runtimeUntil, time.Second)
}
@@ -20,6 +20,8 @@ type tokenRefreshAccountRepo struct {
setErrorCalls int
clearTempCalls int
setTempUnschedCalls int
lastErrorMessage string
lastTempUnschedReason string
lastAccount *Account
updateErr error
}
@@ -51,6 +53,7 @@ func (r *tokenRefreshAccountRepo) UpdateCredentials(ctx context.Context, id int6
func (r *tokenRefreshAccountRepo) SetError(ctx context.Context, id int64, errorMsg string) error {
r.setErrorCalls++
r.lastErrorMessage = errorMsg
return nil
}
@@ -61,6 +64,7 @@ func (r *tokenRefreshAccountRepo) ClearTempUnschedulable(ctx context.Context, id
func (r *tokenRefreshAccountRepo) SetTempUnschedulable(ctx context.Context, id int64, until time.Time, reason string) error {
r.setTempUnschedCalls++
r.lastTempUnschedReason = reason
return nil
}
@@ -76,9 +80,13 @@ func (s *tokenCacheInvalidatorStub) InvalidateToken(ctx context.Context, account
type tempUnschedCacheStub struct {
deleteCalls int
setCalls int
lastState *TempUnschedState
}
func (s *tempUnschedCacheStub) SetTempUnsched(ctx context.Context, accountID int64, state *TempUnschedState) error {
s.setCalls++
s.lastState = state
return nil
}
+4
View File
@@ -54,6 +54,10 @@ export interface GrokQuotaSnapshot {
entitlement_status?: string
status_code?: number
headers?: Record<string, string>
headers_observed: boolean
observation_source?: string
last_probe_at?: string
last_headers_seen_at?: string
updated_at: string
}
@@ -383,11 +383,14 @@
{{ t('admin.accounts.usageWindow.grokRetryAfter', { time: grokRetryAfterLabel }) }}
</div>
<div v-if="grokQuotaUnknown" class="text-[10px] text-gray-500 dark:text-gray-400">
{{ t('admin.accounts.usageWindow.grokUnknown') }}
{{ grokQuotaUnknownLabel }}
</div>
<div v-else-if="usageInfo.error" class="truncate text-xs text-amber-600 dark:text-amber-400 max-w-[200px]" :title="usageInfo.error">
{{ usageErrorLabel }}
</div>
<div v-if="grokQuotaStatusLine" class="text-[10px] text-gray-500 dark:text-gray-400">
{{ grokQuotaStatusLine }}
</div>
<GrokQuotaProbeCell :account="account" />
</div>
<div v-else class="text-xs text-gray-400">-</div>
@@ -584,7 +587,7 @@ import { adminAPI } from '@/api/admin'
import type { Account, AccountUsageInfo, GeminiCredentials, WindowStats } from '@/types'
import { buildOpenAIUsageRefreshKey } from '@/utils/accountUsageRefresh'
import { enqueueUsageRequest } from '@/utils/usageLoadQueue'
import { formatCompactNumber } from '@/utils/format'
import { formatCompactNumber, formatRelativeTime } from '@/utils/format'
import UsageProgressBar from './UsageProgressBar.vue'
import AccountQuotaInfo from './AccountQuotaInfo.vue'
import OpenAIQuotaResetCell from './OpenAIQuotaResetCell.vue'
@@ -1031,6 +1034,34 @@ const grokQuotaUnknown = computed(() => {
if (grokRequestQuotaBar.value || grokTokenQuotaBar.value) return false
return usageInfo.value?.grok_quota_snapshot_state !== 'observed'
})
const grokQuotaUnknownLabel = computed(() => {
return usageInfo.value?.grok_quota_snapshot_state === 'no_headers'
? t('admin.accounts.usageWindow.grokNoHeaders')
: t('admin.accounts.usageWindow.grokUnknown')
})
const grokQuotaStatusLine = computed(() => {
if (props.account.platform !== 'grok') return null
const parts: string[] = []
const status = usageInfo.value?.grok_last_status_code
if (status) {
parts.push(t('admin.accounts.usageWindow.grokLastStatus', { status }))
}
if (usageInfo.value?.grok_last_quota_probe_at) {
parts.push(
t('admin.accounts.usageWindow.grokLastProbe', {
time: formatRelativeTime(usageInfo.value.grok_last_quota_probe_at)
})
)
}
if (usageInfo.value?.grok_last_headers_seen_at) {
parts.push(
t('admin.accounts.usageWindow.grokLastHeadersSeen', {
time: formatRelativeTime(usageInfo.value.grok_last_headers_seen_at)
})
)
}
return parts.length > 0 ? parts.join(' | ') : null
})
const grokLocalUsage = computed(() => usageInfo.value?.grok_local_usage || props.todayStats || null)
const grokEntitlementLabel = computed(() => {
const status = (usageInfo.value?.grok_entitlement_status || '').trim()
+3
View File
@@ -4159,6 +4159,9 @@ export default {
grokResetUnsupported: 'Reset unsupported',
grokResetUnsupportedTooltip: 'xAI does not expose reset credits for Grok OAuth accounts',
grokNoHeaders: 'No quota headers observed',
grokLastStatus: 'Status {status}',
grokLastProbe: 'Probe {time}',
grokLastHeadersSeen: 'Headers {time}',
passiveSampled: 'Passive',
activeQuery: 'Query'
},
+3
View File
@@ -3417,6 +3417,9 @@ export default {
grokResetUnsupported: '不支持重置',
grokResetUnsupportedTooltip: 'xAI 未向 Grok OAuth 账号开放重置额度接口',
grokNoHeaders: '未观察到配额响应头',
grokLastStatus: '状态 {status}',
grokLastProbe: '探测 {time}',
grokLastHeadersSeen: '响应头 {time}',
passiveSampled: '被动采样',
activeQuery: '查询'
},
+3
View File
@@ -969,6 +969,9 @@ export interface AccountUsageInfo {
grok_retry_after_seconds?: number | null
grok_entitlement_status?: string
grok_quota_snapshot_state?: string
grok_last_quota_probe_at?: string
grok_last_headers_seen_at?: string
grok_last_status_code?: number
grok_local_usage?: WindowStats | null
ai_credits?: Array<{
credit_type?: string