Sub2API v1.0 - AI API 网关(二开初始版本,基于上游 Wei-Shaw/sub2api)
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

This commit is contained in:
李建琦
2026-08-21 18:30:13 +08:00
commit 6d655c9903
3584 changed files with 1270640 additions and 0 deletions
+458
View File
@@ -0,0 +1,458 @@
package xai
import (
"encoding/json"
"fmt"
"math"
"net/http"
"strconv"
"strings"
"time"
)
const (
// CLI client identity required by cli-chat-proxy billing endpoints.
CLITokenAuthHeader = "x-xai-token-auth"
CLITokenAuthValue = "xai-grok-cli"
CLIClientVersionHeader = "x-grok-client-version"
// CLIClientVersion is the one place the pinned Grok CLI version lives. The
// repository and service layers build their own client identity from it, so
// one bump here covers OAuth traffic and billing probes together.
// Keep in sync with https://x.ai/cli/stable.
CLIClientVersion = "0.2.114"
// billingCLIUserAgent is the legacy pager/shell UA used by billing probes.
// Distinct from CLIUserAgent() in cli_identity.go (workspace-style UA).
billingCLIUserAgent = "grok-pager/" + CLIClientVersion + " grok-shell/" + CLIClientVersion + " (macos; aarch64)"
BillingWeeklyPath = "/billing?format=credits"
BillingMonthlyPath = "/billing"
SuperGrokLimitCents = 15_000 // $150.00
SuperGrokHeavyLimitCents = 150_000 // $1,500.00
)
// BillingPeriod describes the current weekly/monthly window.
type BillingPeriod struct {
Type string `json:"type,omitempty"`
Start string `json:"start,omitempty"`
End string `json:"end,omitempty"`
}
// BillingProductUsage is per-product usage inside the weekly credits window.
type BillingProductUsage struct {
Product string `json:"product,omitempty"`
UsagePercent *float64 `json:"usagePercent,omitempty"`
}
// BillingConfig is the nested config object from /v1/billing responses.
// Weekly (`?format=credits`) and monthly (`/billing`) share this shape; absolute
// money fields typically appear on the credits (prepaid/on-demand) or monthly
// (limit/used) responses.
type BillingConfig struct {
CurrentPeriod *BillingPeriod `json:"currentPeriod,omitempty"`
CreditUsagePercent *float64 `json:"creditUsagePercent,omitempty"`
ProductUsage []BillingProductUsage `json:"productUsage,omitempty"`
MonthlyLimit json.RawMessage `json:"monthlyLimit,omitempty"`
Used json.RawMessage `json:"used,omitempty"`
OnDemandCap json.RawMessage `json:"onDemandCap,omitempty"`
OnDemandUsed json.RawMessage `json:"onDemandUsed,omitempty"`
PrepaidBalance json.RawMessage `json:"prepaidBalance,omitempty"`
IsUnifiedBillingUser bool `json:"isUnifiedBillingUser,omitempty"`
TopUpMethod string `json:"topUpMethod,omitempty"`
BillingPeriodStart string `json:"billingPeriodStart,omitempty"`
BillingPeriodEnd string `json:"billingPeriodEnd,omitempty"`
}
// BillingPayload is the top-level body from /v1/billing.
type BillingPayload struct {
Config *BillingConfig `json:"config,omitempty"`
}
// BillingProductSummary is a normalized product usage row for UI.
type BillingProductSummary struct {
Product string `json:"product"`
UsagePercent *float64 `json:"usage_percent,omitempty"`
}
// BillingSummary is the merged weekly + monthly billing view.
// Cents fields remain the authoritative monthly numbers; dollar fields are the
// operator-facing absolute money view (prepaid / on-demand / monthly $).
type BillingSummary struct {
PeriodType string `json:"period_type,omitempty"` // weekly | monthly | unknown
UsagePercent *float64 `json:"usage_percent,omitempty"`
PeriodStart string `json:"period_start,omitempty"`
PeriodEnd string `json:"period_end,omitempty"`
ProductUsage []BillingProductSummary `json:"product_usage,omitempty"`
MonthlyLimitCents *float64 `json:"monthly_limit_cents,omitempty"`
UsedCents *float64 `json:"used_cents,omitempty"`
IncludedUsedCents *float64 `json:"included_used_cents,omitempty"`
BillingPeriodStart string `json:"billing_period_start,omitempty"`
BillingPeriodEnd string `json:"billing_period_end,omitempty"`
UsedPercent *float64 `json:"used_percent,omitempty"`
// Absolute money (USD). Prepaid/on-demand come from credits probe as dollars.
// MonthlyLimit/MonthlyUsed are cents/100 for consistent $ display.
PrepaidBalance *float64 `json:"prepaid_balance,omitempty"`
MonthlyLimit *float64 `json:"monthly_limit,omitempty"`
MonthlyUsed *float64 `json:"monthly_used,omitempty"`
OnDemandCap *float64 `json:"on_demand_cap,omitempty"`
OnDemandUsed *float64 `json:"on_demand_used,omitempty"`
TopUpMethod string `json:"top_up_method,omitempty"`
IsUnifiedBillingUser bool `json:"is_unified_billing_user,omitempty"`
Plan string `json:"plan,omitempty"` // SuperGrok | SuperGrok Heavy | ""
StatusCode int `json:"status_code,omitempty"`
WeeklyStatusCode int `json:"weekly_status_code,omitempty"`
MonthlyStatusCode int `json:"monthly_status_code,omitempty"`
Source string `json:"source,omitempty"`
FetchedAt string `json:"fetched_at,omitempty"`
UpdatedAt string `json:"updated_at,omitempty"`
WeeklyUpdatedAt string `json:"weekly_updated_at,omitempty"`
MonthlyUpdatedAt string `json:"monthly_updated_at,omitempty"`
Partial bool `json:"partial,omitempty"`
FailedWindows []string `json:"failed_windows,omitempty"`
}
// BuildBillingURL builds weekly or monthly billing URL against the CLI chat proxy.
func BuildBillingURL(formatCredits bool) string {
base := strings.TrimRight(DefaultCLIBaseURL, "/")
if formatCredits {
return base + BillingWeeklyPath
}
return base + BillingMonthlyPath
}
// BuildBillingURLWithValidator builds the weekly or monthly billing URL against
// the caller-resolved base URL, applying the caller's outbound URL trust policy
// first. Accounts forwarding through a custom upstream keep their billing
// probes on the same upstream.
func BuildBillingURLWithValidator(baseURL string, formatCredits bool, validator BaseURLValidator) (string, error) {
validatedBaseURL, err := validatedBaseURLWithValidator(baseURL, validator)
if err != nil {
return "", fmt.Errorf("invalid base url: %w", err)
}
if formatCredits {
return validatedBaseURL + BillingWeeklyPath, nil
}
return validatedBaseURL + BillingMonthlyPath, nil
}
// ApplyCLIBillingHeaders sets Authorization + CLI identity headers for billing GETs.
func ApplyCLIBillingHeaders(req *http.Request, accessToken string) {
if req == nil {
return
}
token := strings.TrimSpace(accessToken)
if token != "" {
req.Header.Set("Authorization", "Bearer "+token)
}
req.Header.Set("Accept", "application/json")
req.Header.Set("Content-Type", "application/json")
req.Header.Set(CLITokenAuthHeader, CLITokenAuthValue)
req.Header.Set(CLIClientVersionHeader, CLIClientVersion)
req.Header.Set("User-Agent", billingCLIUserAgent)
}
// ParseBillingPayload unmarshals a billing API response body.
func ParseBillingPayload(body []byte) (*BillingPayload, error) {
if len(body) == 0 {
return nil, fmt.Errorf("empty billing body")
}
var payload BillingPayload
if err := json.Unmarshal(body, &payload); err != nil {
return nil, err
}
return &payload, nil
}
// BuildBillingSummary normalizes a billing config into a UI-friendly summary.
func BuildBillingSummary(config *BillingConfig) *BillingSummary {
if config == nil {
return nil
}
summary := &BillingSummary{}
period := config.CurrentPeriod
periodType := resolvePeriodType(period)
creditUsage := cloneFloat(config.CreditUsagePercent)
// Weekly period bounds must not fall back to monthly billing period ends —
// that would park accounts on a multi-week horizon when weekly UsagePercent
// is high (scheduler seven_day uses PeriodEnd).
periodStart := ""
periodEnd := ""
if period != nil {
periodStart = strings.TrimSpace(period.Start)
periodEnd = strings.TrimSpace(period.End)
}
products := make([]BillingProductSummary, 0, len(config.ProductUsage))
for _, item := range config.ProductUsage {
product := strings.TrimSpace(item.Product)
if product == "" {
continue
}
products = append(products, BillingProductSummary{
Product: product,
UsagePercent: cloneFloat(item.UsagePercent),
})
}
monthlyLimit := parseCentValue(config.MonthlyLimit)
used := parseCentValue(config.Used)
// Absolute money on credits responses is dollar-denominated ({"val": 12}).
// Monthly limit/used are cents (same as MonthlyLimitCents / UsedCents).
prepaid := parseCentValue(config.PrepaidBalance)
onDemandCap := parseCentValue(config.OnDemandCap)
onDemandUsed := parseCentValue(config.OnDemandUsed)
billingStart := strings.TrimSpace(config.BillingPeriodStart)
billingEnd := strings.TrimSpace(config.BillingPeriodEnd)
var includedUsed *float64
if used != nil {
if monthlyLimit != nil && *monthlyLimit > 0 {
v := math.Min(*used, *monthlyLimit)
includedUsed = &v
} else {
includedUsed = cloneFloat(used)
}
}
var usedPercent *float64
if monthlyLimit != nil && *monthlyLimit > 0 && includedUsed != nil {
v := (*includedUsed / *monthlyLimit) * 100
usedPercent = &v
}
hasWeekly := creditUsage != nil || periodType == "weekly" || len(products) > 0 || prepaid != nil || onDemandCap != nil || onDemandUsed != nil
hasMonthly := monthlyLimit != nil || used != nil || (!hasWeekly && billingEnd != "")
if !hasWeekly && !hasMonthly {
return nil
}
if hasWeekly {
if periodType == "unknown" {
periodType = "weekly"
}
summary.PeriodType = periodType
summary.UsagePercent = creditUsage
summary.PeriodStart = periodStart
summary.PeriodEnd = periodEnd
} else {
// Monthly-only: do not put monthly % into UsagePercent (weekly bar field).
// Frontend weekly bar only renders when PeriodType == weekly.
summary.PeriodType = "monthly"
summary.PeriodStart = billingStart
summary.PeriodEnd = billingEnd
}
summary.ProductUsage = products
summary.MonthlyLimitCents = monthlyLimit
summary.UsedCents = used
summary.IncludedUsedCents = includedUsed
if hasMonthly {
summary.BillingPeriodStart = billingStart
summary.BillingPeriodEnd = billingEnd
}
summary.UsedPercent = usedPercent
summary.PrepaidBalance = prepaid
if onDemandCap != nil {
summary.OnDemandCap = onDemandCap
}
if onDemandUsed != nil {
summary.OnDemandUsed = onDemandUsed
}
// Expose monthly cents as dollars for UI absolute rows.
if monthlyLimit != nil {
v := *monthlyLimit / 100
summary.MonthlyLimit = &v
}
if used != nil {
v := *used / 100
summary.MonthlyUsed = &v
}
summary.TopUpMethod = strings.TrimSpace(config.TopUpMethod)
summary.IsUnifiedBillingUser = config.IsUnifiedBillingUser
summary.Plan = resolvePlan(monthlyLimit)
return summary
}
// MergeBillingProbeResult updates successful billing domains while retaining
// the previous value for any domain that could not be refreshed.
func MergeBillingProbeResult(previous, weekly, monthly *BillingSummary, weeklyOK, monthlyOK bool) *BillingSummary {
var out BillingSummary
if previous != nil {
out = *previous
previousUpdatedAt := previous.UpdatedAt
if previousUpdatedAt == "" {
previousUpdatedAt = previous.FetchedAt
}
if out.WeeklyUpdatedAt == "" && (out.UsagePercent != nil || len(out.ProductUsage) > 0) {
out.WeeklyUpdatedAt = previousUpdatedAt
}
if out.MonthlyUpdatedAt == "" && (out.MonthlyLimitCents != nil || out.UsedPercent != nil) {
out.MonthlyUpdatedAt = previousUpdatedAt
}
}
now := time.Now().UTC().Format(time.RFC3339)
if weeklyOK && weekly != nil {
out.PeriodType = weekly.PeriodType
out.UsagePercent = weekly.UsagePercent
out.PeriodStart = weekly.PeriodStart
out.PeriodEnd = weekly.PeriodEnd
out.ProductUsage = weekly.ProductUsage
// Absolute prepaid / on-demand usually ride the credits (weekly) response.
if weekly.PrepaidBalance != nil {
out.PrepaidBalance = weekly.PrepaidBalance
}
if weekly.OnDemandCap != nil {
out.OnDemandCap = weekly.OnDemandCap
}
if weekly.OnDemandUsed != nil {
out.OnDemandUsed = weekly.OnDemandUsed
}
if weekly.TopUpMethod != "" {
out.TopUpMethod = weekly.TopUpMethod
}
if weekly.IsUnifiedBillingUser {
out.IsUnifiedBillingUser = true
}
out.WeeklyUpdatedAt = now
}
if monthlyOK && monthly != nil {
if out.PeriodType == "" {
out.PeriodType = "monthly"
}
out.MonthlyLimitCents = monthly.MonthlyLimitCents
out.UsedCents = monthly.UsedCents
out.IncludedUsedCents = monthly.IncludedUsedCents
out.BillingPeriodStart = monthly.BillingPeriodStart
out.BillingPeriodEnd = monthly.BillingPeriodEnd
out.UsedPercent = monthly.UsedPercent
out.MonthlyLimit = monthly.MonthlyLimit
out.MonthlyUsed = monthly.MonthlyUsed
// Monthly probe may also carry on-demand cap when credits omitted it.
if monthly.OnDemandCap != nil && out.OnDemandCap == nil {
out.OnDemandCap = monthly.OnDemandCap
}
if monthly.OnDemandUsed != nil && out.OnDemandUsed == nil {
out.OnDemandUsed = monthly.OnDemandUsed
}
out.Plan = monthly.Plan
out.MonthlyUpdatedAt = now
}
out.Partial = !weeklyOK || !monthlyOK
out.FailedWindows = nil
if !weeklyOK {
out.FailedWindows = append(out.FailedWindows, "weekly")
}
if !monthlyOK {
out.FailedWindows = append(out.FailedWindows, "monthly")
}
if !weeklyOK && !monthlyOK && previous == nil {
return nil
}
return &out
}
// StampBillingSummary sets fetch metadata.
func StampBillingSummary(summary *BillingSummary, statusCode int, source string) *BillingSummary {
if summary == nil {
return nil
}
now := time.Now().UTC().Format(time.RFC3339)
summary.StatusCode = statusCode
summary.Source = source
summary.FetchedAt = now
summary.UpdatedAt = now
return summary
}
func resolvePeriodType(period *BillingPeriod) string {
if period == nil {
return "unknown"
}
raw := strings.ToLower(strings.TrimSpace(period.Type))
if strings.Contains(raw, "weekly") {
return "weekly"
}
if strings.Contains(raw, "monthly") {
return "monthly"
}
return "unknown"
}
func resolvePlan(monthlyLimitCents *float64) string {
if monthlyLimitCents == nil {
return ""
}
// Allow small float noise.
limit := math.Round(*monthlyLimitCents)
switch limit {
case SuperGrokLimitCents:
return "SuperGrok"
case SuperGrokHeavyLimitCents:
return "SuperGrok Heavy"
default:
return ""
}
}
func parseCentValue(raw json.RawMessage) *float64 {
if len(raw) == 0 || string(raw) == "null" {
return nil
}
// Object form: {"val": 123}
var obj struct {
Val any `json:"val"`
}
if err := json.Unmarshal(raw, &obj); err == nil && obj.Val != nil {
return anyToFloat(obj.Val)
}
// Bare number / string
var n any
if err := json.Unmarshal(raw, &n); err != nil {
return nil
}
return anyToFloat(n)
}
func anyToFloat(v any) *float64 {
switch n := v.(type) {
case float64:
return &n
case float32:
f := float64(n)
return &f
case int:
f := float64(n)
return &f
case int64:
f := float64(n)
return &f
case json.Number:
f, err := n.Float64()
if err != nil {
return nil
}
return &f
case string:
s := strings.TrimSpace(n)
if s == "" {
return nil
}
f, err := strconv.ParseFloat(s, 64)
if err != nil {
return nil
}
return &f
default:
return nil
}
}
func cloneFloat(v *float64) *float64 {
if v == nil {
return nil
}
f := *v
return &f
}
+176
View File
@@ -0,0 +1,176 @@
package xai
import (
"encoding/json"
"net/http"
"testing"
"github.com/stretchr/testify/require"
)
func TestBuildBillingURL(t *testing.T) {
t.Parallel()
require.Equal(t, "https://cli-chat-proxy.grok.com/v1/billing?format=credits", BuildBillingURL(true))
require.Equal(t, "https://cli-chat-proxy.grok.com/v1/billing", BuildBillingURL(false))
}
func TestBuildBillingURLWithValidator(t *testing.T) {
t.Parallel()
weeklyURL, err := BuildBillingURLWithValidator(DefaultCLIBaseURL, true, ValidateTrustedBaseURL)
require.NoError(t, err)
require.Equal(t, "https://cli-chat-proxy.grok.com/v1/billing?format=credits", weeklyURL)
monthlyURL, err := BuildBillingURLWithValidator("https://relay.example.test/v1", false, ValidateBaseURL)
require.NoError(t, err)
require.Equal(t, "https://relay.example.test/v1/billing", monthlyURL)
_, err = BuildBillingURLWithValidator("https://relay.example.test/v1", true, ValidateTrustedBaseURL)
require.Error(t, err)
}
func TestApplyCLIBillingHeaders(t *testing.T) {
t.Parallel()
req, err := http.NewRequest(http.MethodGet, BuildBillingURL(true), nil)
require.NoError(t, err)
ApplyCLIBillingHeaders(req, " token ")
require.Equal(t, "Bearer token", req.Header.Get("Authorization"))
require.Equal(t, CLITokenAuthValue, req.Header.Get(CLITokenAuthHeader))
require.Equal(t, CLIClientVersion, req.Header.Get(CLIClientVersionHeader))
require.Equal(t, "grok-pager/"+CLIClientVersion+" grok-shell/"+CLIClientVersion+" (macos; aarch64)", req.UserAgent())
}
func TestBuildBillingSummaryWeeklyAndMonthly(t *testing.T) {
t.Parallel()
weeklyBody := []byte(`{
"config": {
"currentPeriod": {"type":"WEEKLY","start":"2026-07-09T03:25:00Z","end":"2026-07-16T03:25:00Z"},
"creditUsagePercent": 2.0,
"productUsage": [{"product":"Api","usagePercent":2.0}],
"prepaidBalance": {"val": 12},
"onDemandCap": {"val": 100},
"onDemandUsed": {"val": 5},
"isUnifiedBillingUser": true
}
}`)
monthlyBody := []byte(`{
"config": {
"monthlyLimit": {"val": 15000},
"used": {"val": 78},
"billingPeriodStart": "2026-07-01T00:00:00Z",
"billingPeriodEnd": "2026-08-01T00:00:00Z"
}
}`)
weeklyPayload, err := ParseBillingPayload(weeklyBody)
require.NoError(t, err)
monthlyPayload, err := ParseBillingPayload(monthlyBody)
require.NoError(t, err)
weekly := BuildBillingSummary(weeklyPayload.Config)
monthly := BuildBillingSummary(monthlyPayload.Config)
require.NotNil(t, weekly)
require.NotNil(t, monthly)
require.Equal(t, "weekly", weekly.PeriodType)
require.InDelta(t, 2.0, *weekly.UsagePercent, 1e-9)
require.Equal(t, "Api", weekly.ProductUsage[0].Product)
require.InDelta(t, 12, *weekly.PrepaidBalance, 1e-9)
require.InDelta(t, 100, *weekly.OnDemandCap, 1e-9)
require.InDelta(t, 5, *weekly.OnDemandUsed, 1e-9)
require.True(t, weekly.IsUnifiedBillingUser)
require.Equal(t, "SuperGrok", monthly.Plan)
require.InDelta(t, 15000, *monthly.MonthlyLimitCents, 1e-9)
require.InDelta(t, 78, *monthly.UsedCents, 1e-9)
require.InDelta(t, 0.52, *monthly.UsedPercent, 1e-2)
require.InDelta(t, 150, *monthly.MonthlyLimit, 1e-9)
require.InDelta(t, 0.78, *monthly.MonthlyUsed, 1e-9)
merged := MergeBillingProbeResult(nil, weekly, monthly, true, true)
require.Equal(t, "weekly", merged.PeriodType)
require.InDelta(t, 2.0, *merged.UsagePercent, 1e-9)
require.Equal(t, "SuperGrok", merged.Plan)
require.InDelta(t, 15000, *merged.MonthlyLimitCents, 1e-9)
require.Equal(t, "2026-08-01T00:00:00Z", merged.BillingPeriodEnd)
require.InDelta(t, 12, *merged.PrepaidBalance, 1e-9)
require.InDelta(t, 100, *merged.OnDemandCap, 1e-9)
require.InDelta(t, 150, *merged.MonthlyLimit, 1e-9)
require.InDelta(t, 0.78, *merged.MonthlyUsed, 1e-9)
}
func TestParseCentValueBareNumber(t *testing.T) {
t.Parallel()
raw, _ := json.Marshal(15000)
v := parseCentValue(raw)
require.NotNil(t, v)
require.InDelta(t, 15000, *v, 1e-9)
}
func TestBuildBillingSummaryMonthlyOnlyKeepsWeeklyUsageEmpty(t *testing.T) {
t.Parallel()
payload, err := ParseBillingPayload([]byte(`{"config":{"monthlyLimit":{"val":15000},"used":{"val":7500},"billingPeriodStart":"2026-07-01T00:00:00Z","billingPeriodEnd":"2026-08-01T00:00:00Z"}}`))
require.NoError(t, err)
summary := BuildBillingSummary(payload.Config)
require.NotNil(t, summary)
require.Equal(t, "monthly", summary.PeriodType)
require.Nil(t, summary.UsagePercent)
require.InDelta(t, 50, *summary.UsedPercent, 1e-9)
}
func TestBuildBillingSummaryWeeklyDoesNotInheritMonthlyPeriodEnd(t *testing.T) {
t.Parallel()
// Weekly usage without currentPeriod.end must not copy billingPeriodEnd (monthly).
payload, err := ParseBillingPayload([]byte(`{
"config": {
"creditUsagePercent": 95.0,
"productUsage": [{"product":"Api","usagePercent":95.0}],
"billingPeriodStart": "2026-07-01T00:00:00Z",
"billingPeriodEnd": "2026-08-01T00:00:00Z",
"monthlyLimit": {"val": 15000},
"used": {"val": 1000}
}
}`))
require.NoError(t, err)
summary := BuildBillingSummary(payload.Config)
require.NotNil(t, summary)
require.Equal(t, "weekly", summary.PeriodType)
require.Equal(t, "", summary.PeriodEnd)
require.Equal(t, "2026-08-01T00:00:00Z", summary.BillingPeriodEnd)
}
func TestMergeBillingProbeResultRetainsFailedWindow(t *testing.T) {
t.Parallel()
previous := &BillingSummary{
PeriodType: "weekly",
UsagePercent: floatPointer(100),
PeriodEnd: "2026-07-16T00:00:00Z",
MonthlyLimitCents: floatPointer(15000),
UsedPercent: floatPointer(20),
BillingPeriodEnd: "2026-08-01T00:00:00Z",
WeeklyUpdatedAt: "2026-07-10T00:00:00Z",
MonthlyUpdatedAt: "2026-07-10T00:00:00Z",
FailedWindows: []string{"monthly"},
}
monthly := &BillingSummary{
PeriodType: "monthly",
MonthlyLimitCents: floatPointer(15000),
UsedPercent: floatPointer(30),
BillingPeriodEnd: "2026-08-01T00:00:00Z",
}
merged := MergeBillingProbeResult(previous, nil, monthly, false, true)
require.Equal(t, "weekly", merged.PeriodType)
require.InDelta(t, 100, *merged.UsagePercent, 1e-9)
require.Equal(t, previous.WeeklyUpdatedAt, merged.WeeklyUpdatedAt)
require.InDelta(t, 30, *merged.UsedPercent, 1e-9)
require.NotEqual(t, previous.MonthlyUpdatedAt, merged.MonthlyUpdatedAt)
require.True(t, merged.Partial)
require.Equal(t, []string{"weekly"}, merged.FailedWindows)
require.Equal(t, []string{"monthly"}, previous.FailedWindows)
}
func floatPointer(value float64) *float64 {
return &value
}
+79
View File
@@ -0,0 +1,79 @@
package xai
import (
"net/http"
"os"
"strings"
"golang.org/x/mod/semver"
)
// Fixed Grok Build / CLI-chat-proxy client identity.
// These values are intentionally pinned in-binary (not scraped from live CLI).
// Operators may bump the version via XAI_GROK_CLI_VERSION without a release.
const (
// CLIProxyHost is the hostname that requires the official CLI identity headers.
CLIProxyHost = "cli-chat-proxy.grok.com"
// CLIStableVersion is the known-good minimum client version accepted by cli-chat-proxy.
CLIStableVersion = "0.2.93"
// CLIVersionEnv is the optional operator override for CLIStableVersion.
CLIVersionEnv = "XAI_GROK_CLI_VERSION"
// CLITokenAuth is required by cli-chat-proxy for Grok Build OAuth tokens.
CLITokenAuth = "xai-grok-cli"
// CLIClientIdentifier is the x-grok-client-identifier value used by Grok shell/CLI.
CLIClientIdentifier = "grok-shell"
// CLIClientMode is used by billing / quota probes on the CLI surface.
CLIClientMode = "cli"
)
// ResolveCLIVersion returns a supported CLI client version.
// Empty or invalid overrides fall back to CLIClientVersion (the pinned
// preferred client pin in billing.go). CLIStableVersion is only the minimum
// accepted by IsSupportedCLIVersion, not the default identity we advertise.
func ResolveCLIVersion() string {
version := strings.TrimSpace(os.Getenv(CLIVersionEnv))
if !IsSupportedCLIVersion(version) {
return CLIClientVersion
}
return version
}
// IsSupportedCLIVersion reports whether version is a valid semver string at or
// above CLIStableVersion (prereleases below a higher release are rejected when
// they compare less than the stable pin).
func IsSupportedCLIVersion(version string) bool {
canonical := "v" + version
minimum := "v" + CLIStableVersion
return semver.IsValid(canonical) &&
semver.Canonical(canonical) == canonical &&
semver.Compare(canonical, minimum) >= 0
}
// CLIUserAgent builds the workspace-style User-Agent for a CLI client version.
func CLIUserAgent(version string) string {
if strings.TrimSpace(version) == "" {
version = CLIClientVersion
}
return "xai-grok-workspace/" + version
}
// ApplyCLIProxyHeaders stamps the fixed Grok CLI identity when the request
// targets cli-chat-proxy. Direct api.x.ai traffic is left unchanged.
func ApplyCLIProxyHeaders(req *http.Request) {
if req == nil || req.URL == nil || !strings.EqualFold(strings.TrimSpace(req.URL.Hostname()), CLIProxyHost) {
return
}
if req.Header == nil {
req.Header = make(http.Header)
}
version := ResolveCLIVersion()
req.Header.Set("X-XAI-Token-Auth", CLITokenAuth)
req.Header.Set("x-grok-client-version", version)
req.Header.Set("x-grok-client-identifier", CLIClientIdentifier)
req.Header.Set("User-Agent", CLIUserAgent(version))
}
@@ -0,0 +1,67 @@
package xai
import (
"net/http"
"testing"
"github.com/stretchr/testify/require"
)
func TestResolveCLIVersionDefaultsToPinnedClientVersion(t *testing.T) {
t.Setenv(CLIVersionEnv, "")
// Default advertise pin is CLIClientVersion; CLIStableVersion is only the floor.
require.Equal(t, CLIClientVersion, ResolveCLIVersion())
require.True(t, IsSupportedCLIVersion(CLIClientVersion))
require.True(t, IsSupportedCLIVersion(CLIStableVersion))
}
func TestResolveCLIVersionAcceptsValidOverride(t *testing.T) {
t.Setenv(CLIVersionEnv, "0.2.95-alpha.1")
require.Equal(t, "0.2.95-alpha.1", ResolveCLIVersion())
}
func TestResolveCLIVersionRejectsUnsafeOrTooOld(t *testing.T) {
for _, version := range []string{
"0.2.92",
"0.2.93-beta.1",
"0.2.95\r\nX-Injected: true",
"0.2.093",
"0.3",
"1",
} {
t.Run(version, func(t *testing.T) {
t.Setenv(CLIVersionEnv, version)
require.Equal(t, CLIClientVersion, ResolveCLIVersion())
})
}
}
func TestApplyCLIProxyHeaders(t *testing.T) {
t.Setenv(CLIVersionEnv, "")
req, err := http.NewRequest(http.MethodPost, "https://cli-chat-proxy.grok.com/v1/responses", nil)
require.NoError(t, err)
req.Header.Set("User-Agent", "sub2api-grok/1.0")
ApplyCLIProxyHeaders(req)
require.Equal(t, CLIClientVersion, req.Header.Get("x-grok-client-version"))
require.Equal(t, CLIClientIdentifier, req.Header.Get("x-grok-client-identifier"))
require.Equal(t, CLITokenAuth, req.Header.Get("X-XAI-Token-Auth"))
require.Equal(t, CLIUserAgent(CLIClientVersion), req.Header.Get("User-Agent"))
}
func TestApplyCLIProxyHeadersLeavesAPIHostUnchanged(t *testing.T) {
t.Setenv(CLIVersionEnv, "0.2.95")
req, err := http.NewRequest(http.MethodPost, "https://api.x.ai/v1/responses", nil)
require.NoError(t, err)
req.Header.Set("User-Agent", "sub2api-grok/1.0")
ApplyCLIProxyHeaders(req)
require.Empty(t, req.Header.Get("x-grok-client-version"))
require.Empty(t, req.Header.Get("x-grok-client-identifier"))
require.Empty(t, req.Header.Get("X-XAI-Token-Auth"))
require.Equal(t, "sub2api-grok/1.0", req.Header.Get("User-Agent"))
}
+307
View File
@@ -0,0 +1,307 @@
package xai
import (
"strings"
"sync/atomic"
)
// runtimeMappingOpts holds operator-configured defaults applied when Grok
// accounts leave credentials.model_mapping empty. Updated from settings.
var runtimeMappingOpts atomic.Value // ModelMappingOptions
var runtimeMappingVersion atomic.Uint64
func init() {
runtimeMappingOpts.Store(ModelMappingOptions{})
runtimeMappingVersion.Store(1)
}
// SetRuntimeModelMappingOptions updates process-wide defaults used by
// DefaultModelMapping (e.g. after settings load). Safe for concurrent use.
func SetRuntimeModelMappingOptions(opts ModelMappingOptions) {
runtimeMappingOpts.Store(opts)
runtimeMappingVersion.Add(1)
}
// RuntimeModelMappingVersion changes whenever runtime mapping options change.
// Account-level caches include it so settings updates take effect without a restart.
func RuntimeModelMappingVersion() uint64 {
return runtimeMappingVersion.Load()
}
// RuntimeModelMappingOptions returns the last options set via SetRuntimeModelMappingOptions.
func RuntimeModelMappingOptions() ModelMappingOptions {
if v := runtimeMappingOpts.Load(); v != nil {
if opts, ok := v.(ModelMappingOptions); ok {
return opts
}
}
return ModelMappingOptions{}
}
// Model describes an xAI model in OpenAI-compatible /models shape.
type Model struct {
ID string `json:"id"`
Object string `json:"object"`
Type string `json:"type,omitempty"`
Created int64 `json:"created,omitempty"`
OwnedBy string `json:"owned_by"`
DisplayName string `json:"display_name,omitempty"`
}
// DefaultTextModel is the built-in fallback for empty model fields and Grok
// text aliases (e.g. "grok", "grok-latest"). Operators may override the runtime
// default via settings key grok_default_text_model.
const DefaultTextModel = "grok-4.5"
// Official Imagine model IDs (https://docs.x.ai/docs/models).
const (
DefaultImagineImageQualityModel = "grok-imagine-image-quality"
DefaultImagineImageFastModel = "grok-imagine-image"
DefaultImagineVideoModel = "grok-imagine-video"
DefaultImagineVideo15LegacyModel = "grok-imagine-video-1.5"
DefaultImagineVideo15Model = "grok-imagine-video-1.5-preview"
)
// ModelMappingOptions controls optional expansions of the default mapping.
// Cross-client wildcards (gpt-*/claude-*) default ON via settings
// grok_cross_client_model_map_enabled so Codex/Claude clients keep working
// against Grok groups (map to DefaultText / grok-4.5). Operators may disable.
type ModelMappingOptions struct {
// DefaultText is the target for empty models and optional cross-client maps.
// Empty → DefaultTextModel (grok-4.5).
DefaultText string
// EnableCrossClientMap merges gpt-*/codex-*/o*/claude-* → DefaultText.
EnableCrossClientMap bool
}
func (o ModelMappingOptions) defaultText() string {
if t := strings.TrimSpace(o.DefaultText); t != "" {
return t
}
return DefaultTextModel
}
var defaultModels = []Model{
// Text
{ID: "grok-4.6", Object: "model", Type: "model", OwnedBy: "xai", DisplayName: "Grok 4.6"},
{ID: "grok-4.5", Object: "model", Type: "model", OwnedBy: "xai", DisplayName: "Grok 4.5"},
{ID: "grok-4.3", Object: "model", Type: "model", OwnedBy: "xai", DisplayName: "Grok 4.3"},
{ID: "grok-3-mini", Object: "model", Type: "model", OwnedBy: "xai", DisplayName: "Grok 3 Mini"},
{ID: "grok-3-mini-fast", Object: "model", Type: "model", OwnedBy: "xai", DisplayName: "Grok 3 Mini Fast"},
{ID: "grok-build-0.1", Object: "model", Type: "model", OwnedBy: "xai", DisplayName: "Grok Build 0.1"},
{ID: "grok-composer-2.5-fast", Object: "model", Type: "model", OwnedBy: "xai", DisplayName: "Grok Composer 2.5 Fast"},
{ID: "grok-4.20-0309-reasoning", Object: "model", Type: "model", OwnedBy: "xai", DisplayName: "Grok 4.20 Reasoning"},
{ID: "grok-4.20-0309-non-reasoning", Object: "model", Type: "model", OwnedBy: "xai", DisplayName: "Grok 4.20 Non Reasoning"},
{ID: "grok-4.20-multi-agent-0309", Object: "model", Type: "model", OwnedBy: "xai", DisplayName: "Grok 4.20 Multi Agent"},
// Imagine
{ID: DefaultImagineImageQualityModel, Object: "model", Type: "model", OwnedBy: "xai", DisplayName: "Grok Imagine Image Quality"},
{ID: DefaultImagineImageFastModel, Object: "model", Type: "model", OwnedBy: "xai", DisplayName: "Grok Imagine Image"},
{ID: DefaultImagineVideoModel, Object: "model", Type: "model", OwnedBy: "xai", DisplayName: "Grok Imagine Video"},
{ID: DefaultImagineVideo15Model, Object: "model", Type: "model", OwnedBy: "xai", DisplayName: "Grok Imagine Video 1.5 Preview"},
{ID: DefaultImagineVideo15LegacyModel, Object: "model", Type: "model", OwnedBy: "xai", DisplayName: "Grok Imagine Video 1.5 Legacy"},
}
// grokTextResponsesModelAliases is the source of truth for Grok text models
// accepted by the Responses path: client-facing / undated aliases → canonical
// upstream ID. Used by DefaultModelMapping and IsGrokTextResponsesModelID.
var grokTextResponsesModelAliases = map[string]string{
"grok": DefaultTextModel,
"grok-latest": DefaultTextModel,
"grok-4.6": "grok-4.6",
"grok-4.6-latest": "grok-4.6",
"grok-4.5": DefaultTextModel,
"grok-4.5-latest": DefaultTextModel,
"grok-4.3": "grok-4.3",
"grok-4.3-latest": "grok-4.3",
"grok-3-mini": "grok-3-mini",
"grok-3-mini-fast": "grok-3-mini-fast",
"grok-build": "grok-build-0.1",
"grok-build-latest": DefaultTextModel,
"grok-build-0.1": "grok-build-0.1",
"grok-composer-2.5-fast": "grok-composer-2.5-fast",
"grok-composer": "grok-composer-2.5-fast",
"composer-2.5": "grok-composer-2.5-fast",
"grok-4.20-reasoning": "grok-4.20-0309-reasoning",
"grok-4.20-0309-reasoning": "grok-4.20-0309-reasoning",
"grok-4.20-non-reasoning": "grok-4.20-0309-non-reasoning",
"grok-4.20-0309-non-reasoning": "grok-4.20-0309-non-reasoning",
"grok-4.20-multi-agent": "grok-4.20-multi-agent-0309",
"grok-4.20-multi-agent-latest": "grok-4.20-multi-agent-0309",
"grok-4.20-multi-agent-0309": "grok-4.20-multi-agent-0309",
}
func DefaultModels() []Model {
out := make([]Model, len(defaultModels))
copy(out, defaultModels)
return out
}
func DefaultModelIDs() []string {
models := DefaultModels()
ids := make([]string, 0, len(models))
for _, model := range models {
ids = append(ids, model.ID)
}
return ids
}
// DefaultModelMapping returns native Grok/Imagine identity + aliases, using
// runtime options (default text model / optional cross-client wildcards).
// Does NOT enable gpt-*/claude-* unless SetRuntimeModelMappingOptions enables them.
func DefaultModelMapping() map[string]string {
return ModelMappingWithOptions(RuntimeModelMappingOptions())
}
// ModelMappingWithOptions builds the default Grok mapping with optional
// cross-client wildcards and a configurable default text model.
func ModelMappingWithOptions(opts ModelMappingOptions) map[string]string {
defaultText := opts.defaultText()
mapping := make(map[string]string, len(defaultModels)+len(grokTextResponsesModelAliases)+48)
for _, model := range defaultModels {
mapping[model.ID] = model.ID
}
for alias, canonical := range grokTextResponsesModelAliases {
// Remap aliases that pointed at DefaultTextModel constant to runtime default.
if canonical == DefaultTextModel {
mapping[alias] = defaultText
} else {
mapping[alias] = canonical
}
}
// Imagine aliases / legacy IDs → official catalog.
mapping["grok-imagine"] = DefaultImagineImageQualityModel
mapping["grok-imagine-1"] = DefaultImagineImageQualityModel
// Backward-compatible client alias; xAI exposes image editing through the
// image-quality model rather than a separate grok-imagine-edit model.
mapping["grok-imagine-edit"] = DefaultImagineImageQualityModel
mapping["grok-imagine-image"] = DefaultImagineImageFastModel
mapping["grok-imagine-image-quality"] = DefaultImagineImageQualityModel
// Keep official IDs as identity so client-requested model strings are not
// rewritten on the wire (pricing still canonicalizes 1.5* via CanonicalImagineVideoModel).
mapping["grok-imagine-video"] = DefaultImagineVideoModel
mapping["grok-imagine-video-1.5"] = DefaultImagineVideo15LegacyModel
mapping["grok-imagine-video-1.5-preview"] = DefaultImagineVideo15Model
// Informal alias only:
mapping["grok-video-1.5"] = DefaultImagineVideo15Model
if opts.EnableCrossClientMap {
// Codex / OpenAI Responses client defaults (wildcard patterns).
mapping["gpt-*"] = defaultText
mapping["codex-*"] = defaultText
mapping["o1*"] = defaultText
mapping["o3*"] = defaultText
mapping["o4*"] = defaultText
// Claude Code defaults when operators intentionally enable bridging.
mapping["claude-*"] = defaultText
}
addGrokProviderPrefixedMappings(mapping)
return mapping
}
func addGrokProviderPrefixedMappings(mapping map[string]string) {
snapshot := make(map[string]string, len(mapping))
for key, value := range mapping {
snapshot[key] = value
}
for key, value := range snapshot {
if !isGrokNativeOrAlias(key) {
continue
}
for _, prefix := range []string{"xai/", "x-ai/", "grok/"} {
mapping[prefix+key] = value
}
}
}
func isGrokNativeOrAlias(model string) bool {
model = strings.ToLower(strings.TrimSpace(model))
if model == "" || strings.Contains(model, "*") {
return false
}
return strings.HasPrefix(model, "grok") ||
strings.HasPrefix(model, "imagine") ||
strings.HasPrefix(model, "composer")
}
// StripGrokProviderPrefix removes common provider prefixes accepted for
// xAI/Grok models, returning the native model ID.
func StripGrokProviderPrefix(model string) string {
trimmed := strings.TrimSpace(model)
lower := strings.ToLower(trimmed)
for _, prefix := range []string{"xai/", "x-ai/", "grok/"} {
if strings.HasPrefix(lower, prefix) {
return strings.TrimSpace(trimmed[len(prefix):])
}
}
return trimmed
}
// IsGrokModelID reports whether model looks like a native Grok/xAI model id
// (including aliases). Claude/OpenAI model names return false.
func IsGrokModelID(model string) bool {
normalized := strings.ToLower(StripGrokProviderPrefix(model))
if normalized == "" {
return false
}
if strings.HasPrefix(normalized, "grok") {
return true
}
if strings.HasPrefix(normalized, "imagine") {
return true
}
return false
}
// IsGrokTextResponsesModelID reports whether model is a known Grok text model
// for the Responses API. Imagine image/video and unknown custom ids return false.
func IsGrokTextResponsesModelID(model string) bool {
normalized := strings.ToLower(StripGrokProviderPrefix(model))
_, ok := grokTextResponsesModelAliases[normalized]
return ok
}
// ResolveGrokTextResponsesModelID canonicalizes a Grok text alias before upstream.
// empty or bare aliases that resolve via DefaultTextModel use defaultText when set.
func ResolveGrokTextResponsesModelID(model string, defaultText ...string) string {
fallback := DefaultTextModel
if len(defaultText) > 0 && strings.TrimSpace(defaultText[0]) != "" {
fallback = strings.TrimSpace(defaultText[0])
}
trimmed := strings.TrimSpace(model)
if trimmed == "" {
return fallback
}
normalized := strings.ToLower(StripGrokProviderPrefix(trimmed))
if canonical, ok := grokTextResponsesModelAliases[normalized]; ok {
if canonical == DefaultTextModel {
return fallback
}
return canonical
}
return StripGrokProviderPrefix(trimmed)
}
// ResolveDefaultTextModel returns defaultText (or DefaultTextModel) when model is empty.
func ResolveDefaultTextModel(model string, defaultText ...string) string {
if trimmed := strings.TrimSpace(model); trimmed != "" {
return trimmed
}
if len(defaultText) > 0 && strings.TrimSpace(defaultText[0]) != "" {
return strings.TrimSpace(defaultText[0])
}
return DefaultTextModel
}
// CanonicalImagineVideoModel normalizes video model ids for pricing tables.
// Legacy "grok-imagine-video-1.5" shares the 1.5 price family with preview.
func CanonicalImagineVideoModel(model string) string {
m := strings.ToLower(StripGrokProviderPrefix(model))
switch {
case m == "" || m == DefaultImagineVideoModel || m == "grok-imagine-video-preview":
return DefaultImagineVideoModel
case strings.HasPrefix(m, "grok-imagine-video-1.5") || m == "grok-video-1.5":
return DefaultImagineVideo15Model
default:
return m
}
}
+74
View File
@@ -0,0 +1,74 @@
package xai
import (
"testing"
"github.com/stretchr/testify/require"
)
func TestDefaultModelMappingExcludesCrossClientWildcards(t *testing.T) {
original := RuntimeModelMappingOptions()
t.Cleanup(func() { SetRuntimeModelMappingOptions(original) })
SetRuntimeModelMappingOptions(ModelMappingOptions{})
mapping := DefaultModelMapping()
require.Equal(t, "grok-4.5", mapping["grok"])
require.Equal(t, "grok-4.5", mapping["grok-latest"])
require.Equal(t, "grok-build-0.1", mapping["grok-build"])
require.Equal(t, DefaultTextModel, mapping["grok-build-latest"])
require.Equal(t, DefaultImagineImageQualityModel, mapping["grok-imagine-edit"])
require.Equal(t, DefaultImagineVideo15LegacyModel, mapping["grok-imagine-video-1.5"])
require.Equal(t, DefaultImagineVideo15Model, mapping["grok-imagine-video-1.5-preview"])
require.Equal(t, "grok-4.5", mapping["xai/grok"])
// Cross-vendor wildcards must stay opt-in.
_, hasGPT := mapping["gpt-*"]
_, hasClaude := mapping["claude-*"]
require.False(t, hasGPT)
require.False(t, hasClaude)
}
func TestModelMappingWithOptionsCrossClient(t *testing.T) {
t.Parallel()
mapping := ModelMappingWithOptions(ModelMappingOptions{
DefaultText: "grok-4.3",
EnableCrossClientMap: true,
})
require.Equal(t, "grok-4.3", mapping["grok"])
require.Equal(t, "grok-4.3", mapping["gpt-*"])
require.Equal(t, "grok-4.3", mapping["claude-*"])
require.Equal(t, "grok-4.3", mapping["codex-*"])
}
func TestCanonicalImagineVideoModel(t *testing.T) {
t.Parallel()
require.Equal(t, DefaultImagineVideoModel, CanonicalImagineVideoModel("grok-imagine-video"))
require.Equal(t, DefaultImagineVideo15Model, CanonicalImagineVideoModel("grok-imagine-video-1.5"))
require.Equal(t, DefaultImagineVideo15Model, CanonicalImagineVideoModel("grok-imagine-video-1.5-preview"))
require.Equal(t, DefaultImagineVideo15Model, CanonicalImagineVideoModel("xai/grok-video-1.5"))
require.Equal(t, "grok-imagine-video-2", CanonicalImagineVideoModel("grok-imagine-video-2"))
}
func TestIsGrokModelID(t *testing.T) {
t.Parallel()
require.True(t, IsGrokModelID("grok-4.5"))
require.True(t, IsGrokModelID("grok-4.6"))
require.True(t, IsGrokModelID("x-ai/grok-4.3"))
require.False(t, IsGrokModelID("gpt-5"))
require.False(t, IsGrokModelID("claude-sonnet-4"))
}
func TestDefaultModelsIncludesGrok46(t *testing.T) {
t.Parallel()
ids := DefaultModelIDs()
require.Contains(t, ids, "grok-4.6")
require.Equal(t, "grok-4.6", ResolveGrokTextResponsesModelID("grok-4.6"))
require.Equal(t, "grok-4.6", ResolveGrokTextResponsesModelID("grok-4.6-latest"))
}
func TestResolveGrokTextResponsesModelID(t *testing.T) {
t.Parallel()
require.Equal(t, "grok-4.5", ResolveGrokTextResponsesModelID(""))
require.Equal(t, "grok-4.3", ResolveGrokTextResponsesModelID("grok", "grok-4.3"))
require.Equal(t, "grok-4.20-multi-agent-0309", ResolveGrokTextResponsesModelID("grok-4.20-multi-agent"))
}
+746
View File
@@ -0,0 +1,746 @@
package xai
import (
"context"
"crypto/rand"
"crypto/sha256"
"encoding/base64"
"encoding/hex"
"errors"
"fmt"
"log/slog"
"net/url"
"os"
"strings"
"sync"
"time"
"github.com/Wei-Shaw/sub2api/internal/pkg/redissession"
"github.com/Wei-Shaw/sub2api/internal/util/logredact"
"github.com/Wei-Shaw/sub2api/internal/util/urlvalidator"
"github.com/redis/go-redis/v9"
)
const (
OAuthIssuer = "https://auth.x.ai"
DiscoveryURL = OAuthIssuer + "/.well-known/openid-configuration"
DefaultAuthorizeURL = OAuthIssuer + "/oauth2/authorize"
DefaultTokenURL = OAuthIssuer + "/oauth2/token"
DefaultBaseURL = "https://api.x.ai/v1"
DefaultCLIBaseURL = "https://cli-chat-proxy.grok.com/v1"
DefaultUSEast1BaseURL = "https://us-east-1.api.x.ai/v1"
DefaultUSWest2BaseURL = "https://us-west-2.api.x.ai/v1"
DefaultEUWest1BaseURL = "https://eu-west-1.api.x.ai/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"
EnvAllowUnsafeURLOverrides = "XAI_ALLOW_UNSAFE_URL_OVERRIDES"
EnvUnsafeAllowHighConcurrency = "XAI_GROK_UNSAFE_ALLOW_CONCURRENCY_GT_ONE"
)
var (
oauthEndpointAllowedHosts = []string{"x.ai", "*.x.ai"}
// *.api.x.ai 覆盖 xAI 区域端点(us-east-1/us-west-2/eu-west-1 等),
// 运营方可在端点间手动切换以规避单点不可用。
baseURLAllowedHosts = []string{"api.x.ai", "*.api.x.ai", "cli-chat-proxy.grok.com"}
)
// OAuthSession stores one PKCE OAuth flow.
type OAuthSession struct {
State string `json:"state"`
CodeVerifier string `json:"code_verifier"`
CodeChallenge string `json:"code_challenge"`
ClientID string `json:"client_id,omitempty"`
Scope string `json:"scope,omitempty"`
ProxyURL string `json:"proxy_url,omitempty"`
RedirectURI string `json:"redirect_uri"`
CreatedAt time.Time `json:"created_at"`
mu sync.Mutex
consumed bool
}
func (s *OAuthSession) TryConsume() bool {
if s == nil {
return false
}
s.mu.Lock()
defer s.mu.Unlock()
if s.consumed {
return false
}
s.consumed = true
return true
}
// SessionStore manages xAI OAuth sessions with an optional Redis backend.
type SessionStore struct {
mu sync.RWMutex
sessions map[string]*OAuthSession
localOnly map[string]struct{}
stopOnce sync.Once
stopCh chan struct{}
remote *redissession.Store
}
type oauthSessionDTO struct {
State string `json:"state"`
CodeVerifier string `json:"code_verifier"`
CodeChallenge string `json:"code_challenge"`
ClientID string `json:"client_id,omitempty"`
Scope string `json:"scope,omitempty"`
ProxyURL string `json:"proxy_url,omitempty"`
RedirectURI string `json:"redirect_uri"`
CreatedAt time.Time `json:"created_at"`
}
func NewSessionStore() *SessionStore {
store := &SessionStore{
sessions: make(map[string]*OAuthSession),
localOnly: make(map[string]struct{}),
stopCh: make(chan struct{}),
}
go store.cleanup()
return store
}
func NewRedisSessionStore(rdb *redis.Client) *SessionStore {
store := NewSessionStore()
if rdb != nil {
store.remote = redissession.New(rdb, "oauth:session:xai", SessionTTL)
}
return store
}
func (s *SessionStore) Set(sessionID string, session *OAuthSession) {
if session == nil {
return
}
var remoteErr error
if s != nil && s.remote != nil {
remoteErr = s.remote.Set(context.Background(), sessionID, oauthSessionDTO{
State: session.State, CodeVerifier: session.CodeVerifier, CodeChallenge: session.CodeChallenge,
ClientID: session.ClientID, Scope: session.Scope, ProxyURL: session.ProxyURL,
RedirectURI: session.RedirectURI, CreatedAt: session.CreatedAt,
})
}
s.mu.Lock()
defer s.mu.Unlock()
s.sessions[sessionID] = session
if remoteErr != nil {
s.localOnly[sessionID] = struct{}{}
slog.Warn("xai oauth session Redis write failed; using process-local fallback", "error", remoteErr)
} else {
delete(s.localOnly, sessionID)
}
}
func (s *SessionStore) Get(sessionID string) (*OAuthSession, bool) {
if s.isLocalOnly(sessionID) {
return s.getMemory(sessionID)
}
if s != nil && s.remote != nil {
var dto oauthSessionDTO
ok, err := s.remote.Get(context.Background(), sessionID, &dto)
if err != nil || !ok || time.Since(dto.CreatedAt) > SessionTTL {
return nil, false
}
session := &OAuthSession{
State: dto.State, CodeVerifier: dto.CodeVerifier, CodeChallenge: dto.CodeChallenge,
ClientID: dto.ClientID, Scope: dto.Scope, ProxyURL: dto.ProxyURL,
RedirectURI: dto.RedirectURI, CreatedAt: dto.CreatedAt,
}
s.mu.Lock()
s.sessions[sessionID] = session
s.mu.Unlock()
return session, true
}
return s.getMemory(sessionID)
}
func (s *SessionStore) getMemory(sessionID string) (*OAuthSession, bool) {
s.mu.RLock()
defer s.mu.RUnlock()
session, ok := s.sessions[sessionID]
if !ok {
return nil, false
}
if time.Since(session.CreatedAt) > SessionTTL {
return nil, false
}
return session, true
}
func (s *SessionStore) Delete(sessionID string) {
if s != nil && s.remote != nil {
_ = s.remote.Delete(context.Background(), sessionID)
}
s.mu.Lock()
defer s.mu.Unlock()
delete(s.sessions, sessionID)
delete(s.localOnly, sessionID)
}
func (s *SessionStore) TryConsumeSession(sessionID string) bool {
if s == nil {
return false
}
if s.isLocalOnly(sessionID) {
return s.tryConsumeMemory(sessionID)
}
if s.remote != nil {
ok, err := s.remote.TryConsume(context.Background(), sessionID)
return err == nil && ok
}
return s.tryConsumeMemory(sessionID)
}
func (s *SessionStore) isLocalOnly(sessionID string) bool {
s.mu.RLock()
defer s.mu.RUnlock()
_, ok := s.localOnly[sessionID]
return ok
}
func (s *SessionStore) tryConsumeMemory(sessionID string) bool {
session, ok := s.getMemory(sessionID)
return ok && session.TryConsume()
}
func (s *SessionStore) Stop() {
s.stopOnce.Do(func() {
close(s.stopCh)
})
}
func (s *SessionStore) cleanup() {
ticker := time.NewTicker(5 * time.Minute)
defer ticker.Stop()
for {
select {
case <-s.stopCh:
return
case <-ticker.C:
s.mu.Lock()
for id, session := range s.sessions {
if time.Since(session.CreatedAt) > SessionTTL {
delete(s.sessions, id)
delete(s.localOnly, id)
}
}
s.mu.Unlock()
}
}
}
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)
}
func EffectiveScope() string {
return envOrDefault(EnvScope, DefaultScope)
}
func EffectiveRedirectURI(override string) string {
if trimmed := strings.TrimSpace(override); trimmed != "" {
return trimmed
}
return envOrDefault(EnvRedirectURI, DefaultRedirectURI)
}
func EffectiveBaseURL(override string) string {
if trimmed := strings.TrimSpace(override); trimmed != "" {
return strings.TrimRight(trimmed, "/")
}
return strings.TrimRight(envOrDefault(EnvBaseURL, DefaultBaseURL), "/")
}
func ValidatedBaseURL(override string) (string, error) {
return ValidateBaseURL(EffectiveBaseURL(override))
}
// BaseURLValidator applies the caller's outbound URL trust policy before xAI
// endpoint paths are appended. The service layer uses this for API-key accounts
// so the global security.url_allowlist policy remains the single source of
// truth; OAuth callers keep using the strict trusted-host validator.
type BaseURLValidator func(string) (string, error)
func validatedBaseURLWithValidator(override string, validator BaseURLValidator) (string, error) {
if validator == nil {
return ValidatedBaseURL(override)
}
raw := EffectiveBaseURL(override)
validated, err := validator(raw)
if err != nil {
return "", err
}
return normalizeKnownBaseURLPath(validated)
}
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; OAuth URLs use trusted-host allowlists; API-key base URLs require public HTTPS 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{
AllowPrivate: false,
})
if err != nil {
return "", err
}
return normalizeKnownBaseURLPath(normalized)
}
func ValidateTrustedBaseURL(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)
}
// normalizeKnownBaseURLPath 规范化 base URL 的 path 部分:
// - 官方主机固定使用 /v1 前缀(空 path 自动补齐,其余 path 拒绝);
// - 其他主机保留管理员配置的任意 path 前缀(第三方转发地址常见
// /xxx/v1 之类的路由前缀),空 path 仍按惯例补 /v1。
//
// 所有主机统一禁止 userinfo/query/fragment,并去除尾部斜杠。
func normalizeKnownBaseURLPath(raw string) (string, error) {
parsed, err := url.Parse(raw)
if err != nil || parsed.Scheme == "" || parsed.Host == "" {
return "", errors.New("invalid base URL")
}
if parsed.User != nil {
return "", errors.New("base URL must not include userinfo")
}
if parsed.ForceQuery || parsed.RawQuery != "" {
return "", errors.New("base URL must not include a query")
}
if parsed.Fragment != "" {
return "", errors.New("base URL must not include a fragment")
}
path := strings.TrimRight(parsed.Path, "/")
if path == "" {
parsed.Path = "/v1"
parsed.RawPath = ""
return strings.TrimRight(parsed.String(), "/"), nil
}
if path != "/v1" && IsOfficialBaseURLHost(parsed.Hostname()) {
return "", fmt.Errorf("base URL path must be /v1")
}
parsed.Path = path
parsed.RawPath = ""
return strings.TrimRight(parsed.String(), "/"), nil
}
// IsOfficialBaseURLHost 报告 host 是否属于官方 API / 区域 API / CLI 网关主机。
func IsOfficialBaseURLHost(host string) bool {
host = strings.ToLower(strings.TrimSpace(host))
for _, allowed := range baseURLAllowedHosts {
if strings.HasPrefix(allowed, "*.") {
suffix := strings.TrimPrefix(allowed, "*.")
if host == suffix || strings.HasSuffix(host, "."+suffix) {
return true
}
continue
}
if host == allowed {
return true
}
}
return false
}
// IsParseableBaseURL 报告 raw 是否能解析出 host。
// 供读取路径判定存量脏数据:无法解析的值应回落默认端点,而不是把流量发往未定义目标。
func IsParseableBaseURL(raw string) bool {
trimmed := strings.TrimSpace(raw)
if trimmed == "" {
return false
}
parsed, err := url.Parse(trimmed)
return err == nil && parsed.Host != ""
}
// IsOfficialBaseURL 报告 raw 是否指向官方主机(api.x.ai / *.api.x.ai 区域端点 / CLI 网关),
// 容忍存量凭证中的历史变体(大小写、显式 443 端口、百分号编码 path 等)。
// 无法解析的值一并视为官方,调用方据此回落默认端点而不是把流量发往未定义目标。
func IsOfficialBaseURL(raw string) bool {
trimmed := strings.TrimSpace(raw)
if trimmed == "" {
return true
}
parsed, err := url.Parse(trimmed)
if err != nil || parsed.Host == "" {
return true
}
return IsOfficialBaseURLHost(parsed.Hostname())
}
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
}
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 {
return nil, err
}
return b, nil
}
func GenerateState() (string, error) {
bytes, err := GenerateRandomBytes(32)
if err != nil {
return "", err
}
return hex.EncodeToString(bytes), nil
}
func GenerateNonce() (string, error) {
bytes, err := GenerateRandomBytes(16)
if err != nil {
return "", err
}
return hex.EncodeToString(bytes), nil
}
func GenerateSessionID() (string, error) {
bytes, err := GenerateRandomBytes(16)
if err != nil {
return "", err
}
return hex.EncodeToString(bytes), nil
}
func GenerateCodeVerifier() (string, error) {
bytes, err := GenerateRandomBytes(32)
if err != nil {
return "", err
}
return base64URLEncode(bytes), nil
}
func GenerateCodeChallenge(verifier string) string {
hash := sha256.Sum256([]byte(verifier))
return base64URLEncode(hash[:])
}
func base64URLEncode(data []byte) string {
return strings.TrimRight(base64.URLEncoding.EncodeToString(data), "=")
}
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")
params.Set("client_id", EffectiveClientID())
params.Set("redirect_uri", redirectURI)
params.Set("scope", EffectiveScope())
params.Set("state", state)
params.Set("nonce", nonce)
params.Set("code_challenge", codeChallenge)
params.Set("code_challenge_method", "S256")
params.Set("plan", "generic")
params.Set("referrer", "sub2api")
return fmt.Sprintf("%s?%s", authorizeURL, params.Encode()), nil
}
// AuthorizationInput is a parsed manual OAuth callback input.
type AuthorizationInput struct {
Code string
State string
RequiresState bool
}
// ParseAuthorizationInput accepts a full callback URL, query string, or bare code.
func ParseAuthorizationInput(raw string) AuthorizationInput {
trimmed := strings.TrimSpace(raw)
if trimmed == "" {
return AuthorizationInput{}
}
if parsed, err := url.Parse(trimmed); err == nil && parsed != nil {
values := parsed.Query()
if code := strings.TrimSpace(values.Get("code")); code != "" {
return AuthorizationInput{
Code: code,
State: strings.TrimSpace(values.Get("state")),
RequiresState: true,
}
}
}
queryCandidate := strings.TrimPrefix(trimmed, "?")
if strings.Contains(queryCandidate, "=") {
if values, err := url.ParseQuery(queryCandidate); err == nil {
if code := strings.TrimSpace(values.Get("code")); code != "" {
return AuthorizationInput{
Code: code,
State: strings.TrimSpace(values.Get("state")),
RequiresState: true,
}
}
}
}
return AuthorizationInput{Code: trimmed}
}
func BuildResponsesURL(baseURL string) (string, error) {
return BuildResponsesURLWithValidator(baseURL, nil)
}
func BuildResponsesURLWithValidator(baseURL string, validator BaseURLValidator) (string, error) {
validatedBaseURL, err := validatedBaseURLWithValidator(baseURL, validator)
if err != nil {
return "", fmt.Errorf("invalid base url: %w", err)
}
return validatedBaseURL + "/responses", nil
}
func BuildChatCompletionsURL(baseURL string) (string, error) {
return BuildChatCompletionsURLWithValidator(baseURL, nil)
}
func BuildChatCompletionsURLWithValidator(baseURL string, validator BaseURLValidator) (string, error) {
validatedBaseURL, err := validatedBaseURLWithValidator(baseURL, validator)
if err != nil {
return "", fmt.Errorf("invalid base url: %w", err)
}
return validatedBaseURL + "/chat/completions", nil
}
func BuildImagesGenerationsURL(baseURL string) (string, error) {
return BuildImagesGenerationsURLWithValidator(baseURL, nil)
}
func BuildImagesGenerationsURLWithValidator(baseURL string, validator BaseURLValidator) (string, error) {
validatedBaseURL, err := validatedBaseURLWithValidator(baseURL, validator)
if err != nil {
return "", fmt.Errorf("invalid base url: %w", err)
}
return validatedBaseURL + "/images/generations", nil
}
func BuildImagesEditsURL(baseURL string) (string, error) {
return BuildImagesEditsURLWithValidator(baseURL, nil)
}
func BuildImagesEditsURLWithValidator(baseURL string, validator BaseURLValidator) (string, error) {
validatedBaseURL, err := validatedBaseURLWithValidator(baseURL, validator)
if err != nil {
return "", fmt.Errorf("invalid base url: %w", err)
}
return validatedBaseURL + "/images/edits", nil
}
func BuildVideosGenerationsURL(baseURL string) (string, error) {
return BuildVideosGenerationsURLWithValidator(baseURL, nil)
}
func BuildVideosGenerationsURLWithValidator(baseURL string, validator BaseURLValidator) (string, error) {
validatedBaseURL, err := validatedBaseURLWithValidator(baseURL, validator)
if err != nil {
return "", fmt.Errorf("invalid base url: %w", err)
}
return validatedBaseURL + "/videos/generations", nil
}
func BuildVideosEditsURL(baseURL string) (string, error) {
return BuildVideosEditsURLWithValidator(baseURL, nil)
}
func BuildVideosEditsURLWithValidator(baseURL string, validator BaseURLValidator) (string, error) {
validatedBaseURL, err := validatedBaseURLWithValidator(baseURL, validator)
if err != nil {
return "", fmt.Errorf("invalid base url: %w", err)
}
return validatedBaseURL + "/videos/edits", nil
}
func BuildVideosExtensionsURL(baseURL string) (string, error) {
return BuildVideosExtensionsURLWithValidator(baseURL, nil)
}
func BuildVideosExtensionsURLWithValidator(baseURL string, validator BaseURLValidator) (string, error) {
validatedBaseURL, err := validatedBaseURLWithValidator(baseURL, validator)
if err != nil {
return "", fmt.Errorf("invalid base url: %w", err)
}
return validatedBaseURL + "/videos/extensions", nil
}
func BuildVideoURL(baseURL, requestID string) (string, error) {
return BuildVideoURLWithValidator(baseURL, requestID, nil)
}
func BuildVideoURLWithValidator(baseURL, requestID string, validator BaseURLValidator) (string, error) {
validatedBaseURL, err := validatedBaseURLWithValidator(baseURL, validator)
if err != nil {
return "", fmt.Errorf("invalid base url: %w", err)
}
requestID = strings.TrimSpace(requestID)
if requestID == "" {
return "", fmt.Errorf("request id is required")
}
// requestID 由客户端提供并拼进上游 URL 的 path。PathEscape 之外再要求它不是
// 纯点片段、不含控制字符,保证它只能是一个普通的路径片段。
if requestID == "." || requestID == ".." || strings.ContainsAny(requestID, "\x00\r\n") {
return "", fmt.Errorf("invalid request id")
}
return validatedBaseURL + "/videos/" + url.PathEscape(requestID), nil
}
// TokenResponse represents xAI OAuth token responses.
type TokenResponse struct {
AccessToken string `json:"access_token"`
RefreshToken string `json:"refresh_token,omitempty"`
IDToken string `json:"id_token,omitempty"`
TokenType string `json:"token_type,omitempty"`
ExpiresIn int64 `json:"expires_in,omitempty"`
Scope string `json:"scope,omitempty"`
}
@@ -0,0 +1,35 @@
//go:build unit
package xai
import (
"context"
"testing"
"time"
"github.com/alicebob/miniredis/v2"
"github.com/redis/go-redis/v9"
"github.com/stretchr/testify/require"
)
func TestSessionStoreRedisFallbackIsLimitedToFailedWrites(t *testing.T) {
mr := miniredis.RunT(t)
client := redis.NewClient(&redis.Options{Addr: mr.Addr(), MaxRetries: -1})
t.Cleanup(func() { _ = client.Close() })
store := NewRedisSessionStore(client)
defer store.Stop()
session := func(state string) *OAuthSession { return &OAuthSession{State: state, CreatedAt: time.Now()} }
store.Set("remote", session("remote"))
require.NoError(t, store.remote.Delete(context.Background(), "remote"))
_, ok := store.Get("remote")
require.False(t, ok, "a remote miss must not revive the stale local copy")
mr.Close()
store.Set("local-only", session("local"))
got, ok := store.Get("local-only")
require.True(t, ok)
require.Equal(t, "local", got.State)
require.True(t, store.TryConsumeSession("local-only"))
require.False(t, store.TryConsumeSession("local-only"))
}
+375
View File
@@ -0,0 +1,375 @@
//go:build unit
package xai
import (
"net/url"
"testing"
"github.com/Wei-Shaw/sub2api/internal/util/urlvalidator"
"github.com/stretchr/testify/require"
)
func TestParseAuthorizationInput(t *testing.T) {
t.Parallel()
tests := []struct {
name string
raw string
wantCode string
wantState string
wantRequiresState bool
}{
{
name: "full callback url",
raw: "http://127.0.0.1:56121/callback?code=abc123&state=state456",
wantCode: "abc123",
wantState: "state456",
wantRequiresState: true,
},
{
name: "query string",
raw: "?code=abc123&state=state456",
wantCode: "abc123",
wantState: "state456",
wantRequiresState: true,
},
{
name: "full callback url missing state",
raw: "http://127.0.0.1:56121/callback?code=abc123",
wantCode: "abc123",
wantRequiresState: true,
},
{
name: "query string missing state",
raw: "code=abc123",
wantCode: "abc123",
wantRequiresState: true,
},
{
name: "bare code",
raw: "abc123",
wantCode: "abc123",
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
t.Parallel()
got := ParseAuthorizationInput(tt.raw)
require.Equal(t, tt.wantCode, got.Code)
require.Equal(t, tt.wantState, got.State)
require.Equal(t, tt.wantRequiresState, got.RequiresState)
})
}
}
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, 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)
values := parsed.Query()
require.Equal(t, "https", parsed.Scheme)
require.Equal(t, "auth.example.test", parsed.Host)
require.Equal(t, "/oauth2/authorize", parsed.Path)
require.Equal(t, "code", values.Get("response_type"))
require.Equal(t, "client-id", values.Get("client_id"))
require.Equal(t, "http://127.0.0.1:56121/callback", values.Get("redirect_uri"))
require.Equal(t, "openid profile offline_access api:access", values.Get("scope"))
require.Equal(t, "state", values.Get("state"))
require.Equal(t, "nonce", values.Get("nonce"))
require.Equal(t, "challenge", values.Get("code_challenge"))
require.Equal(t, "S256", values.Get("code_challenge_method"))
require.Equal(t, "generic", values.Get("plan"))
require.Equal(t, "sub2api", values.Get("referrer"))
}
func TestValidateXAIURLsAllowOfficialOAuthAndGatewayHosts(t *testing.T) {
authorizeURL, err := ValidateOAuthEndpointURL(DefaultAuthorizeURL)
require.NoError(t, err)
require.Equal(t, DefaultAuthorizeURL, authorizeURL)
tokenURL, err := ValidateOAuthEndpointURL(DefaultTokenURL)
require.NoError(t, err)
require.Equal(t, DefaultTokenURL, tokenURL)
baseURL, err := ValidateBaseURL(DefaultBaseURL)
require.NoError(t, err)
require.Equal(t, DefaultBaseURL, baseURL)
cliBaseURL, err := ValidateBaseURL(DefaultCLIBaseURL)
require.NoError(t, err)
require.Equal(t, DefaultCLIBaseURL, cliBaseURL)
baseURLNoPath, err := ValidateBaseURL("https://api.x.ai")
require.NoError(t, err)
require.Equal(t, DefaultBaseURL, baseURLNoPath)
chatURL, err := BuildChatCompletionsURL(DefaultCLIBaseURL + "/")
require.NoError(t, err)
require.Equal(t, DefaultCLIBaseURL+"/chat/completions", chatURL)
}
func TestBuildGrokMediaURLs(t *testing.T) {
imagesURL, err := BuildImagesGenerationsURL(DefaultBaseURL + "/")
require.NoError(t, err)
require.Equal(t, DefaultBaseURL+"/images/generations", imagesURL)
editsURL, err := BuildImagesEditsURL(DefaultBaseURL)
require.NoError(t, err)
require.Equal(t, DefaultBaseURL+"/images/edits", editsURL)
videosURL, err := BuildVideosGenerationsURL(DefaultBaseURL)
require.NoError(t, err)
require.Equal(t, DefaultBaseURL+"/videos/generations", videosURL)
videoEditsURL, err := BuildVideosEditsURL(DefaultBaseURL)
require.NoError(t, err)
require.Equal(t, DefaultBaseURL+"/videos/edits", videoEditsURL)
videoExtensionsURL, err := BuildVideosExtensionsURL(DefaultBaseURL)
require.NoError(t, err)
require.Equal(t, DefaultBaseURL+"/videos/extensions", videoExtensionsURL)
videoURL, err := BuildVideoURL(DefaultBaseURL, "req 123")
require.NoError(t, err)
require.Equal(t, DefaultBaseURL+"/videos/req%20123", videoURL)
_, err = BuildVideoURL(DefaultBaseURL, " ")
require.Error(t, err)
}
func TestValidateXAIURLsRejectUntrustedOAuthAndUnsafeBaseURLsByDefault(t *testing.T) {
_, err := ValidateOAuthEndpointURL("https://auth.example.test/oauth2/token")
require.Error(t, err)
_, err = ValidateBaseURL("http://127.0.0.1:8080/v1")
require.Error(t, err)
_, err = ValidateBaseURL("https://api.x.ai/custom")
require.Error(t, err)
}
func TestValidateBaseURLAllowsPublicThirdPartyGrokAPI(t *testing.T) {
baseURL, err := ValidateBaseURL("https://grok.example.test/v1/")
require.NoError(t, err)
require.Equal(t, "https://grok.example.test/v1", baseURL)
_, err = ValidateTrustedBaseURL("https://grok.example.test/v1")
require.Error(t, err)
}
func TestValidateBaseURLPathPrefixPolicy(t *testing.T) {
// 非官方主机保留管理员配置的任意 path 前缀。
prefixed, err := ValidateBaseURL("https://relay.example.test/xai/v1/")
require.NoError(t, err)
require.Equal(t, "https://relay.example.test/xai/v1", prefixed)
deepPrefixed, err := ValidateBaseURL("https://relay.example.test/tenant-a/proxy")
require.NoError(t, err)
require.Equal(t, "https://relay.example.test/tenant-a/proxy", deepPrefixed)
// 空 path 仍按惯例补 /v1,保持既有配置兼容。
rootOnly, err := ValidateBaseURL("https://relay.example.test")
require.NoError(t, err)
require.Equal(t, "https://relay.example.test/v1", rootOnly)
// 官方主机固定 /v1 前缀。
_, err = ValidateBaseURL("https://api.x.ai/xai/v1")
require.Error(t, err)
_, err = ValidateBaseURL("https://cli-chat-proxy.grok.com/other")
require.Error(t, err)
}
func TestIsOfficialBaseURL(t *testing.T) {
official := []string{
"",
" ",
DefaultBaseURL,
DefaultCLIBaseURL,
"https://api.x.ai",
"HTTPS://API.X.AI:443/",
"https://api.x.ai:0443/v1",
"https://api.x.ai/%76%31",
"https://api.x.ai:8443/v1",
"HTTPS://CLI-CHAT-PROXY.GROK.COM:443/%76%31/",
"::invalid::url", // 无法解析的值按官方处理,回落默认端点
}
for _, raw := range official {
require.True(t, IsOfficialBaseURL(raw), "expected official: %q", raw)
}
custom := []string{
"https://relay.example.test/v1",
"https://relay.example.test/xai/v1",
"http://relay.example.test/v1",
"https://grok.com.evil.example.test/v1",
"https://api.x.ai.evil.example.test/v1", // 后缀伪装不属于 *.api.x.ai
}
for _, raw := range custom {
require.False(t, IsOfficialBaseURL(raw), "expected custom: %q", raw)
}
}
func TestRegionalAPIEndpointsAreOfficialAndTrusted(t *testing.T) {
regional := []string{
"https://us-east-1.api.x.ai/v1",
"https://us-west-2.api.x.ai/v1",
"https://eu-west-1.api.x.ai/v1",
}
for _, raw := range regional {
require.True(t, IsOfficialBaseURL(raw), "expected official: %q", raw)
validated, err := ValidateTrustedBaseURL(raw)
require.NoError(t, err, "trusted validation should accept regional endpoint %q", raw)
require.Equal(t, raw, validated)
}
// 区域端点作为官方主机同样强制 /v1 path
_, err := ValidateTrustedBaseURL("https://us-east-1.api.x.ai/other")
require.Error(t, err)
}
func TestValidateBaseURLsRejectEmptyQueryDelimiter(t *testing.T) {
_, err := ValidateBaseURL("https://grok.example.test/v1?")
require.Error(t, err)
_, err = ValidateTrustedBaseURL("https://api.x.ai/v1?")
require.Error(t, err)
}
func TestBuildResponsesURLWithValidatorUsesCallerPolicy(t *testing.T) {
validator := func(raw string) (string, error) {
return urlvalidator.ValidateURLFormat(raw, true)
}
target, err := BuildResponsesURLWithValidator("http://grok.example.test/v1/", validator)
require.NoError(t, err)
require.Equal(t, "http://grok.example.test/v1/responses", target)
}
func TestValidateTrustedBaseURLAcceptsOfficialRegionalHosts(t *testing.T) {
for _, raw := range []string{DefaultUSEast1BaseURL, DefaultUSWest2BaseURL, DefaultEUWest1BaseURL} {
got, err := ValidateTrustedBaseURL(raw)
require.NoError(t, err, raw)
require.Equal(t, raw, got)
}
}
func TestBuildResponsesURLPreservesUnsafeOverrideCustomPath(t *testing.T) {
t.Setenv(EnvAllowUnsafeURLOverrides, "true")
target, err := BuildResponsesURL("http://localhost:8080/custom")
require.NoError(t, err)
require.Equal(t, "http://localhost:8080/custom/responses", target)
}
func TestBuildResponsesURLWithValidatorRejectsBaseURLComponents(t *testing.T) {
permissive := func(raw string) (string, error) { return raw, nil }
tests := []struct {
name string
raw string
}{
{name: "userinfo", raw: "https://user:secret@grok.example.test/v1"},
{name: "query", raw: "https://grok.example.test/v1?token=secret"},
{name: "empty query delimiter", raw: "https://grok.example.test/v1?"},
{name: "fragment", raw: "https://grok.example.test/v1#secret"},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
_, err := BuildResponsesURLWithValidator(tt.raw, permissive)
require.Error(t, err)
require.NotContains(t, err.Error(), "secret")
})
}
}
func TestValidateXAIURLsAllowUnsafeDevOverride(t *testing.T) {
t.Setenv(EnvAllowUnsafeURLOverrides, "true")
tokenURL, err := ValidateOAuthEndpointURL("http://127.0.0.1:8080/oauth2/token")
require.NoError(t, err)
require.Equal(t, "http://127.0.0.1:8080/oauth2/token", tokenURL)
baseURL, err := ValidateBaseURL("http://127.0.0.1:8080/v1/")
require.NoError(t, err)
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")
require.Contains(t, report.ProxyPolicy, "API-key base URLs require public HTTPS")
}
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) {
original := RuntimeModelMappingOptions()
t.Cleanup(func() { SetRuntimeModelMappingOptions(original) })
SetRuntimeModelMappingOptions(ModelMappingOptions{})
mapping := DefaultModelMapping()
require.Equal(t, "grok-4.5", mapping["grok"])
require.Equal(t, "grok-4.5", mapping["grok-latest"])
require.Equal(t, "grok-4.6", mapping["grok-4.6"])
require.Equal(t, "grok-4.6", mapping["grok-4.6-latest"])
require.Equal(t, "grok-4.5", mapping["grok-4.5"])
require.Equal(t, "grok-4.5", mapping["grok-4.5-latest"])
require.Equal(t, "grok-build-0.1", mapping["grok-build"])
require.Equal(t, "grok-4.5", mapping["grok-build-latest"])
require.Equal(t, "grok-composer-2.5-fast", mapping["grok-composer"])
require.Equal(t, "grok-composer-2.5-fast", mapping["composer-2.5"])
require.Equal(t, "grok-4.20-0309-reasoning", mapping["grok-4.20-reasoning"])
require.Equal(t, "grok-4.20-0309-non-reasoning", mapping["grok-4.20-non-reasoning"])
require.Equal(t, "grok-4.20-multi-agent-0309", mapping["grok-4.20-multi-agent-0309"])
require.Equal(t, DefaultImagineImageQualityModel, mapping["grok-imagine"])
require.Equal(t, DefaultImagineImageFastModel, mapping["grok-imagine-image"])
require.Equal(t, DefaultImagineImageQualityModel, mapping["grok-imagine-image-quality"])
require.Equal(t, DefaultImagineImageQualityModel, mapping["grok-imagine-edit"])
require.Equal(t, DefaultImagineVideoModel, mapping["grok-imagine-video"])
require.Equal(t, DefaultImagineVideo15LegacyModel, mapping["grok-imagine-video-1.5"])
require.Equal(t, DefaultImagineVideo15Model, mapping["grok-imagine-video-1.5-preview"])
_, hasGPT := mapping["gpt-*"]
require.False(t, hasGPT, "cross-client wildcards must be opt-in")
}
+263
View File
@@ -0,0 +1,263 @@
package xai
import (
"net/http"
"strconv"
"strings"
"time"
)
// GrokFreeRolling24hTokenLimit is the operator soft-gate nominal Free allowance
// (rolling 24h). Soft-gate default matches this; upstream header limits may
// still report historical 1M/2M Free snapshots.
const GrokFreeRolling24hTokenLimit int64 = 500_000
var grokFreeRolling24hTokenLimits = map[int64]struct{}{
GrokFreeRolling24hTokenLimit: {},
1_000_000: {}, // Observed Free limit variants.
2_000_000: {}, // Legacy Free limit observed before July 2026.
}
func IsGrokFreeRolling24hTokenLimit(limit int64) bool {
_, ok := grokFreeRolling24hTokenLimits[limit]
return ok
}
type QuotaWindow struct {
Limit *int64 `json:"limit,omitempty"`
Remaining *int64 `json:"remaining,omitempty"`
ResetUnix *int64 `json:"reset_unix,omitempty"`
ResetAt string `json:"reset_at,omitempty"`
}
type QuotaSnapshot struct {
Requests *QuotaWindow `json:"requests,omitempty"`
Tokens *QuotaWindow `json:"tokens,omitempty"`
RetryAfterSeconds *int `json:"retry_after_seconds,omitempty"`
SubscriptionTier string `json:"subscription_tier,omitempty"`
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"`
// Model is the upstream id that produced these rate-limit headers.
Model string `json:"model,omitempty"`
// PlanFrom45Responses is inferred from a grok-4.5 Responses window
// (8300/53M = Heavy). Carried across later non-4.5 overwrites.
PlanFrom45Responses string `json:"plan_from_45_responses,omitempty"`
PlanFrom45ResponsesAt string `json:"plan_from_45_responses_at,omitempty"`
}
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",
"x-ratelimit-reset-requests",
"x-ratelimit-limit-tokens",
"x-ratelimit-remaining-tokens",
"x-ratelimit-reset-tokens",
"x-rate-limit-limit-requests",
"x-rate-limit-remaining-requests",
"x-rate-limit-reset-requests",
"x-rate-limit-limit-tokens",
"x-rate-limit-remaining-tokens",
"x-rate-limit-reset-tokens",
"retry-after",
"x-subscription-tier",
"xai-subscription-tier",
"x-xai-subscription-tier",
"x-xai-user-tier",
"xai-user-tier",
"xai-tier",
"x-user-tier",
"x-plan-tier",
"x-subscription-plan",
"x-entitlement-status",
"xai-entitlement-status",
"x-xai-entitlement-status",
"x-xai-user-entitlement-status",
"x-user-entitlement-status",
}
func ParseQuotaHeaders(headers http.Header, statusCode int) *QuotaSnapshot {
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),
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
}
snapshot.SubscriptionTier = firstHeader(headers,
"xai-subscription-tier",
"x-subscription-tier",
"x-xai-subscription-tier",
"x-xai-user-tier",
"xai-user-tier",
"xai-tier",
"x-user-tier",
"x-plan-tier",
"x-subscription-plan",
)
snapshot.EntitlementStatus = firstHeader(headers,
"xai-entitlement-status",
"x-entitlement-status",
"x-xai-entitlement-status",
"x-xai-user-entitlement-status",
"x-user-entitlement-status",
)
for _, name := range quotaHeaderAllowlist {
if value := strings.TrimSpace(headers.Get(name)); value != "" {
snapshot.Headers[name] = value
}
}
if snapshot.Requests == nil &&
snapshot.Tokens == nil &&
snapshot.RetryAfterSeconds == nil &&
snapshot.SubscriptionTier == "" &&
snapshot.EntitlementStatus == "" &&
len(snapshot.Headers) == 0 {
if keepEmpty {
return snapshot
}
return nil
}
snapshot.HeadersObserved = true
snapshot.LastHeadersSeenAt = now
return snapshot
}
func parseQuotaWindow(headers http.Header, dimension string) *QuotaWindow {
limitHeader := firstHeader(headers,
"x-ratelimit-limit-"+dimension,
"x-rate-limit-limit-"+dimension,
)
remainingHeader := firstHeader(headers,
"x-ratelimit-remaining-"+dimension,
"x-rate-limit-remaining-"+dimension,
)
resetHeader := firstHeader(headers,
"x-ratelimit-reset-"+dimension,
"x-rate-limit-reset-"+dimension,
)
window := &QuotaWindow{
Limit: parseInt64Ptr(limitHeader),
Remaining: parseInt64Ptr(remainingHeader),
}
if reset := parseResetHeader(resetHeader); reset != nil {
window.ResetUnix = reset
window.ResetAt = time.Unix(*reset, 0).UTC().Format(time.RFC3339)
}
if window.Limit == nil && window.Remaining == nil && window.ResetUnix == nil {
return nil
}
return window
}
func parseResetHeader(raw string) *int64 {
raw = strings.TrimSpace(raw)
if raw == "" {
return nil
}
if value, err := strconv.ParseInt(raw, 10, 64); err == nil {
// xAI (and OpenAI-compatible upstreams) may express the reset as a
// millisecond epoch, a second epoch, or a *relative* number of seconds
// until reset (e.g. "60"). Disambiguate by magnitude, mirroring the
// Kiro reset parser, so a relative "60" is not misread as 1970-01-01.
switch {
case value >= 1_000_000_000_000: // milliseconds epoch → seconds
value = value / 1000
case value >= 1_000_000_000: // already a plausible unix-seconds epoch (>= 2001-09)
// keep as-is
default: // relative seconds from now
value = time.Now().Unix() + value
}
return &value
}
if duration, err := time.ParseDuration(raw); err == nil && duration > 0 {
if duration < time.Second {
duration = time.Second
}
value := time.Now().Add(duration).Unix()
return &value
}
if t, err := time.Parse(time.RFC3339, raw); err == nil {
value := t.Unix()
return &value
}
return nil
}
func parseRetryAfter(raw string) *int {
raw = strings.TrimSpace(raw)
if raw == "" {
return nil
}
if value, err := strconv.Atoi(raw); err == nil {
return &value
}
if t, err := http.ParseTime(raw); err == nil {
seconds := int(time.Until(t).Seconds())
if seconds < 0 {
seconds = 0
}
return &seconds
}
return nil
}
func parseInt64Ptr(raw string) *int64 {
raw = strings.TrimSpace(raw)
if raw == "" {
return nil
}
value, err := strconv.ParseInt(raw, 10, 64)
if err != nil {
return nil
}
return &value
}
func firstHeader(headers http.Header, names ...string) string {
for _, name := range names {
if value := strings.TrimSpace(headers.Get(name)); value != "" {
return value
}
}
return ""
}
+172
View File
@@ -0,0 +1,172 @@
//go:build unit
package xai
import (
"net/http"
"testing"
"time"
"github.com/stretchr/testify/require"
)
func TestParseQuotaHeaders(t *testing.T) {
t.Parallel()
headers := http.Header{}
headers.Set("x-ratelimit-limit-requests", "100")
headers.Set("x-ratelimit-remaining-requests", "25")
headers.Set("x-ratelimit-reset-requests", "1893456000")
headers.Set("x-ratelimit-limit-tokens", "1000000")
headers.Set("x-ratelimit-remaining-tokens", "750000")
headers.Set("retry-after", "60")
headers.Set("xai-subscription-tier", "supergrok")
headers.Set("xai-entitlement-status", "active")
headers.Set("authorization", "should-not-be-copied")
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)
require.Equal(t, "2030-01-01T00:00:00Z", snapshot.Requests.ResetAt)
require.Equal(t, int64(1000000), *snapshot.Tokens.Limit)
require.Equal(t, int64(750000), *snapshot.Tokens.Remaining)
require.Equal(t, 60, *snapshot.RetryAfterSeconds)
require.Equal(t, "supergrok", snapshot.SubscriptionTier)
require.Equal(t, "active", snapshot.EntitlementStatus)
require.Contains(t, snapshot.Headers, "x-ratelimit-limit-requests")
require.NotContains(t, snapshot.Headers, "authorization")
}
func TestParseQuotaHeadersAcceptsXAITierAliases(t *testing.T) {
t.Parallel()
headers := http.Header{}
headers.Set("x-xai-user-tier", "supergrok-heavy")
headers.Set("x-xai-user-entitlement-status", "enabled")
snapshot := ParseQuotaHeaders(headers, http.StatusOK)
require.NotNil(t, snapshot)
require.True(t, snapshot.HeadersObserved)
require.Equal(t, "supergrok-heavy", snapshot.SubscriptionTier)
require.Equal(t, "enabled", snapshot.EntitlementStatus)
require.Equal(t, "supergrok-heavy", snapshot.Headers["x-xai-user-tier"])
require.Equal(t, "enabled", snapshot.Headers["x-xai-user-entitlement-status"])
}
func TestParseQuotaHeadersAcceptsRateLimitAliases(t *testing.T) {
t.Parallel()
headers := http.Header{}
headers.Set("x-rate-limit-limit-tokens", "500000")
headers.Set("x-rate-limit-remaining-tokens", "100")
headers.Set("x-rate-limit-reset-tokens", "1893456000")
snapshot := ParseQuotaHeaders(headers, http.StatusOK)
require.NotNil(t, snapshot)
require.NotNil(t, snapshot.Tokens)
require.Equal(t, int64(500000), *snapshot.Tokens.Limit)
require.Equal(t, int64(100), *snapshot.Tokens.Remaining)
require.Equal(t, int64(1893456000), *snapshot.Tokens.ResetUnix)
require.Contains(t, snapshot.Headers, "x-rate-limit-limit-tokens")
}
func TestParseResetHeaderRelativeSecondsNotMisreadAsEpoch(t *testing.T) {
t.Parallel()
headers := http.Header{}
// xAI may return the reset window as a relative number of seconds ("60").
// It must resolve to ~now+60s, NOT 1970-01-01 (epoch 60).
headers.Set("x-ratelimit-reset-requests", "60")
headers.Set("x-ratelimit-remaining-requests", "0")
before := time.Now().Unix()
snapshot := ParseQuotaHeaders(headers, http.StatusTooManyRequests)
require.NotNil(t, snapshot)
require.NotNil(t, snapshot.Requests)
require.NotNil(t, snapshot.Requests.ResetUnix)
got := *snapshot.Requests.ResetUnix
require.GreaterOrEqual(t, got, before+59)
require.LessOrEqual(t, got, time.Now().Unix()+61)
}
func TestParseResetHeaderDurationWindow(t *testing.T) {
t.Parallel()
headers := http.Header{}
headers.Set("x-ratelimit-reset-requests", "6m0s")
headers.Set("x-ratelimit-remaining-requests", "0")
before := time.Now().Unix()
snapshot := ParseQuotaHeaders(headers, http.StatusTooManyRequests)
require.NotNil(t, snapshot)
require.NotNil(t, snapshot.Requests)
require.NotNil(t, snapshot.Requests.ResetUnix)
got := *snapshot.Requests.ResetUnix
require.GreaterOrEqual(t, got, before+359)
require.LessOrEqual(t, got, time.Now().Unix()+361)
}
func TestParseResetHeaderSubsecondDurationCeilsToFutureSecond(t *testing.T) {
t.Parallel()
headers := http.Header{}
headers.Set("x-rate-limit-reset-tokens", "250ms")
headers.Set("x-rate-limit-remaining-tokens", "0")
before := time.Now().Unix()
snapshot := ParseQuotaHeaders(headers, http.StatusTooManyRequests)
require.NotNil(t, snapshot)
require.NotNil(t, snapshot.Tokens)
require.NotNil(t, snapshot.Tokens.ResetUnix)
require.GreaterOrEqual(t, *snapshot.Tokens.ResetUnix, before)
require.LessOrEqual(t, *snapshot.Tokens.ResetUnix, time.Now().Unix()+2)
}
func TestParseResetHeaderMillisecondsEpochNormalized(t *testing.T) {
t.Parallel()
headers := http.Header{}
headers.Set("x-ratelimit-reset-tokens", "1893456000000") // ms epoch
headers.Set("x-ratelimit-remaining-tokens", "0")
snapshot := ParseQuotaHeaders(headers, http.StatusTooManyRequests)
require.NotNil(t, snapshot)
require.NotNil(t, snapshot.Tokens)
require.Equal(t, int64(1893456000), *snapshot.Tokens.ResetUnix)
}
func TestParseQuotaHeadersReturnsNilForMissingHeaders(t *testing.T) {
t.Parallel()
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)
}
func TestIsGrokFreeRolling24hTokenLimit(t *testing.T) {
t.Parallel()
require.True(t, IsGrokFreeRolling24hTokenLimit(GrokFreeRolling24hTokenLimit))
require.True(t, IsGrokFreeRolling24hTokenLimit(500_000))
require.True(t, IsGrokFreeRolling24hTokenLimit(1_000_000), "observed Free limit variants remain classifiable")
require.True(t, IsGrokFreeRolling24hTokenLimit(2_000_000), "legacy snapshots remain classifiable")
require.False(t, IsGrokFreeRolling24hTokenLimit(3_000_000))
}
+448
View File
@@ -0,0 +1,448 @@
package xai
import (
"context"
"encoding/base64"
"encoding/json"
"errors"
"fmt"
"io"
"net/http"
"net/http/cookiejar"
"net/url"
"sort"
"strconv"
"strings"
"time"
)
const (
SSOBuildScope = "openid profile email offline_access grok-cli:access api:access conversations:read conversations:write"
SSOAccountsURL = "https://accounts.x.ai/"
SSODeviceURL = OAuthIssuer + "/oauth2/device/code"
SSOVerifyURL = OAuthIssuer + "/oauth2/device/verify"
SSOApproveURL = OAuthIssuer + "/oauth2/device/approve"
SSOTokenURL = OAuthIssuer + "/oauth2/token"
SSOConversionTimeout = 90 * time.Second
ssoMaxAuthBody = 2 << 20
ssoMaxTokenLength = 16 << 10
ssoDefaultUA = "Mozilla/5.0 (Macintosh; Intel Mac OS X 10_15_7) AppleWebKit/537.36 (KHTML, like Gecko) Chrome/131.0.0.0 Safari/537.36"
ssoDefaultTokenTTL = 6 * time.Hour
)
var (
ErrSSOUnauthorized = errors.New("xai sso unauthorized")
ErrSSOAuthorizationDenied = errors.New("xai device authorization denied")
)
type SSOHTTPError struct{ Status int }
func (e SSOHTTPError) Error() string { return fmt.Sprintf("xAI OAuth HTTP %d", e.Status) }
type SSODeviceHTTPClient interface {
Do(*http.Request) (*http.Response, error)
}
type SSODeviceOptions struct {
HTTPClient SSODeviceHTTPClient
UserAgent string
Sleep func(context.Context, time.Duration) error
}
type ssoDeviceFlow struct {
client SSODeviceHTTPClient
userAgent string
cookieJar http.CookieJar
sleep func(context.Context, time.Duration) error
}
func ConvertSSOToBuild(ctx context.Context, ssoToken string, opts *SSODeviceOptions) (*TokenResponse, error) {
ssoToken = NormalizeSSOToken(ssoToken)
if ssoToken == "" {
return nil, ErrSSOUnauthorized
}
if opts == nil {
opts = &SSODeviceOptions{}
}
client := opts.HTTPClient
if client == nil {
client = &http.Client{
Timeout: SSOConversionTimeout,
CheckRedirect: func(*http.Request, []*http.Request) error {
return http.ErrUseLastResponse
},
}
}
userAgent := strings.TrimSpace(opts.UserAgent)
if userAgent == "" {
userAgent = ssoDefaultUA
}
sleep := opts.Sleep
if sleep == nil {
sleep = sleepContext
}
jar, err := cookiejar.New(nil)
if err != nil {
return nil, err
}
seedSSOCookies(jar, ssoToken)
flow := &ssoDeviceFlow{
client: client,
userAgent: userAgent,
cookieJar: jar,
sleep: sleep,
}
return flow.convert(ctx)
}
func (f *ssoDeviceFlow) convert(ctx context.Context) (*TokenResponse, error) {
status, finalURL, _, err := f.do(ctx, http.MethodGet, SSOAccountsURL, nil)
if err != nil {
return nil, err
}
if status == http.StatusUnauthorized || strings.Contains(finalURL, "sign-in") || strings.Contains(finalURL, "sign-up") {
return nil, ErrSSOUnauthorized
}
if status < 200 || status >= 400 {
return nil, fmt.Errorf("validate Grok Web SSO: %w", SSOHTTPError{Status: status})
}
status, _, body, err := f.do(ctx, http.MethodPost, SSODeviceURL, url.Values{
"client_id": {DefaultClientID},
"scope": {SSOBuildScope},
})
if err != nil {
return nil, err
}
if status < 200 || status >= 300 {
return nil, fmt.Errorf("start xAI device flow: %w", SSOHTTPError{Status: status})
}
var device struct {
DeviceCode string `json:"device_code"`
UserCode string `json:"user_code"`
VerificationURIComplete string `json:"verification_uri_complete"`
Interval int `json:"interval"`
ExpiresIn int `json:"expires_in"`
}
if err := json.Unmarshal(body, &device); err != nil {
return nil, fmt.Errorf("parse xAI device flow response: %w", err)
}
if device.DeviceCode == "" || device.UserCode == "" || !safeXAIAuthURL(device.VerificationURIComplete) {
return nil, errors.New("xAI device flow response is incomplete")
}
if device.Interval <= 0 {
device.Interval = 5
}
if device.ExpiresIn <= 0 {
device.ExpiresIn = 1800
}
status, _, _, err = f.do(ctx, http.MethodGet, device.VerificationURIComplete, nil)
if err != nil {
return nil, err
}
if status < 200 || status >= 400 {
return nil, fmt.Errorf("open xAI device verification page: %w", SSOHTTPError{Status: status})
}
status, finalURL, _, err = f.do(ctx, http.MethodPost, SSOVerifyURL, url.Values{"user_code": {device.UserCode}})
if err != nil {
return nil, err
}
if status < 200 || status >= 400 {
return nil, fmt.Errorf("verify xAI device code: %w", SSOHTTPError{Status: status})
}
if !strings.Contains(finalURL, "consent") {
return nil, errors.New("xAI device verification did not reach consent page")
}
status, finalURL, _, err = f.do(ctx, http.MethodPost, SSOApproveURL, url.Values{
"user_code": {device.UserCode},
"action": {"allow"},
"principal_type": {"User"},
"principal_id": {""},
})
if err != nil {
return nil, err
}
if status < 200 || status >= 400 {
return nil, fmt.Errorf("approve xAI device code: %w", SSOHTTPError{Status: status})
}
if !strings.Contains(finalURL, "done") {
return nil, errors.New("xAI device approval did not reach done page")
}
return f.pollToken(ctx, device.DeviceCode, time.Duration(device.Interval)*time.Second, time.Duration(device.ExpiresIn)*time.Second)
}
func (f *ssoDeviceFlow) pollToken(ctx context.Context, deviceCode string, interval, expiresIn time.Duration) (*TokenResponse, error) {
if interval < time.Second {
interval = time.Second
}
deadline := time.Now().Add(minDuration(expiresIn, 75*time.Second))
for time.Now().Before(deadline) {
if err := f.sleep(ctx, interval); err != nil {
return nil, err
}
status, _, body, err := f.do(ctx, http.MethodPost, SSOTokenURL, url.Values{
"grant_type": {"urn:ietf:params:oauth:grant-type:device_code"},
"client_id": {DefaultClientID},
"device_code": {deviceCode},
})
if err != nil {
return nil, err
}
var payload struct {
AccessToken string `json:"access_token"`
RefreshToken string `json:"refresh_token"`
IDToken string `json:"id_token"`
TokenType string `json:"token_type"`
ExpiresIn int64 `json:"expires_in"`
Scope string `json:"scope"`
Error string `json:"error"`
ErrorDescription string `json:"error_description"`
}
if err := json.Unmarshal(body, &payload); err != nil {
return nil, fmt.Errorf("parse xAI token response: %w", err)
}
if status >= 200 && status < 300 && payload.AccessToken != "" {
if payload.ExpiresIn <= 0 {
payload.ExpiresIn = int64(ssoDefaultTokenTTL.Seconds())
}
if payload.TokenType == "" {
payload.TokenType = "Bearer"
}
return &TokenResponse{
AccessToken: payload.AccessToken,
RefreshToken: payload.RefreshToken,
IDToken: payload.IDToken,
TokenType: payload.TokenType,
ExpiresIn: payload.ExpiresIn,
Scope: payload.Scope,
}, nil
}
switch payload.Error {
case "authorization_pending":
continue
case "slow_down":
interval += 5 * time.Second
continue
case "access_denied", "expired_token":
return nil, ErrSSOAuthorizationDenied
default:
if status >= 400 {
return nil, fmt.Errorf("xAI token polling failed (%s): %w", firstNonEmpty(payload.ErrorDescription, payload.Error), SSOHTTPError{Status: status})
}
return nil, fmt.Errorf("xAI token polling failed: %s", firstNonEmpty(payload.ErrorDescription, payload.Error, strconv.Itoa(status)))
}
}
return nil, errors.New("xAI device flow token polling timed out")
}
func (f *ssoDeviceFlow) do(ctx context.Context, method, endpoint string, form url.Values) (int, string, []byte, error) {
if !safeXAIAuthURL(endpoint) {
return 0, "", nil, errors.New("xAI OAuth URL is not trusted")
}
currentURL := endpoint
currentMethod := method
currentForm := form
for redirects := 0; redirects <= 8; redirects++ {
var body io.Reader
if currentForm != nil {
body = strings.NewReader(currentForm.Encode())
}
request, err := http.NewRequestWithContext(ctx, currentMethod, currentURL, body)
if err != nil {
return 0, currentURL, nil, err
}
request.Header.Set("Accept", "application/json, text/html;q=0.9, */*;q=0.8")
request.Header.Set("Accept-Language", "zh-CN,zh;q=0.9,en;q=0.8")
request.Header.Set("User-Agent", f.userAgent)
if cookie := f.cookieHeader(request.URL); cookie != "" {
request.Header.Set("Cookie", cookie)
}
if currentForm != nil {
request.Header.Set("Content-Type", "application/x-www-form-urlencoded")
}
response, err := f.client.Do(request)
if err != nil {
return 0, currentURL, nil, err
}
f.captureCookies(request.URL, response)
data, readErr := io.ReadAll(io.LimitReader(response.Body, ssoMaxAuthBody+1))
_ = response.Body.Close()
if readErr != nil {
return response.StatusCode, currentURL, nil, readErr
}
if len(data) > ssoMaxAuthBody {
return response.StatusCode, currentURL, nil, errors.New("xAI OAuth response exceeds 2 MiB")
}
if response.StatusCode < 300 || response.StatusCode > 399 {
return response.StatusCode, currentURL, data, nil
}
location := strings.TrimSpace(response.Header.Get("Location"))
if location == "" {
return response.StatusCode, currentURL, data, errors.New("xAI OAuth redirect missing Location")
}
base, _ := url.Parse(currentURL)
next, err := url.Parse(location)
if err != nil {
return response.StatusCode, currentURL, data, err
}
currentURL = base.ResolveReference(next).String()
if !safeXAIAuthURL(currentURL) {
return response.StatusCode, currentURL, data, errors.New("xAI OAuth redirected to untrusted host")
}
if response.StatusCode == http.StatusSeeOther || ((response.StatusCode == http.StatusMovedPermanently || response.StatusCode == http.StatusFound) && currentMethod != http.MethodGet && currentMethod != http.MethodHead) {
currentMethod = http.MethodGet
currentForm = nil
}
}
return 0, currentURL, nil, errors.New("xAI OAuth redirected too many times")
}
func seedSSOCookies(jar http.CookieJar, token string) {
if jar == nil {
return
}
for _, rawURL := range []string{SSOAccountsURL, OAuthIssuer + "/"} {
target, err := url.Parse(rawURL)
if err != nil {
continue
}
jar.SetCookies(target, []*http.Cookie{
{Name: "sso", Value: token, Path: "/", Secure: true, HttpOnly: true},
{Name: "sso-rw", Value: token, Path: "/", Secure: true, HttpOnly: true},
})
}
}
func (f *ssoDeviceFlow) captureCookies(requestURL *url.URL, response *http.Response) {
if f == nil || f.cookieJar == nil || requestURL == nil || response == nil {
return
}
cookies := make([]*http.Cookie, 0)
for _, cookie := range response.Cookies() {
name := strings.TrimSpace(cookie.Name)
value := strings.TrimSpace(cookie.Value)
if name == "" || len(name) > 128 || len(value) > 16384 || strings.ContainsAny(name+value, "\r\n\x00") {
continue
}
cookie.Name = name
cookie.Value = value
cookies = append(cookies, cookie)
}
f.cookieJar.SetCookies(requestURL, cookies)
}
func (f *ssoDeviceFlow) cookieHeader(requestURL *url.URL) string {
if f == nil || f.cookieJar == nil || requestURL == nil {
return ""
}
cookies := f.cookieJar.Cookies(requestURL)
sort.Slice(cookies, func(i, j int) bool { return cookies[i].Name < cookies[j].Name })
parts := make([]string, 0, len(cookies))
for _, cookie := range cookies {
parts = append(parts, cookie.Name+"="+cookie.Value)
}
return strings.Join(parts, "; ")
}
func safeXAIAuthURL(raw string) bool {
parsed, err := url.Parse(raw)
if err != nil || parsed.User != nil || parsed.Hostname() == "" {
return false
}
if AllowUnsafeURLOverrides() {
return parsed.Scheme != "" && parsed.Host != ""
}
if parsed.Scheme != "https" {
return false
}
host := strings.ToLower(parsed.Hostname())
return host == "x.ai" || strings.HasSuffix(host, ".x.ai")
}
func NormalizeSSOToken(value string) string {
value = strings.TrimSpace(value)
if strings.HasPrefix(strings.ToLower(value), "cookie:") {
value = strings.TrimSpace(value[len("cookie:"):])
}
for _, part := range strings.Split(value, ";") {
name, token, found := strings.Cut(strings.TrimSpace(part), "=")
if !found {
continue
}
switch strings.ToLower(strings.TrimSpace(name)) {
case "sso", "sso-rw":
return sanitizeSSOToken(token)
}
}
if token, _, found := strings.Cut(value, ";"); found {
value = strings.TrimSpace(token)
}
return sanitizeSSOToken(value)
}
func sanitizeSSOToken(value string) string {
value = strings.NewReplacer("\r", "", "\n", "", "\x00", "").Replace(strings.TrimSpace(value))
if len(value) > ssoMaxTokenLength {
return ""
}
return value
}
func DecodeJWTClaims(token string) map[string]any {
parts := strings.Split(token, ".")
if len(parts) < 2 {
return nil
}
payload, err := base64.RawURLEncoding.DecodeString(parts[1])
if err != nil {
return nil
}
var claims map[string]any
if err := json.Unmarshal(payload, &claims); err != nil {
return nil
}
return claims
}
func JWTClaimString(claims map[string]any, key string) string {
value, _ := claims[key].(string)
return strings.TrimSpace(value)
}
func sleepContext(ctx context.Context, d time.Duration) error {
timer := time.NewTimer(d)
defer timer.Stop()
select {
case <-ctx.Done():
return ctx.Err()
case <-timer.C:
return nil
}
}
func minDuration(a, b time.Duration) time.Duration {
if a <= 0 {
return b
}
if a < b {
return a
}
return b
}
func firstNonEmpty(values ...string) string {
for _, value := range values {
if strings.TrimSpace(value) != "" {
return strings.TrimSpace(value)
}
}
return ""
}
+143
View File
@@ -0,0 +1,143 @@
//go:build unit
package xai
import (
"context"
"io"
"net/http"
"net/http/cookiejar"
"net/url"
"strings"
"testing"
"time"
"github.com/stretchr/testify/require"
)
type ssoDeviceFakeClient struct {
t *testing.T
tokenCalls int
cookieHeaders []string
}
func (c *ssoDeviceFakeClient) Do(req *http.Request) (*http.Response, error) {
c.cookieHeaders = append(c.cookieHeaders, req.Header.Get("Cookie"))
switch req.URL.String() {
case SSOAccountsURL:
require.Equal(c.t, http.MethodGet, req.Method)
return ssoDeviceResponse(http.StatusOK, http.Header{"Set-Cookie": {"session=web-session; Domain=x.ai; Path=/"}}, `{}`), nil
case SSODeviceURL:
require.Equal(c.t, http.MethodPost, req.Method)
values := readSSODeviceForm(c.t, req)
require.Equal(c.t, DefaultClientID, values.Get("client_id"))
require.Equal(c.t, SSOBuildScope, values.Get("scope"))
return ssoDeviceResponse(http.StatusOK, http.Header{"Set-Cookie": {"csrf=csrf-token; Path=/"}}, `{"device_code":"device-1","user_code":"USER-1","verification_uri_complete":"https://auth.x.ai/oauth2/device/complete","interval":1,"expires_in":60}`), nil
case "https://auth.x.ai/oauth2/device/complete":
require.Equal(c.t, http.MethodGet, req.Method)
return ssoDeviceResponse(http.StatusOK, nil, `<html>ok</html>`), nil
case SSOVerifyURL:
require.Equal(c.t, http.MethodPost, req.Method)
values := readSSODeviceForm(c.t, req)
require.Equal(c.t, "USER-1", values.Get("user_code"))
return ssoDeviceResponse(http.StatusFound, http.Header{"Location": {"/oauth2/device/consent"}}, ``), nil
case "https://auth.x.ai/oauth2/device/consent":
require.Equal(c.t, http.MethodGet, req.Method)
return ssoDeviceResponse(http.StatusOK, nil, `<html>consent</html>`), nil
case SSOApproveURL:
require.Equal(c.t, http.MethodPost, req.Method)
values := readSSODeviceForm(c.t, req)
require.Equal(c.t, "USER-1", values.Get("user_code"))
require.Equal(c.t, "allow", values.Get("action"))
require.Equal(c.t, "User", values.Get("principal_type"))
return ssoDeviceResponse(http.StatusSeeOther, http.Header{"Location": {"/oauth2/device/done"}}, ``), nil
case "https://auth.x.ai/oauth2/device/done":
require.Equal(c.t, http.MethodGet, req.Method)
return ssoDeviceResponse(http.StatusOK, nil, `<html>done</html>`), nil
case SSOTokenURL:
require.Equal(c.t, http.MethodPost, req.Method)
c.tokenCalls++
values := readSSODeviceForm(c.t, req)
require.Equal(c.t, "urn:ietf:params:oauth:grant-type:device_code", values.Get("grant_type"))
require.Equal(c.t, "device-1", values.Get("device_code"))
return ssoDeviceResponse(http.StatusOK, nil, `{"access_token":"access-token","refresh_token":"refresh-token","id_token":"id-token","token_type":"Bearer","expires_in":3600,"scope":"`+SSOBuildScope+`"}`), nil
default:
c.t.Fatalf("unexpected request: %s %s", req.Method, req.URL.String())
return nil, nil
}
}
func TestConvertSSOToBuildCompletesDeviceFlow(t *testing.T) {
t.Setenv(EnvClientID, "")
client := &ssoDeviceFakeClient{t: t}
token, err := ConvertSSOToBuild(context.Background(), "sso=sso-token; ignored=1", &SSODeviceOptions{
HTTPClient: client,
Sleep: func(context.Context, time.Duration) error {
return nil
},
})
require.NoError(t, err)
require.Equal(t, "access-token", token.AccessToken)
require.Equal(t, "refresh-token", token.RefreshToken)
require.Equal(t, "id-token", token.IDToken)
require.Equal(t, SSOBuildScope, token.Scope)
require.Equal(t, 1, client.tokenCalls)
require.Contains(t, client.cookieHeaders[0], "sso=sso-token")
require.Contains(t, client.cookieHeaders[0], "sso-rw=sso-token")
require.Contains(t, client.cookieHeaders[len(client.cookieHeaders)-1], "session=web-session")
require.Contains(t, client.cookieHeaders[len(client.cookieHeaders)-1], "csrf=csrf-token")
}
func TestNormalizeSSOTokenAcceptsCookieHeader(t *testing.T) {
require.Equal(t, "token-1", NormalizeSSOToken("Cookie: foo=bar; sso=token-1; sso-rw=token-2"))
require.Equal(t, "token-2", NormalizeSSOToken("sso-rw=token-2; foo=bar"))
require.Equal(t, "raw-token", NormalizeSSOToken(" raw-token ; ignored=1"))
require.Empty(t, NormalizeSSOToken(strings.Repeat("x", ssoMaxTokenLength+1)))
}
func TestSSODeviceCookieJarHonorsDomainAndPath(t *testing.T) {
jar, err := cookiejar.New(nil)
require.NoError(t, err)
flow := &ssoDeviceFlow{cookieJar: jar}
accountsURL, err := url.Parse("https://accounts.x.ai/")
require.NoError(t, err)
authURL, err := url.Parse("https://auth.x.ai/oauth2/device/verify")
require.NoError(t, err)
flow.captureCookies(accountsURL, ssoDeviceResponse(http.StatusOK, http.Header{"Set-Cookie": {
"host-only=accounts; Path=/",
"shared=all-xai; Domain=x.ai; Path=/",
"narrow=oauth-only; Domain=x.ai; Path=/oauth2",
}}, ""))
authCookies := flow.cookieHeader(authURL)
require.NotContains(t, authCookies, "host-only=accounts")
require.Contains(t, authCookies, "shared=all-xai")
require.Contains(t, authCookies, "narrow=oauth-only")
accountsCookies := flow.cookieHeader(accountsURL)
require.Contains(t, accountsCookies, "host-only=accounts")
require.Contains(t, accountsCookies, "shared=all-xai")
require.NotContains(t, accountsCookies, "narrow=oauth-only")
}
func ssoDeviceResponse(status int, header http.Header, body string) *http.Response {
if header == nil {
header = http.Header{}
}
return &http.Response{
StatusCode: status,
Header: header,
Body: io.NopCloser(strings.NewReader(body)),
}
}
func readSSODeviceForm(t *testing.T, req *http.Request) url.Values {
t.Helper()
data, err := io.ReadAll(req.Body)
require.NoError(t, err)
values, err := url.ParseQuery(string(data))
require.NoError(t, err)
return values
}
@@ -0,0 +1,279 @@
package xai
import (
"encoding/json"
"strconv"
"strings"
"time"
)
// GrokQuotaSignalMaxAge bounds how long a grok-4.5 Responses window can
// influence SuperGrok vs Heavy inference.
const GrokQuotaSignalMaxAge = 24 * time.Hour
const (
grok45ResponsesModel = "grok-4.5"
grokHeavyQuotaRequestLimit int64 = 8_300
grokHeavyQuotaTokenLimit int64 = 53_000_000
)
// MapJWTSubscriptionTier maps prod_auth.SubscriptionTier numeric JWT claims
// to stable snake_case keys used by Grok Build / Mixpanel.
func MapJWTSubscriptionTier(tier uint64) string {
switch tier {
case 0:
return "free"
case 1:
return "supergrok"
case 2:
return "x_basic"
case 3:
return "x_premium"
case 4:
return "x_premium_plus"
case 5:
return "supergrok_heavy"
case 6:
return "supergrok_lite"
case 7:
return "supergrok_plus"
default:
return strconv.FormatUint(tier, 10)
}
}
// NormalizeSubscriptionTier canonicalizes display names, /user strings, and
// JWT-derived keys onto the same snake_case identifiers.
func NormalizeSubscriptionTier(raw string) string {
t := strings.ToLower(strings.TrimSpace(raw))
t = strings.ReplaceAll(t, "-", "_")
t = strings.Join(strings.Fields(t), "_")
switch t {
case "free", "grok_free", "grokfree", "free_tier", "freetier", "grok_basic", "grokbasic":
return "free"
case "supergrok", "grokpro":
return "supergrok"
case "supergrok_lite", "supergroklite":
return "supergrok_lite"
case "supergrok_heavy", "supergrokheavy":
return "supergrok_heavy"
case "supergrok_pro", "supergrokpro":
return "supergrok_pro"
case "supergrok_plus", "supergrokplus":
return "supergrok_plus"
case "x_basic", "xbasic", "basic":
return "x_basic"
case "x_premium", "xpremium":
return "x_premium"
case "x_premium_plus", "xpremiumplus", "x_premium+":
return "x_premium_plus"
default:
return t
}
}
// SubscriptionTierFromJWT decodes an access token payload (no signature check)
// and maps the numeric or string `tier` claim.
func SubscriptionTierFromJWT(jwt string) string {
claims := DecodeJWTClaims(jwt)
if claims == nil {
return ""
}
raw, ok := claims["tier"]
if !ok || raw == nil {
return ""
}
switch v := raw.(type) {
case float64:
if v < 0 {
return ""
}
return MapJWTSubscriptionTier(uint64(v))
case json.Number:
n, err := v.Int64()
if err != nil || n < 0 {
return NormalizeSubscriptionTier(v.String())
}
return MapJWTSubscriptionTier(uint64(n))
case string:
trimmed := strings.TrimSpace(v)
if trimmed == "" {
return ""
}
if n, err := strconv.ParseUint(trimmed, 10, 64); err == nil {
return MapJWTSubscriptionTier(n)
}
return NormalizeSubscriptionTier(trimmed)
default:
return ""
}
}
// CanonicalGrokPlan resolves SuperGrok vs Heavy when the provider label is
// ambiguous (SuperGrokPro). JWT numeric claims are applied by the caller first.
// Monthly $150/$1500 limits still win when present.
// Rate-limit windows are only used when they came from grok-4.5 Responses.
func CanonicalGrokPlan(monthlyLimitCents *float64, subscriptionTier string, quota *QuotaSnapshot) string {
if plan := resolvePlan(monthlyLimitCents); plan != "" {
return NormalizeSubscriptionTier(plan)
}
normalized := NormalizeSubscriptionTier(subscriptionTier)
switch normalized {
case "free", "x_basic":
return "free"
case "supergrok_heavy":
return "supergrok_heavy"
case "supergrok_lite":
return "supergrok_lite"
case "supergrok_plus":
return "supergrok_plus"
}
if isAmbiguousGrokPaidPlan(normalized) {
if hint := Grok45ResponsesPlanHint(quota, time.Time{}); hint != "" {
return hint
}
return "supergrok"
}
return ""
}
func isAmbiguousGrokPaidPlan(normalized string) bool {
switch normalized {
case "supergrok", "supergrok_pro", "paid", "pro":
return true
default:
return false
}
}
// IsGrok45ResponsesQuotaModel reports whether model is the grok-4.5 Responses
// id (or a dated grok-4.5-* variant). Empty and other families are false.
func IsGrok45ResponsesQuotaModel(model string) bool {
m := strings.ToLower(strings.TrimSpace(StripGrokProviderPrefix(model)))
return m == grok45ResponsesModel || strings.HasPrefix(m, grok45ResponsesModel+"-")
}
// Grok45ResponsesPlanHint returns SuperGrok / Heavy inferred from a grok-4.5
// Responses window. Other models' limits are ignored.
func Grok45ResponsesPlanHint(quota *QuotaSnapshot, now time.Time) string {
if quota == nil {
return ""
}
if plan := NormalizeSubscriptionTier(quota.PlanFrom45Responses); plan == "supergrok" || plan == "supergrok_heavy" {
if isQuotaTimestampFresh(quota.PlanFrom45ResponsesAt, now) {
return plan
}
}
if !IsGrok45ResponsesQuotaModel(quota.Model) || !IsQuotaSnapshotFresh(quota, now) {
return ""
}
if quotaLooksLikeGrokHeavy(quota) {
return "supergrok_heavy"
}
return ""
}
// ApplyGrok45ResponsesPlanSignal records a grok-4.5 Heavy/SuperGrok hint, or
// copies the previous 4.5 hint when this observation is a different model.
func (s *QuotaSnapshot) ApplyGrok45ResponsesPlanSignal(prev *QuotaSnapshot) {
if s == nil {
return
}
observedAt := firstNonEmptyQuotaTime(s.LastHeadersSeenAt, s.UpdatedAt)
if IsGrok45ResponsesQuotaModel(s.Model) && quotaHasLimitWindow(s) {
if quotaLooksLikeGrokHeavy(s) {
s.PlanFrom45Responses = "supergrok_heavy"
s.PlanFrom45ResponsesAt = observedAt
return
}
s.PlanFrom45Responses = "supergrok"
s.PlanFrom45ResponsesAt = observedAt
return
}
if prev != nil && strings.TrimSpace(prev.PlanFrom45Responses) != "" {
s.PlanFrom45Responses = prev.PlanFrom45Responses
s.PlanFrom45ResponsesAt = prev.PlanFrom45ResponsesAt
}
}
// QuotaSnapshotObservedAt prefers LastHeadersSeenAt over UpdatedAt so a later
// snapshot rewrite cannot refresh a stale Heavy window.
func QuotaSnapshotObservedAt(snapshot *QuotaSnapshot) (time.Time, bool) {
if snapshot == nil {
return time.Time{}, false
}
return parseQuotaTimestamp(firstNonEmptyQuotaTime(snapshot.LastHeadersSeenAt, snapshot.UpdatedAt))
}
// IsQuotaSnapshotFresh reports whether a quota signal is recent enough to
// distinguish SuperGrok from Heavy.
func IsQuotaSnapshotFresh(snapshot *QuotaSnapshot, now time.Time) bool {
observedAt, ok := QuotaSnapshotObservedAt(snapshot)
if !ok {
return false
}
return isTimeFresh(observedAt, now)
}
func isQuotaTimestampFresh(raw string, now time.Time) bool {
parsed, ok := parseQuotaTimestamp(raw)
if !ok {
return false
}
return isTimeFresh(parsed, now)
}
func parseQuotaTimestamp(raw string) (time.Time, bool) {
raw = strings.TrimSpace(raw)
if raw == "" {
return time.Time{}, false
}
parsed, err := time.Parse(time.RFC3339, raw)
if err != nil {
return time.Time{}, false
}
return parsed, true
}
func isTimeFresh(observedAt, now time.Time) bool {
if now.IsZero() {
now = time.Now()
}
age := now.Sub(observedAt)
return age <= GrokQuotaSignalMaxAge && age >= -5*time.Minute
}
func quotaHasLimitWindow(quota *QuotaSnapshot) bool {
if quota == nil {
return false
}
if quota.Requests != nil && quota.Requests.Limit != nil {
return true
}
return quota.Tokens != nil && quota.Tokens.Limit != nil
}
func quotaLooksLikeGrokHeavy(quota *QuotaSnapshot) bool {
if quota == nil {
return false
}
var requestLimit, tokenLimit int64
if quota.Requests != nil && quota.Requests.Limit != nil {
requestLimit = *quota.Requests.Limit
}
if quota.Tokens != nil && quota.Tokens.Limit != nil {
tokenLimit = *quota.Tokens.Limit
}
return requestLimit >= grokHeavyQuotaRequestLimit || tokenLimit >= grokHeavyQuotaTokenLimit
}
func firstNonEmptyQuotaTime(values ...string) string {
for _, value := range values {
if strings.TrimSpace(value) != "" {
return strings.TrimSpace(value)
}
}
return ""
}
@@ -0,0 +1,144 @@
//go:build unit
package xai
import (
"encoding/base64"
"encoding/json"
"testing"
"time"
"github.com/stretchr/testify/require"
)
func TestMapJWTSubscriptionTierNumber(t *testing.T) {
t.Parallel()
require.Equal(t, "free", MapJWTSubscriptionTier(0))
require.Equal(t, "supergrok", MapJWTSubscriptionTier(1))
require.Equal(t, "x_basic", MapJWTSubscriptionTier(2))
require.Equal(t, "x_premium", MapJWTSubscriptionTier(3))
require.Equal(t, "x_premium_plus", MapJWTSubscriptionTier(4))
require.Equal(t, "supergrok_heavy", MapJWTSubscriptionTier(5))
require.Equal(t, "supergrok_lite", MapJWTSubscriptionTier(6))
require.Equal(t, "supergrok_plus", MapJWTSubscriptionTier(7))
require.Equal(t, "9", MapJWTSubscriptionTier(9))
}
func TestNormalizeSubscriptionTierAliases(t *testing.T) {
t.Parallel()
require.Equal(t, "free", NormalizeSubscriptionTier("Free"))
require.Equal(t, "free", NormalizeSubscriptionTier(" FREE "))
require.Equal(t, "supergrok", NormalizeSubscriptionTier("SuperGrok"))
require.Equal(t, "supergrok_heavy", NormalizeSubscriptionTier("SuperGrok Heavy"))
require.Equal(t, "supergrok_pro", NormalizeSubscriptionTier("SuperGrokPro"))
require.Equal(t, "supergrok_lite", NormalizeSubscriptionTier("SuperGrok Lite"))
require.Equal(t, "supergrok_lite", NormalizeSubscriptionTier("SuperGrokLite"))
require.Equal(t, "x_basic", NormalizeSubscriptionTier("X Basic"))
require.Equal(t, "free", NormalizeSubscriptionTier("free-tier"))
require.Equal(t, "free", NormalizeSubscriptionTier("free_tier"))
require.Equal(t, "free", NormalizeSubscriptionTier("grok-basic"))
require.Equal(t, "free", NormalizeSubscriptionTier("grok_basic"))
require.Equal(t, "supergrok_lite", NormalizeSubscriptionTier("supergrok_lite"))
}
func TestSubscriptionTierFromJWTUsesNumericClaim(t *testing.T) {
t.Parallel()
require.Equal(t, "supergrok_heavy", SubscriptionTierFromJWT(jwtWithClaims(t, map[string]any{"tier": 5})))
require.Equal(t, "free", SubscriptionTierFromJWT(jwtWithClaims(t, map[string]any{"tier": 0})))
require.Equal(t, "supergrok_lite", SubscriptionTierFromJWT(jwtWithClaims(t, map[string]any{"tier": 6})))
require.Empty(t, SubscriptionTierFromJWT(jwtWithClaims(t, map[string]any{"sub": "user"})))
require.Empty(t, SubscriptionTierFromJWT("not-a-jwt"))
}
func TestCanonicalGrokPlanUsesOnlyGrok45ResponsesWindow(t *testing.T) {
t.Parallel()
zero := float64(0)
heavyReq, heavyTok := int64(8300), int64(53_000_000)
superReq, superTok := int64(900), int64(15_000_000)
fresh := time.Now().UTC().Format(time.RFC3339)
stale := time.Now().Add(-GrokQuotaSignalMaxAge - time.Hour).UTC().Format(time.RFC3339)
require.Equal(t, "supergrok", CanonicalGrokPlan(&zero, "SuperGrokPro", nil))
require.Equal(t, "supergrok", CanonicalGrokPlan(&zero, "SuperGrok", nil))
require.Equal(t, "supergrok_heavy", CanonicalGrokPlan(&zero, "SuperGrok Heavy", nil))
require.Empty(t, CanonicalGrokPlan(&zero, "", nil))
require.Equal(t, "free", CanonicalGrokPlan(&zero, "free", nil))
require.Equal(t, "supergrok_heavy", CanonicalGrokPlan(&zero, "SuperGrokPro", &QuotaSnapshot{
Model: "grok-4.5",
Requests: &QuotaWindow{Limit: &heavyReq},
Tokens: &QuotaWindow{Limit: &heavyTok},
LastHeadersSeenAt: fresh,
}))
require.Equal(t, "supergrok", CanonicalGrokPlan(&zero, "SuperGrokPro", &QuotaSnapshot{
Model: "grok-4.6",
Requests: &QuotaWindow{Limit: &heavyReq},
Tokens: &QuotaWindow{Limit: &heavyTok},
LastHeadersSeenAt: fresh,
}))
require.Equal(t, "supergrok", CanonicalGrokPlan(&zero, "SuperGrokPro", &QuotaSnapshot{
Requests: &QuotaWindow{Limit: &heavyReq},
Tokens: &QuotaWindow{Limit: &heavyTok},
LastHeadersSeenAt: fresh,
}))
require.Equal(t, "supergrok", CanonicalGrokPlan(&zero, "SuperGrokPro", &QuotaSnapshot{
Model: "grok-4.5",
Requests: &QuotaWindow{Limit: &superReq},
Tokens: &QuotaWindow{Limit: &superTok},
LastHeadersSeenAt: fresh,
}))
require.Equal(t, "supergrok", CanonicalGrokPlan(&zero, "SuperGrokPro", &QuotaSnapshot{
Model: "grok-4.5",
Requests: &QuotaWindow{Limit: &heavyReq},
LastHeadersSeenAt: stale,
}))
require.Equal(t, "supergrok_heavy", CanonicalGrokPlan(&zero, "SuperGrokPro", &QuotaSnapshot{
Model: "grok-4.6",
Requests: &QuotaWindow{Limit: &superReq},
PlanFrom45Responses: "supergrok_heavy",
PlanFrom45ResponsesAt: fresh,
}))
require.Equal(t, "free", CanonicalGrokPlan(&zero, "free", &QuotaSnapshot{
Model: "grok-4.5",
Requests: &QuotaWindow{Limit: &heavyReq},
LastHeadersSeenAt: fresh,
}))
heavyCents := float64(SuperGrokHeavyLimitCents)
require.Equal(t, "supergrok_heavy", CanonicalGrokPlan(&heavyCents, "SuperGrokPro", nil))
}
func TestApplyGrok45ResponsesPlanSignalCarriesHint(t *testing.T) {
t.Parallel()
heavyReq := int64(8300)
fresh := time.Now().UTC().Format(time.RFC3339)
prev := &QuotaSnapshot{
Model: "grok-4.5",
Requests: &QuotaWindow{Limit: &heavyReq},
LastHeadersSeenAt: fresh,
}
prev.ApplyGrok45ResponsesPlanSignal(nil)
require.Equal(t, "supergrok_heavy", prev.PlanFrom45Responses)
next := &QuotaSnapshot{
Model: "grok-4.6",
Requests: &QuotaWindow{Limit: int64Ptr(100)},
LastHeadersSeenAt: fresh,
}
next.ApplyGrok45ResponsesPlanSignal(prev)
require.Equal(t, "supergrok_heavy", next.PlanFrom45Responses)
require.Equal(t, prev.PlanFrom45ResponsesAt, next.PlanFrom45ResponsesAt)
}
func int64Ptr(v int64) *int64 { return &v }
func jwtWithClaims(t *testing.T, claims map[string]any) string {
t.Helper()
payload, err := json.Marshal(claims)
require.NoError(t, err)
return "header." + base64.RawURLEncoding.EncodeToString(payload) + ".sig"
}