Merge remote-tracking branch 'origin/main' into feat/grok-sso-device-oauth

# Conflicts:
#	frontend/src/api/admin/grok.ts
This commit is contained in:
shaw
2026-07-14 10:19:16 +08:00
147 changed files with 8735 additions and 597 deletions
+9 -9
View File
@@ -17,6 +17,7 @@ import (
"github.com/Wei-Shaw/sub2api/internal/pkg/proxyurl"
"github.com/Wei-Shaw/sub2api/internal/pkg/proxyutil"
"github.com/Wei-Shaw/sub2api/internal/pkg/servertiming"
)
// ForbiddenError 表示上游返回 403 Forbidden
@@ -279,7 +280,6 @@ func NewClient(proxyURL string) (*Client, error) {
}
client.Transport = transport
}
return &Client{
httpClient: client,
}, nil
@@ -341,7 +341,7 @@ func (c *Client) ExchangeCode(ctx context.Context, code, codeVerifier string) (*
}
req.Header.Set("Content-Type", "application/x-www-form-urlencoded")
resp, err := c.httpClient.Do(req)
resp, err := servertiming.Do(c.httpClient, req)
if err != nil {
return nil, fmt.Errorf("token 交换请求失败: %w", err)
}
@@ -383,7 +383,7 @@ func (c *Client) RefreshToken(ctx context.Context, refreshToken string) (*TokenR
}
req.Header.Set("Content-Type", "application/x-www-form-urlencoded")
resp, err := c.httpClient.Do(req)
resp, err := servertiming.Do(c.httpClient, req)
if err != nil {
return nil, fmt.Errorf("token 刷新请求失败: %w", err)
}
@@ -414,7 +414,7 @@ func (c *Client) GetUserInfo(ctx context.Context, accessToken string) (*UserInfo
}
req.Header.Set("Authorization", "Bearer "+accessToken)
resp, err := c.httpClient.Do(req)
resp, err := servertiming.Do(c.httpClient, req)
if err != nil {
return nil, fmt.Errorf("用户信息请求失败: %w", err)
}
@@ -465,7 +465,7 @@ func (c *Client) LoadCodeAssist(ctx context.Context, accessToken string) (*LoadC
req.Header.Set("Content-Type", "application/json")
req.Header.Set("User-Agent", GetUserAgentForContext(ctx))
resp, err := c.httpClient.Do(req)
resp, err := servertiming.Do(c.httpClient, req)
if err != nil {
lastErr = fmt.Errorf("loadCodeAssist 请求失败: %w", err)
if shouldFallbackToNextURL(err, 0) && urlIdx < len(availableURLs)-1 {
@@ -544,7 +544,7 @@ func (c *Client) OnboardUser(ctx context.Context, accessToken, tierID string) (s
req.Header.Set("Content-Type", "application/json")
req.Header.Set("User-Agent", GetUserAgentForContext(ctx))
resp, err := c.httpClient.Do(req)
resp, err := servertiming.Do(c.httpClient, req)
if err != nil {
lastErr = fmt.Errorf("onboardUser 请求失败: %w", err)
if shouldFallbackToNextURL(err, 0) && urlIdx < len(availableURLs)-1 {
@@ -683,7 +683,7 @@ func (c *Client) FetchAvailableModels(ctx context.Context, accessToken, projectI
req.Header.Set("Content-Type", "application/json")
req.Header.Set("User-Agent", GetUserAgentForContext(ctx))
resp, err := fetchClient.Do(req)
resp, err := servertiming.Do(fetchClient, req)
if err != nil {
lastErr = fmt.Errorf("fetchAvailableModels 请求失败: %w", err)
if shouldFallbackToNextURL(err, 0) && urlIdx < len(availableURLs)-1 {
@@ -842,7 +842,7 @@ func (c *Client) SetUserSettings(ctx context.Context, accessToken string) (*SetU
req.Header.Set("X-Goog-Api-Client", "gl-node/22.21.1")
req.Host = "daily-cloudcode-pa.googleapis.com"
resp, err := c.httpClient.Do(req)
resp, err := servertiming.Do(c.httpClient, req)
if err != nil {
return nil, fmt.Errorf("setUserSettings 请求失败: %w", err)
}
@@ -885,7 +885,7 @@ func (c *Client) FetchUserInfo(ctx context.Context, accessToken, projectID strin
req.Header.Set("X-Goog-Api-Client", "gl-node/22.21.1")
req.Host = "daily-cloudcode-pa.googleapis.com"
resp, err := c.httpClient.Do(req)
resp, err := servertiming.Do(c.httpClient, req)
if err != nil {
return nil, fmt.Errorf("fetchUserInfo 请求失败: %w", err)
}
+2
View File
@@ -25,6 +25,7 @@ import (
"github.com/Wei-Shaw/sub2api/internal/pkg/proxyurl"
"github.com/Wei-Shaw/sub2api/internal/pkg/proxyutil"
"github.com/Wei-Shaw/sub2api/internal/pkg/servertiming"
"github.com/Wei-Shaw/sub2api/internal/util/urlvalidator"
)
@@ -92,6 +93,7 @@ func buildClient(opts Options) (*http.Client, error) {
if opts.ValidateResolvedIP && !opts.AllowPrivateHosts {
rt = newValidatedTransport(transport)
}
rt = servertiming.WrapRoundTripper(rt)
return &http.Client{
Transport: rt,
Timeout: opts.Timeout,
@@ -0,0 +1,348 @@
package servertiming
import (
"context"
"fmt"
"sort"
"strconv"
"strings"
"sync"
"time"
)
const (
HeaderName = "Server-Timing"
AdminUIHeader = "X-Admin-UI-Request"
MetricDatabase = "db"
MetricRedis = "redis"
dependencyPrefix = "dep_"
maxMetricNameLength = 48
maxIntervals = 2048
maxHeaderLength = 4096
)
type contextKey struct{}
type interval struct {
start time.Time
end time.Time
}
type metric struct {
count int64
intervals []interval
}
// Collector stores request-scoped timing samples. It is safe for concurrent use.
type Collector struct {
startedAt time.Time
mu sync.Mutex
metrics map[string]*metric
cacheStatus string
}
// New creates a collector whose total duration starts at startedAt.
func New(startedAt time.Time) *Collector {
if startedAt.IsZero() {
startedAt = time.Now()
}
return &Collector{
startedAt: startedAt,
metrics: make(map[string]*metric),
}
}
// WithCollector attaches a collector to a context.
func WithCollector(ctx context.Context, collector *Collector) context.Context {
if ctx == nil {
ctx = context.Background()
}
if collector == nil {
return ctx
}
return context.WithValue(ctx, contextKey{}, collector)
}
// FromContext returns the request timing collector, when one is active.
func FromContext(ctx context.Context) (*Collector, bool) {
if ctx == nil {
return nil, false
}
collector, ok := ctx.Value(contextKey{}).(*Collector)
return collector, ok && collector != nil
}
// Active reports whether timing collection is enabled for this request.
func Active(ctx context.Context) bool {
_, ok := FromContext(ctx)
return ok
}
// Record adds a completed interval and operation count to a metric.
func Record(ctx context.Context, name string, startedAt, endedAt time.Time, count int) {
collector, ok := FromContext(ctx)
if !ok {
return
}
collector.Record(name, startedAt, endedAt, count)
}
// RecordInterval adds timing without incrementing the operation count. It is
// useful when one logical operation has multiple blocking driver calls.
func RecordInterval(ctx context.Context, name string, startedAt, endedAt time.Time) {
collector, ok := FromContext(ctx)
if !ok {
return
}
collector.record(name, startedAt, endedAt, 0)
}
// Record adds a completed interval directly to the collector.
func (c *Collector) Record(name string, startedAt, endedAt time.Time, count int) {
if count <= 0 {
count = 1
}
c.record(name, startedAt, endedAt, count)
}
func (c *Collector) record(name string, startedAt, endedAt time.Time, count int) {
name = normalizeMetricName(name)
if c == nil || name == "" || startedAt.IsZero() || endedAt.Before(startedAt) {
return
}
if count < 0 {
count = 0
}
c.mu.Lock()
m := c.metrics[name]
if m == nil {
m = &metric{}
c.metrics[name] = m
}
m.count += int64(count)
if len(m.intervals) < maxIntervals {
m.intervals = append(m.intervals, interval{start: startedAt, end: endedAt})
}
c.mu.Unlock()
}
// Observe starts a metric span and returns an idempotent completion function.
func Observe(ctx context.Context, name string) func() {
collector, ok := FromContext(ctx)
name = normalizeMetricName(name)
if !ok || name == "" {
return func() {}
}
startedAt := time.Now()
var once sync.Once
return func() {
once.Do(func() {
collector.Record(name, startedAt, time.Now(), 1)
})
}
}
// ObserveDependency starts a named external dependency span.
func ObserveDependency(ctx context.Context, module string) func() {
return Observe(ctx, dependencyMetricName(module))
}
// RecordDependency records a completed external dependency interval.
func RecordDependency(ctx context.Context, module string, startedAt, endedAt time.Time) {
Record(ctx, dependencyMetricName(module), startedAt, endedAt, 1)
}
// SetCacheStatus records the response-cache outcome for the request.
func SetCacheStatus(ctx context.Context, status string) {
collector, ok := FromContext(ctx)
if !ok {
return
}
status = normalizeCacheStatus(status)
if status == "" {
return
}
collector.mu.Lock()
collector.cacheStatus = status
collector.mu.Unlock()
}
// HeaderValue renders a bounded, deterministic Server-Timing header.
func HeaderValue(ctx context.Context, endedAt time.Time, cacheStatus string) string {
collector, ok := FromContext(ctx)
if !ok {
return ""
}
return collector.HeaderValue(endedAt, cacheStatus)
}
// HeaderValue renders a bounded, deterministic Server-Timing header.
func (c *Collector) HeaderValue(endedAt time.Time, cacheStatus string) string {
if c == nil {
return ""
}
if endedAt.IsZero() {
endedAt = time.Now()
}
if endedAt.Before(c.startedAt) {
endedAt = c.startedAt
}
c.mu.Lock()
metrics := make(map[string]metric, len(c.metrics))
allIntervals := make([]interval, 0)
dependencyIntervals := make([]interval, 0)
var dependencyCount int64
for name, source := range c.metrics {
copied := metric{count: source.count, intervals: append([]interval(nil), source.intervals...)}
metrics[name] = copied
allIntervals = append(allIntervals, copied.intervals...)
if strings.HasPrefix(name, dependencyPrefix) {
dependencyIntervals = append(dependencyIntervals, copied.intervals...)
dependencyCount += copied.count
}
}
storedCacheStatus := c.cacheStatus
c.mu.Unlock()
total := endedAt.Sub(c.startedAt)
blocked := unionDuration(allIntervals, c.startedAt, endedAt)
app := total - blocked
if app < 0 {
app = 0
}
cacheStatus = normalizeCacheStatus(cacheStatus)
if cacheStatus == "" {
cacheStatus = normalizeCacheStatus(storedCacheStatus)
}
if cacheStatus == "" {
cacheStatus = "bypass"
}
database := metrics[MetricDatabase]
redisMetric := metrics[MetricRedis]
parts := []string{
"total;dur=" + formatDuration(total),
"app;dur=" + formatDuration(app),
fmt.Sprintf("db;dur=%s;desc=\"queries=%d\"", formatDuration(unionDuration(database.intervals, c.startedAt, endedAt)), database.count),
fmt.Sprintf("redis;dur=%s;desc=\"commands=%d\"", formatDuration(unionDuration(redisMetric.intervals, c.startedAt, endedAt)), redisMetric.count),
"cache;desc=\"" + cacheStatus + "\"",
fmt.Sprintf("deps;dur=%s;desc=\"calls=%d\"", formatDuration(unionDuration(dependencyIntervals, c.startedAt, endedAt)), dependencyCount),
}
dependencyNames := make([]string, 0)
for name := range metrics {
if strings.HasPrefix(name, dependencyPrefix) {
dependencyNames = append(dependencyNames, name)
}
}
sort.Strings(dependencyNames)
for _, name := range dependencyNames {
m := metrics[name]
part := fmt.Sprintf("%s;dur=%s;desc=\"calls=%d\"", name, formatDuration(unionDuration(m.intervals, c.startedAt, endedAt)), m.count)
candidate := strings.Join(append(parts, part), ", ")
if len(candidate) > maxHeaderLength {
break
}
parts = append(parts, part)
}
return strings.Join(parts, ", ")
}
func dependencyMetricName(module string) string {
module = normalizeMetricName(module)
module = strings.TrimPrefix(module, dependencyPrefix)
if module == "" {
module = "http"
}
return dependencyPrefix + module
}
func normalizeMetricName(name string) string {
name = strings.ToLower(strings.TrimSpace(name))
if name == "" {
return ""
}
var b strings.Builder
b.Grow(min(len(name), maxMetricNameLength))
for _, r := range name {
if b.Len() >= maxMetricNameLength {
break
}
switch {
case r >= 'a' && r <= 'z', r >= '0' && r <= '9':
_, _ = b.WriteRune(r)
case r == '_' || r == '-':
_ = b.WriteByte('_')
}
}
return strings.Trim(b.String(), "_")
}
func normalizeCacheStatus(status string) string {
switch strings.ToLower(strings.TrimSpace(status)) {
case "hit":
return "hit"
case "miss":
return "miss"
case "bypass":
return "bypass"
default:
return ""
}
}
func unionDuration(intervals []interval, lowerBound, upperBound time.Time) time.Duration {
if len(intervals) == 0 || !upperBound.After(lowerBound) {
return 0
}
normalized := make([]interval, 0, len(intervals))
for _, item := range intervals {
start := item.start
end := item.end
if start.Before(lowerBound) {
start = lowerBound
}
if end.After(upperBound) {
end = upperBound
}
if end.After(start) {
normalized = append(normalized, interval{start: start, end: end})
}
}
if len(normalized) == 0 {
return 0
}
sort.Slice(normalized, func(i, j int) bool {
return normalized[i].start.Before(normalized[j].start)
})
currentStart := normalized[0].start
currentEnd := normalized[0].end
var total time.Duration
for _, item := range normalized[1:] {
if !item.start.After(currentEnd) {
if item.end.After(currentEnd) {
currentEnd = item.end
}
continue
}
total += currentEnd.Sub(currentStart)
currentStart = item.start
currentEnd = item.end
}
total += currentEnd.Sub(currentStart)
return total
}
func formatDuration(value time.Duration) string {
if value < 0 {
value = 0
}
return strconv.FormatFloat(float64(value)/float64(time.Millisecond), 'f', 1, 64)
}
@@ -0,0 +1,129 @@
package servertiming
import (
"context"
"fmt"
"strings"
"sync"
"testing"
"time"
)
func TestCollectorHeaderValueAggregatesIntervals(t *testing.T) {
startedAt := time.Unix(100, 0)
collector := New(startedAt)
collector.Record(MetricDatabase, startedAt.Add(10*time.Millisecond), startedAt.Add(40*time.Millisecond), 2)
collector.Record(MetricRedis, startedAt.Add(30*time.Millisecond), startedAt.Add(50*time.Millisecond), 3)
collector.Record(dependencyMetricName("openai"), startedAt.Add(70*time.Millisecond), startedAt.Add(100*time.Millisecond), 1)
collector.Record(dependencyMetricName("github"), startedAt.Add(60*time.Millisecond), startedAt.Add(90*time.Millisecond), 1)
got := collector.HeaderValue(startedAt.Add(120*time.Millisecond), "miss")
want := `total;dur=120.0, app;dur=40.0, db;dur=30.0;desc="queries=2", redis;dur=20.0;desc="commands=3", cache;desc="miss", deps;dur=40.0;desc="calls=2", dep_github;dur=30.0;desc="calls=1", dep_openai;dur=30.0;desc="calls=1"`
if got != want {
t.Fatalf("HeaderValue() = %q, want %q", got, want)
}
}
func TestRecordIntervalDoesNotIncrementCount(t *testing.T) {
startedAt := time.Unix(200, 0)
collector := New(startedAt)
ctx := WithCollector(context.Background(), collector)
Record(ctx, MetricDatabase, startedAt.Add(10*time.Millisecond), startedAt.Add(20*time.Millisecond), 1)
RecordInterval(ctx, MetricDatabase, startedAt.Add(30*time.Millisecond), startedAt.Add(40*time.Millisecond))
header := HeaderValue(ctx, startedAt.Add(100*time.Millisecond), "hit")
if !strings.Contains(header, `db;dur=20.0;desc="queries=1"`) {
t.Fatalf("header %q does not contain one query with both blocking intervals", header)
}
if !strings.Contains(header, "app;dur=80.0") {
t.Fatalf("header %q does not subtract the interval union from app time", header)
}
}
func TestCollectorCacheStatusFallback(t *testing.T) {
startedAt := time.Unix(300, 0)
collector := New(startedAt)
ctx := WithCollector(context.Background(), collector)
SetCacheStatus(ctx, " HIT ")
if got := HeaderValue(ctx, startedAt.Add(time.Millisecond), "invalid"); !strings.Contains(got, `cache;desc="hit"`) {
t.Fatalf("HeaderValue() = %q, want stored cache hit", got)
}
other := New(startedAt)
if got := other.HeaderValue(startedAt.Add(time.Millisecond), "invalid"); !strings.Contains(got, `cache;desc="bypass"`) {
t.Fatalf("HeaderValue() = %q, want cache bypass", got)
}
}
func TestCollectorSanitizesDependencyMetric(t *testing.T) {
startedAt := time.Unix(400, 0)
collector := New(startedAt)
ctx := WithCollector(context.Background(), collector)
RecordDependency(ctx, "GitHub API\r\nInjected;dur=999", startedAt, startedAt.Add(time.Millisecond))
header := HeaderValue(ctx, startedAt.Add(2*time.Millisecond), "bypass")
if strings.ContainsAny(header, "\r\n") || strings.Contains(header, ";dur=999") {
t.Fatalf("unsafe metric content reached header: %q", header)
}
if !strings.Contains(header, "dep_githubapiinjecteddur999;dur=1.0") {
t.Fatalf("sanitized dependency metric missing from header: %q", header)
}
}
func TestCollectorBoundsHeaderLength(t *testing.T) {
startedAt := time.Unix(500, 0)
collector := New(startedAt)
for i := 0; i < 300; i++ {
collector.Record(
dependencyMetricName(fmt.Sprintf("module_%03d_with_a_deliberately_long_name", i)),
startedAt,
startedAt.Add(time.Millisecond),
1,
)
}
header := collector.HeaderValue(startedAt.Add(2*time.Millisecond), "bypass")
if len(header) > maxHeaderLength {
t.Fatalf("header length = %d, want <= %d", len(header), maxHeaderLength)
}
if !strings.Contains(header, "total;dur=2.0") || !strings.Contains(header, "deps;dur=1.0") {
t.Fatalf("bounded header lost fixed metrics: %q", header)
}
}
func TestCollectorConcurrentRecording(t *testing.T) {
startedAt := time.Now()
collector := New(startedAt)
ctx := WithCollector(context.Background(), collector)
const workers = 25
const recordsPerWorker = 100
var wg sync.WaitGroup
wg.Add(workers)
for i := 0; i < workers; i++ {
go func() {
defer wg.Done()
for j := 0; j < recordsPerWorker; j++ {
Record(ctx, MetricDatabase, startedAt, startedAt.Add(time.Microsecond), 1)
}
}()
}
wg.Wait()
header := HeaderValue(ctx, startedAt.Add(time.Millisecond), "bypass")
want := fmt.Sprintf(`queries=%d`, workers*recordsPerWorker)
if !strings.Contains(header, want) {
t.Fatalf("header %q does not contain %q", header, want)
}
}
func TestContextHelpersHandleMissingCollector(t *testing.T) {
if Active(context.Background()) {
t.Fatal("context without collector reported active")
}
if got := HeaderValue(context.Background(), time.Now(), "hit"); got != "" {
t.Fatalf("HeaderValue() = %q without collector, want empty", got)
}
}
+104
View File
@@ -0,0 +1,104 @@
package servertiming
import (
"context"
"net/http"
"strings"
"time"
)
type dependencyModuleKey struct{}
type timingRoundTripper struct {
base http.RoundTripper
}
// WithDependencyModule overrides the safe module name used for an outbound call.
func WithDependencyModule(ctx context.Context, module string) context.Context {
if ctx == nil {
ctx = context.Background()
}
module = strings.TrimPrefix(normalizeMetricName(module), dependencyPrefix)
if module == "" {
return ctx
}
return context.WithValue(ctx, dependencyModuleKey{}, module)
}
// WrapRoundTripper records outbound response-header latency for active requests.
func WrapRoundTripper(base http.RoundTripper) http.RoundTripper {
if base == nil {
base = http.DefaultTransport
}
if _, ok := base.(*timingRoundTripper); ok {
return base
}
return &timingRoundTripper{base: base}
}
// InstrumentClient returns a shallow client copy with an instrumented transport.
func InstrumentClient(client *http.Client) *http.Client {
if client == nil {
client = &http.Client{}
}
copyClient := *client
copyClient.Transport = WrapRoundTripper(copyClient.Transport)
return &copyClient
}
// Do records response-header latency without changing the client's transport
// type. Use it for clients whose callers inspect or configure *http.Transport.
func Do(client *http.Client, req *http.Request) (*http.Response, error) {
if client == nil {
client = http.DefaultClient
}
if req == nil || !Active(req.Context()) {
return client.Do(req)
}
startedAt := time.Now()
response, err := client.Do(req)
RecordDependency(req.Context(), dependencyModule(req), startedAt, time.Now())
return response, err
}
func (t *timingRoundTripper) RoundTrip(req *http.Request) (*http.Response, error) {
if req == nil || !Active(req.Context()) {
return t.base.RoundTrip(req)
}
startedAt := time.Now()
response, err := t.base.RoundTrip(req)
RecordDependency(req.Context(), dependencyModule(req), startedAt, time.Now())
return response, err
}
func dependencyModule(req *http.Request) string {
if req != nil {
if module, ok := req.Context().Value(dependencyModuleKey{}).(string); ok && module != "" {
return module
}
}
if req == nil || req.URL == nil {
return "http"
}
host := strings.ToLower(req.URL.Hostname())
switch {
case strings.Contains(host, "github"):
return "github"
case strings.Contains(host, "openai"):
return "openai"
case strings.Contains(host, "anthropic"):
return "anthropic"
case strings.Contains(host, "generativelanguage") || strings.Contains(host, "gemini"):
return "gemini"
case strings.Contains(host, "cloudcode") || strings.Contains(host, "antigravity"):
return "antigravity"
case strings.Contains(host, "googleapis") || strings.Contains(host, "google"):
return "google"
case strings.Contains(host, "amazonaws") || strings.Contains(host, "cloudflarestorage") || strings.Contains(host, "s3"):
return "s3"
case strings.Contains(host, "stripe") || strings.Contains(host, "airwallex") || strings.Contains(host, "alipay") || strings.Contains(host, "wechatpay") || strings.Contains(host, "paypal"):
return "payment"
default:
return "http"
}
}
@@ -0,0 +1,168 @@
package servertiming
import (
"context"
"io"
"net/http"
"strings"
"testing"
"time"
)
type roundTripFunc func(*http.Request) (*http.Response, error)
func (f roundTripFunc) RoundTrip(req *http.Request) (*http.Response, error) {
return f(req)
}
type trackingBody struct {
read bool
}
func (b *trackingBody) Read(_ []byte) (int, error) {
b.read = true
return 0, io.EOF
}
func (b *trackingBody) Close() error { return nil }
func TestWrapRoundTripperRecordsResponseHeaderLatency(t *testing.T) {
startedAt := time.Now()
collector := New(startedAt)
body := &trackingBody{}
baseCalled := false
base := roundTripFunc(func(req *http.Request) (*http.Response, error) {
baseCalled = true
return &http.Response{
StatusCode: http.StatusOK,
Body: body,
Header: make(http.Header),
Request: req,
}, nil
})
req, err := http.NewRequestWithContext(WithCollector(context.Background(), collector), http.MethodGet, "https://api.github.com/repos/example/project", nil)
if err != nil {
t.Fatal(err)
}
resp, err := WrapRoundTripper(base).RoundTrip(req)
if err != nil {
t.Fatal(err)
}
defer func() { _ = resp.Body.Close() }()
if !baseCalled {
t.Fatal("base RoundTripper was not called")
}
if body.read {
t.Fatal("RoundTripper instrumentation read the response body; timing must stop at response headers")
}
header := collector.HeaderValue(time.Now(), "bypass")
if !strings.Contains(header, `dep_github;dur=`) || !strings.Contains(header, `deps;dur=`) {
t.Fatalf("dependency metrics missing from header: %q", header)
}
}
func TestWrapRoundTripperUsesContextModuleOverride(t *testing.T) {
collector := New(time.Now())
ctx := WithDependencyModule(WithCollector(context.Background(), collector), "data-managementd")
req, err := http.NewRequestWithContext(ctx, http.MethodGet, "https://private.example.test/path", nil)
if err != nil {
t.Fatal(err)
}
base := roundTripFunc(func(req *http.Request) (*http.Response, error) {
return &http.Response{StatusCode: http.StatusOK, Body: http.NoBody, Header: make(http.Header), Request: req}, nil
})
if _, err := WrapRoundTripper(base).RoundTrip(req); err != nil {
t.Fatal(err)
}
header := collector.HeaderValue(time.Now(), "bypass")
if !strings.Contains(header, "dep_data_managementd") {
t.Fatalf("module override missing from header: %q", header)
}
if strings.Contains(header, "private.example") {
t.Fatalf("raw host leaked into header: %q", header)
}
}
func TestWrapRoundTripperSkipsInactiveContext(t *testing.T) {
called := false
base := roundTripFunc(func(req *http.Request) (*http.Response, error) {
called = true
return &http.Response{StatusCode: http.StatusNoContent, Body: http.NoBody, Header: make(http.Header), Request: req}, nil
})
req, err := http.NewRequest(http.MethodGet, "https://api.openai.com/v1/models", nil)
if err != nil {
t.Fatal(err)
}
if _, err := WrapRoundTripper(base).RoundTrip(req); err != nil {
t.Fatal(err)
}
if !called {
t.Fatal("inactive request did not reach base RoundTripper")
}
}
func TestDoRecordsWithoutChangingTransportType(t *testing.T) {
base := roundTripFunc(func(req *http.Request) (*http.Response, error) {
return &http.Response{StatusCode: http.StatusOK, Body: http.NoBody, Header: make(http.Header), Request: req}, nil
})
client := &http.Client{Transport: base}
collector := New(time.Now())
req, err := http.NewRequestWithContext(WithCollector(context.Background(), collector), http.MethodGet, "https://api.openai.com/v1/models", nil)
if err != nil {
t.Fatal(err)
}
if _, err := Do(client, req); err != nil {
t.Fatal(err)
}
if _, ok := client.Transport.(roundTripFunc); !ok {
t.Fatalf("Do changed client transport type to %T", client.Transport)
}
if header := collector.HeaderValue(time.Now(), "bypass"); !strings.Contains(header, "dep_openai;dur=") {
t.Fatalf("dependency metric missing from header: %q", header)
}
}
func TestDependencyModuleClassification(t *testing.T) {
tests := map[string]string{
"https://api.github.com/repos/a/b": "github",
"https://api.openai.com/v1/models": "openai",
"https://api.anthropic.com/v1/messages": "anthropic",
"https://generativelanguage.googleapis.com/v1/models": "gemini",
"https://cloudcode-pa.googleapis.com/v1internal": "antigravity",
"https://storage.googleapis.com/bucket/object": "google",
"https://bucket.s3.amazonaws.com/object": "s3",
"https://api.stripe.com/v1/refunds": "payment",
"https://dependency.example.test/path": "http",
}
for rawURL, want := range tests {
req, err := http.NewRequest(http.MethodGet, rawURL, nil)
if err != nil {
t.Fatalf("NewRequest(%q): %v", rawURL, err)
}
if got := dependencyModule(req); got != want {
t.Errorf("dependencyModule(%q) = %q, want %q", rawURL, got, want)
}
}
}
func TestClientInstrumentationDoesNotMutateOriginal(t *testing.T) {
base := roundTripFunc(func(req *http.Request) (*http.Response, error) {
return &http.Response{StatusCode: http.StatusOK, Body: http.NoBody, Header: make(http.Header), Request: req}, nil
})
original := &http.Client{Transport: base, Timeout: time.Second}
instrumented := InstrumentClient(original)
if instrumented == original {
t.Fatal("InstrumentClient returned the original client")
}
if _, ok := original.Transport.(roundTripFunc); !ok {
t.Fatalf("InstrumentClient mutated the original transport to %T", original.Transport)
}
if instrumented.Timeout != original.Timeout {
t.Fatal("InstrumentClient did not preserve client settings")
}
if WrapRoundTripper(instrumented.Transport) != instrumented.Transport {
t.Fatal("WrapRoundTripper wrapped an already instrumented transport twice")
}
}
+372
View File
@@ -0,0 +1,372 @@
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"
// Keep in sync with https://x.ai/cli/stable.
CLIClientVersion = "0.2.93"
CLIUserAgent = "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.
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"`
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.
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"`
Plan string `json:"plan,omitempty"` // SuperGrok | SuperGrok Heavy | ""
StatusCode int `json:"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
}
// 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", CLIUserAgent)
}
// 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)
periodStart := ""
periodEnd := ""
if period != nil {
periodStart = strings.TrimSpace(period.Start)
periodEnd = strings.TrimSpace(period.End)
}
if periodStart == "" {
periodStart = strings.TrimSpace(config.BillingPeriodStart)
}
if periodEnd == "" {
periodEnd = strings.TrimSpace(config.BillingPeriodEnd)
}
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)
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
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.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
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.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
}
+127
View File
@@ -0,0 +1,127 @@
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 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
}