feat(grok): stream idle 换号、team+model 冷却与 RT 刷新打散

- Grok 流默认启用上游读空闲超时,首包前返回可 failover 错误并短冷却
- 同 team_id 账号在 429/free-usage 后对同模型共享进程内冷却,粘连与选号均过滤
- TokenRefreshService 已覆盖 Grok;NeedsRefresh 增加按账号稳定 jitter 防 stampede
This commit is contained in:
IanShaw027
2026-08-07 13:56:38 +08:00
parent dfca6246eb
commit ec9e733606
9 changed files with 427 additions and 3 deletions
@@ -0,0 +1,38 @@
package service
import (
"fmt"
"strings"
"time"
)
// Default Grok stream idle when gateway.stream_data_interval_timeout is 0.
// Long enough for slow thinking models, short enough to release hung sockets.
const defaultGrokStreamIdleTimeout = 180 * time.Second
// Shorter cool after a Grok stream-idle failure so the account can re-enter soon
// but is not immediately re-picked in a tight failover loop.
const grokStreamIdleCooldown = 2 * time.Minute
// resolveGrokStreamIdleTimeout returns the effective upstream-read idle timeout
// for Grok streams. Prefers the global gateway setting when positive; otherwise
// applies a Grok-only default so hung SSE bodies still fail over.
func resolveGrokStreamIdleTimeout(cfgStreamIntervalSec int) time.Duration {
if cfgStreamIntervalSec > 0 {
return time.Duration(cfgStreamIntervalSec) * time.Second
}
return defaultGrokStreamIdleTimeout
}
// grokStreamIdleFailoverError builds a pre-commit/handler-visible failover so
// the gateway can switch OAuth accounts after a hung Grok upstream stream.
func grokStreamIdleFailoverError(account *Account, idle time.Duration) *UpstreamFailoverError {
msg := fmt.Sprintf("Grok stream idle timeout after %s with no upstream data", idle.Round(time.Second))
return &UpstreamFailoverError{
StatusCode: 502,
ResponseBody: []byte(`{"error":{"code":"empty_upstream","message":"` + strings.ReplaceAll(msg, `"`, `'`) + `"}}`),
SafeToFailoverAfterWrite: true,
// Allow pool-mode retries; normal OAuth switches account via handler.
RetryableOnSameAccount: account != nil && account.IsPoolMode(),
}
}
@@ -0,0 +1,25 @@
//go:build unit
package service
import (
"testing"
"time"
"github.com/stretchr/testify/require"
)
func TestResolveGrokStreamIdleTimeout(t *testing.T) {
require.Equal(t, 90*time.Second, resolveGrokStreamIdleTimeout(90))
require.Equal(t, defaultGrokStreamIdleTimeout, resolveGrokStreamIdleTimeout(0))
require.Equal(t, defaultGrokStreamIdleTimeout, resolveGrokStreamIdleTimeout(-1))
}
func TestGrokStreamIdleFailoverError(t *testing.T) {
account := &Account{ID: 1, Platform: PlatformGrok, Type: AccountTypeOAuth}
err := grokStreamIdleFailoverError(account, 180*time.Second)
require.NotNil(t, err)
require.Equal(t, 502, err.StatusCode)
require.True(t, err.SafeToFailoverAfterWrite)
require.Contains(t, string(err.ResponseBody), "empty_upstream")
}
@@ -0,0 +1,148 @@
package service
import (
"crypto/sha256"
"encoding/hex"
"strings"
"sync"
"time"
)
// In-memory team+model rate-limit overlay for Grok OAuth. When xAI rate-limits
// one account in a team for a model, sibling accounts sharing team_id skip the
// same model until the cooldown expires (mirrors grok2api teamModelRateLimit).
//
// Process-local only: multi-instance deployments each learn the block from their
// own 429s. Prefer short TTLs so drift self-heals.
type grokTeamModelRateLimit struct {
Until time.Time
}
type grokTeamModelRateLimitStore struct {
mu sync.Mutex
items map[string]grokTeamModelRateLimit
}
var globalGrokTeamModelRateLimits = &grokTeamModelRateLimitStore{
items: make(map[string]grokTeamModelRateLimit),
}
const (
grokTeamRateLimitDefaultTTL = 10 * time.Minute
grokTeamRateLimitMaxTTL = time.Hour
grokTeamRateLimitMinTTL = 30 * time.Second
)
func grokTeamFingerprint(teamID string) string {
teamID = strings.TrimSpace(teamID)
if teamID == "" {
return ""
}
sum := sha256.Sum256([]byte(strings.ToLower(teamID)))
return hex.EncodeToString(sum[:8])
}
func grokTeamModelRateLimitKey(teamFingerprint, model string) string {
return teamFingerprint + "|" + strings.ToLower(strings.TrimSpace(model))
}
func accountGrokTeamID(account *Account) string {
if account == nil {
return ""
}
return strings.TrimSpace(account.GetCredential("team_id"))
}
// markGrokTeamModelRateLimit records that this team+model pair should be skipped
// until until. No-op when team_id or model is empty.
func markGrokTeamModelRateLimit(account *Account, model string, until time.Time) {
if account == nil || !account.IsGrokOAuth() {
return
}
fp := grokTeamFingerprint(accountGrokTeamID(account))
model = strings.TrimSpace(model)
if fp == "" || model == "" || until.IsZero() {
return
}
now := time.Now()
if !until.After(now) {
until = now.Add(grokTeamRateLimitDefaultTTL)
}
maxUntil := now.Add(grokTeamRateLimitMaxTTL)
if until.After(maxUntil) {
until = maxUntil
}
key := grokTeamModelRateLimitKey(fp, model)
globalGrokTeamModelRateLimits.mu.Lock()
defer globalGrokTeamModelRateLimits.mu.Unlock()
if cur, ok := globalGrokTeamModelRateLimits.items[key]; ok && cur.Until.After(until) {
return
}
globalGrokTeamModelRateLimits.items[key] = grokTeamModelRateLimit{Until: until}
// Opportunistic prune of expired entries.
for k, v := range globalGrokTeamModelRateLimits.items {
if !v.Until.After(now) {
delete(globalGrokTeamModelRateLimits.items, k)
}
}
}
// isGrokTeamModelRateLimited reports whether the account's team is currently
// blocked for the requested model.
func isGrokTeamModelRateLimited(account *Account, model string, now time.Time) bool {
if account == nil || !account.IsGrokOAuth() {
return false
}
fp := grokTeamFingerprint(accountGrokTeamID(account))
model = strings.TrimSpace(model)
if fp == "" || model == "" {
return false
}
key := grokTeamModelRateLimitKey(fp, model)
globalGrokTeamModelRateLimits.mu.Lock()
defer globalGrokTeamModelRateLimits.mu.Unlock()
cur, ok := globalGrokTeamModelRateLimits.items[key]
if !ok {
return false
}
if !cur.Until.After(now) {
delete(globalGrokTeamModelRateLimits.items, key)
return false
}
return true
}
// filterGrokTeamModelRateLimitedAccounts drops candidates whose team is under a
// model-scoped rate-limit cool. Accounts without team_id pass through.
func filterGrokTeamModelRateLimitedAccounts(accounts []Account, model string, now time.Time) []Account {
if len(accounts) == 0 || strings.TrimSpace(model) == "" {
return accounts
}
out := accounts[:0]
kept := false
for i := range accounts {
if isGrokTeamModelRateLimited(&accounts[i], model, now) {
continue
}
out = append(out, accounts[i])
kept = true
}
if !kept && len(out) == 0 {
// All filtered — return empty (caller treats as no capacity).
return nil
}
return out
}
// resolveGrokTeamRateLimitUntil derives a team cool window from an observed
// account rate-limit reset, with sane clamps.
func resolveGrokTeamRateLimitUntil(resetAt, now time.Time) time.Time {
if resetAt.After(now.Add(grokTeamRateLimitMinTTL)) {
maxUntil := now.Add(grokTeamRateLimitMaxTTL)
if resetAt.After(maxUntil) {
return maxUntil
}
return resetAt
}
return now.Add(grokTeamRateLimitDefaultTTL)
}
@@ -0,0 +1,59 @@
//go:build unit
package service
import (
"testing"
"time"
"github.com/stretchr/testify/require"
)
func TestGrokTeamModelRateLimit_MarksAndFiltersSiblings(t *testing.T) {
// Isolate from other tests by using unique team ids.
team := "team-test-" + time.Now().Format("150405.000")
a1 := &Account{
ID: 101, Platform: PlatformGrok, Type: AccountTypeOAuth,
Credentials: map[string]any{"team_id": team},
}
a2 := &Account{
ID: 102, Platform: PlatformGrok, Type: AccountTypeOAuth,
Credentials: map[string]any{"team_id": team},
}
other := &Account{
ID: 103, Platform: PlatformGrok, Type: AccountTypeOAuth,
Credentials: map[string]any{"team_id": team + "-other"},
}
noTeam := &Account{
ID: 104, Platform: PlatformGrok, Type: AccountTypeOAuth,
Credentials: map[string]any{},
}
now := time.Now()
markGrokTeamModelRateLimit(a1, "grok-4.5", now.Add(5*time.Minute))
require.True(t, isGrokTeamModelRateLimited(a1, "grok-4.5", now))
require.True(t, isGrokTeamModelRateLimited(a2, "grok-4.5", now), "sibling with same team must cool")
require.False(t, isGrokTeamModelRateLimited(a2, "grok-4.3", now), "other model stays pickable")
require.False(t, isGrokTeamModelRateLimited(other, "grok-4.5", now))
require.False(t, isGrokTeamModelRateLimited(noTeam, "grok-4.5", now))
filtered := filterGrokTeamModelRateLimitedAccounts([]Account{*a1, *a2, *other, *noTeam}, "grok-4.5", now)
require.Len(t, filtered, 2)
ids := []int64{filtered[0].ID, filtered[1].ID}
require.Contains(t, ids, int64(103))
require.Contains(t, ids, int64(104))
}
func TestGrokTeamModelRateLimit_Expires(t *testing.T) {
team := "team-expire-" + time.Now().Format("150405.000")
a := &Account{
ID: 201, Platform: PlatformGrok, Type: AccountTypeOAuth,
Credentials: map[string]any{"team_id": team},
}
past := time.Now().Add(-time.Minute)
markGrokTeamModelRateLimit(a, "grok-4.5", past)
// mark clamps expired until into default TTL from "now" — use direct store inject via past+recheck
// After mark with past, resolveGrokTeamRateLimitUntil path isn't used; mark uses now+default when until not after now.
require.True(t, isGrokTeamModelRateLimited(a, "grok-4.5", time.Now()))
}
@@ -3,12 +3,25 @@ package service
import (
"context"
"errors"
"hash/fnv"
"strings"
"time"
)
// Base warm window: refresh when access token lifetime remaining is below this.
// Grok access tokens are typically ~1h; refreshing up to 1h early keeps the pool
// warm for request path cache misses.
const grokTokenRefreshSkew = time.Hour
// Stampede spread: each account's effective warm window is reduced by a
// deterministic offset in [0, grokTokenRefreshJitterMax] so co-imported accounts
// do not all refresh in the same TokenRefreshService cycle (grok2api-style
// RefreshDueAt scatter).
const grokTokenRefreshJitterMax = 3 * time.Minute
// Floor so jitter cannot shrink the window below a useful threshold.
const grokTokenRefreshSkewMin = 30 * time.Minute
type GrokTokenRefresher struct {
grokOAuthService GrokOAuthTokenService
}
@@ -40,9 +53,35 @@ func (r *GrokTokenRefresher) NeedsRefresh(account *Account, refreshWindow time.D
if refreshWindow < grokTokenRefreshSkew {
refreshWindow = grokTokenRefreshSkew
}
// Deterministic per-account jitter: spread warm refreshes without random
// non-determinism in tests (hash of account id).
refreshWindow = grokTokenRefreshWindowWithJitter(account.ID, refreshWindow)
return time.Until(*expiresAt) < refreshWindow
}
// grokTokenRefreshWindowWithJitter returns refreshWindow minus a stable offset
// in [0, jitterMax] based on accountID. Result is never below grokTokenRefreshSkewMin
// when the base window is at least that large.
func grokTokenRefreshWindowWithJitter(accountID int64, refreshWindow time.Duration) time.Duration {
if accountID <= 0 || refreshWindow <= grokTokenRefreshSkewMin {
return refreshWindow
}
h := fnv.New32a()
var b [8]byte
id := uint64(accountID)
for i := 0; i < 8; i++ {
b[i] = byte(id >> (8 * i))
}
_, _ = h.Write(b[:])
// Jitter in [0, grokTokenRefreshJitterMax).
jitter := time.Duration(h.Sum32()%uint32(grokTokenRefreshJitterMax/time.Second)) * time.Second
out := refreshWindow - jitter
if out < grokTokenRefreshSkewMin {
return grokTokenRefreshSkewMin
}
return out
}
func (r *GrokTokenRefresher) Refresh(ctx context.Context, account *Account) (map[string]any, error) {
if r == nil || r.grokOAuthService == nil {
return nil, errors.New("grok oauth service is not configured")
@@ -0,0 +1,50 @@
//go:build unit
package service
import (
"testing"
"time"
"github.com/stretchr/testify/require"
)
func TestGrokTokenRefreshWindowWithJitter_StableAndBounded(t *testing.T) {
base := grokTokenRefreshSkew
w1 := grokTokenRefreshWindowWithJitter(42, base)
w2 := grokTokenRefreshWindowWithJitter(42, base)
require.Equal(t, w1, w2, "same account id must yield stable window")
require.GreaterOrEqual(t, w1, grokTokenRefreshSkewMin)
require.LessOrEqual(t, w1, base)
// Different accounts should usually differ (not a hard guarantee for all pairs,
// but for sequential ids the hash spread is good enough to assert inequality
// across a small sample).
seen := map[time.Duration]bool{}
for id := int64(1); id <= 50; id++ {
seen[grokTokenRefreshWindowWithJitter(id, base)] = true
}
require.Greater(t, len(seen), 1, "jitter should spread windows across accounts")
}
func TestGrokTokenRefresher_NeedsRefresh_UsesSkewFloor(t *testing.T) {
refresher := NewGrokTokenRefresher(nil)
// Expires in 50 minutes — within 1h skew, should need refresh.
expires := time.Now().Add(50 * time.Minute).UTC().Format(time.RFC3339)
account := &Account{
ID: 7,
Platform: PlatformGrok,
Type: AccountTypeOAuth,
Credentials: map[string]any{
"access_token": "at",
"refresh_token": "rt",
"expires_at": expires,
},
}
// Pass a tiny window; NeedsRefresh raises it to grokTokenRefreshSkew then jitter.
require.True(t, refresher.NeedsRefresh(account, time.Minute))
// Far future — no refresh.
account.Credentials["expires_at"] = time.Now().Add(3 * time.Hour).UTC().Format(time.RFC3339)
require.False(t, refresher.NeedsRefresh(account, time.Minute))
}
@@ -506,6 +506,11 @@ func (s *defaultOpenAIAccountScheduler) selectBySessionHash(
_ = s.service.deleteStickySessionAccountID(ctx, req.GroupID, sessionHash)
return nil, false, nil
}
// Team+model cool: sticky must not pin a sibling under the same team 429 window.
if account != nil && isGrokTeamModelRateLimited(account, req.RequestedModel, time.Now()) {
_ = s.service.deleteStickySessionAccountID(ctx, req.GroupID, sessionHash)
return nil, false, nil
}
escapeCfg := s.service.openAIStickyEscapeConfig()
if reason, errorRate, ttft, shouldEscape := s.shouldEscapeStickyAccount(accountID, escapeCfg); shouldEscape {
slog.Info("sticky_escape_triggered",
@@ -1348,6 +1353,16 @@ func (s *defaultOpenAIAccountScheduler) selectByLoadBalance(
if len(accounts) == 0 {
return nil, 0, 0, 0, noAvailableOpenAISelectionError(req.RequestedModel, false, openAISelectionFilterStats{}.summary("grok_free_quota_soft_gate"))
}
// Team+model rate-limit cool: siblings of a 429'd team skip the hot model.
if req.Platform == PlatformGrok {
filtered := filterGrokTeamModelRateLimitedAccounts(accounts, req.RequestedModel, time.Now())
if len(filtered) == 0 && len(accounts) > 0 {
return nil, 0, 0, 0, noAvailableOpenAISelectionError(req.RequestedModel, false, openAISelectionFilterStats{}.summary("grok_team_model_rate_limit"))
}
if filtered != nil {
accounts = filtered
}
}
// require_privacy_set: 获取分组信息
var schedGroup *Group
@@ -161,7 +161,13 @@ func (s *OpenAIGatewayService) forwardGrokResponses(
Kind: kind,
Message: upstreamMsg,
})
s.handleGrokAccountUpstreamError(ctx, account, resp.StatusCode, resp.Header, respBody)
errCtx := withGrokTeamRateLimitModel(ctx, upstreamModel)
s.handleGrokAccountUpstreamError(errCtx, account, resp.StatusCode, resp.Header, respBody)
// 429 / free-usage: stamp team+model cool so sibling accounts skip this model.
if resp.StatusCode == http.StatusTooManyRequests ||
classifyGrokUpstreamFailure(resp.StatusCode, respBody, upstreamModel).Class == GrokFailureFreeUsage {
markGrokTeamModelRateLimit(account, upstreamModel, resolveGrokTeamRateLimitUntil(time.Now().Add(grokTeamRateLimitDefaultTTL), time.Now()))
}
if s.shouldFailoverGrokUpstreamError(resp.StatusCode, respBody) {
return nil, &UpstreamFailoverError{
StatusCode: resp.StatusCode,
@@ -173,7 +179,9 @@ func (s *OpenAIGatewayService) forwardGrokResponses(
return s.handleErrorResponse(ctx, resp, c, account, patchedBody, upstreamModel)
}
s.updateGrokUsageFromResponse(ctx, account, resp.Header, resp.StatusCode)
// Attach model so rate-limit snapshots can fan out a team+model cool.
stateCtx := withGrokTeamRateLimitModel(ctx, upstreamModel)
s.updateGrokUsageFromResponse(stateCtx, account, resp.Header, resp.StatusCode)
var usage *OpenAIUsage
var firstTokenMs *int
@@ -1337,7 +1345,8 @@ func (s *OpenAIGatewayService) rateLimitGrok(ctx context.Context, account *Accou
if s == nil || account == nil {
return
}
resetAt = normalizeGrokRateLimitResetAt(account, resetAt, time.Now())
now := time.Now()
resetAt = normalizeGrokRateLimitResetAt(account, resetAt, now)
runtimeUntil := resetAt
if account.TempUnschedulableUntil != nil && account.TempUnschedulableUntil.After(runtimeUntil) {
@@ -1345,6 +1354,27 @@ func (s *OpenAIGatewayService) rateLimitGrok(ctx context.Context, account *Accou
}
s.BlockAccountScheduling(account, runtimeUntil, "429")
persistGrokRateLimit(ctx, s.accountRepo, account, resetAt)
// Propagate a short team+model cool so sibling OAuth accounts on the same
// xAI team skip the hot model without waiting for each to hit 429 alone.
// Model is taken from the latest request context when available; empty is a
// no-op inside markGrokTeamModelRateLimit.
if model, _ := ctx.Value(grokTeamRateLimitModelContextKey{}).(string); model != "" {
markGrokTeamModelRateLimit(account, model, resolveGrokTeamRateLimitUntil(resetAt, now))
}
}
// grokTeamRateLimitModelContextKey carries the upstream model for team cools.
type grokTeamRateLimitModelContextKey struct{}
// withGrokTeamRateLimitModel attaches the upstream model name for rate-limit
// side effects (team+model cool). Safe when model is empty.
func withGrokTeamRateLimitModel(ctx context.Context, model string) context.Context {
model = strings.TrimSpace(model)
if model == "" || ctx == nil {
return ctx
}
return context.WithValue(ctx, grokTeamRateLimitModelContextKey{}, model)
}
func (s *OpenAIGatewayService) handleGrokAccountUpstreamError(ctx context.Context, account *Account, statusCode int, headers http.Header, responseBody []byte) {
@@ -150,6 +150,16 @@ func (s *OpenAIGatewayService) handleStreamingResponseWithReasoning(ctx context.
if s.cfg != nil && s.cfg.Gateway.StreamDataIntervalTimeout > 0 {
streamInterval = time.Duration(s.cfg.Gateway.StreamDataIntervalTimeout) * time.Second
}
// Grok: always enforce an upstream-read idle so hung SSE bodies fail over
// instead of holding the OAuth slot until the client cancels. Prefer the
// global gateway setting when set; otherwise apply a Grok-only default.
if account != nil && account.Platform == PlatformGrok {
cfgSec := 0
if s.cfg != nil {
cfgSec = s.cfg.Gateway.StreamDataIntervalTimeout
}
streamInterval = resolveGrokStreamIdleTimeout(cfgSec)
}
// 仅监控上游数据间隔超时,不被下游写入阻塞影响
var intervalTicker *time.Ticker
if streamInterval > 0 {
@@ -687,6 +697,16 @@ func (s *OpenAIGatewayService) handleStreamingResponseWithReasoning(ctx context.
if s.rateLimitService != nil {
s.rateLimitService.HandleStreamTimeout(ctx, account, originalModel)
}
// Grok: short cool + account failover when no client-visible bytes
// were committed yet (pre-commit). After output started we keep the
// legacy stream_timeout path so partial SSE is not dual-written.
if account != nil && account.Platform == PlatformGrok {
s.tempUnscheduleGrok(ctx, account, grokStreamIdleCooldown, "grok stream idle timeout")
if !openAIStreamClientOutputStarted(c, clientOutputStarted) && !eventShouldFlush {
_ = resp.Body.Close()
return resultWithUsage(), grokStreamIdleFailoverError(account, streamInterval)
}
}
sendErrorEvent("stream_timeout")
return resultWithUsage(), fmt.Errorf("stream data interval timeout")