Files
sub2api/backend/internal/pkg/xai/billing_test.go
T
shaw 221581400b feat(grok): 支持账号级自定义上游地址与请求头覆写
将账号级请求头覆写从 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 文案同步。
2026-07-15 20:30:23 +08:00

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
}