mirror of
https://github.com/Wei-Shaw/sub2api.git
synced 2026-09-24 16:05:44 +08:00
test: harden grok quota readiness
This commit is contained in:
@@ -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) {
|
||||
|
||||
@@ -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()
|
||||
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user