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
+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)
}