将账号级请求头覆写从 anthropic/openai 的 api_key 账号扩展到 Grok 的 api_key 与 oauth 账号,并放开 Grok OAuth 账号的自定义上游地址(仅作用于 转发端点,授权与 token 刷新链路不变)。 后端: - IsHeaderOverrideEligible 扩展到 Grok(api_key+oauth),禁止名单新增 x-grok-conv-id(逐请求会话路由头)。 - 所有 Grok 上游请求路径接线 ApplyHeaderOverrides(Responses/Chat 桥/ 媒体/配额探测/billing 探测/连通性测试),统一置于内置默认头之后。 - GetGrokBaseURL/GetGrokMediaBaseURL 放开 OAuth 自定义地址:官方地址 视同未定制回落官方网关,仅显式第三方 host 改发转发流量。 - 非官方 host 允许任意 path 前缀,官方 host 仍强制 /v1。 - OAuth base_url 校验按 host 判定官方/自定义,自定义 host 恒受运营方 URL 策略约束,不受 XAI_ALLOW_UNSAFE_URL_OVERRIDES 调试开关放宽。 - 给 GrokQuotaService 注入 config,使配额/billing 探测与转发共用同一 URL 策略。 前端: - Edit 模态为 Grok OAuth 账号新增「自定义上游地址」开关。 - 请求头表单新增 JSON 快速导入与按 JSON 一键复制(复用 useClipboard, 兼容非安全上下文 HTTP 页面)。 - 门控改为平台×类型判定,Bulk 批量 base_url 增加格式校验,zh/en 文案同步。
142 lines
4.9 KiB
Go
142 lines
4.9 KiB
Go
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}]
|
|
}
|
|
}`)
|
|
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.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)
|
|
|
|
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)
|
|
}
|
|
|
|
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 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
|
|
}
|