fix: harden grok oauth gateway paths

This commit is contained in:
Heatherm Huang
2026-06-26 10:36:09 +08:00
parent b3a07aeae7
commit b2e2c7e69c
8 changed files with 229 additions and 19 deletions
@@ -10,6 +10,7 @@ import (
infraerrors "github.com/Wei-Shaw/sub2api/internal/pkg/errors"
"github.com/Wei-Shaw/sub2api/internal/pkg/xai"
"github.com/Wei-Shaw/sub2api/internal/service"
"github.com/Wei-Shaw/sub2api/internal/util/logredact"
"github.com/imroc/req/v3"
)
@@ -105,7 +106,7 @@ func grokOAuthStatusError(code, message string, resp *req.Response) error {
body := ""
if resp != nil {
upstreamStatus = resp.StatusCode
body = resp.String()
body = logredact.RedactText(resp.String())
}
return infraerrors.Newf(statusCode, errorCode, "%s: status %d, body: %s", message, upstreamStatus, body)
}
@@ -85,3 +85,23 @@ func TestGrokOAuthClientRefreshForbiddenClassifiesEntitlement(t *testing.T) {
require.Error(t, err)
require.Contains(t, strings.ToUpper(err.Error()), "GROK_OAUTH_ENTITLEMENT_DENIED")
}
func TestGrokOAuthClientStatusErrorRedactsSensitiveResponseBody(t *testing.T) {
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
w.WriteHeader(http.StatusBadRequest)
_, _ = w.Write([]byte(`{"error":"invalid_grant","access_token":"access-secret","refresh_token":"refresh-secret","code_verifier":"verifier-secret"}`))
}))
defer server.Close()
t.Setenv(xai.EnvTokenURL, server.URL)
client := NewGrokOAuthClient()
_, err := client.RefreshToken(context.Background(), "refresh-secret", "", "client-id")
require.Error(t, err)
errText := err.Error()
require.Contains(t, errText, "status 400")
require.Contains(t, errText, `"refresh_token":"***"`)
require.NotContains(t, errText, "access-secret")
require.NotContains(t, errText, "refresh-secret")
require.NotContains(t, errText, "verifier-secret")
}
+53 -11
View File
@@ -31,7 +31,7 @@ func RegisterGatewayRoutes(
requireGroupAnthropic := middleware.RequireGroupAssignment(settingService, middleware.AnthropicErrorWriter)
requireGroupGoogle := middleware.RequireGroupAssignment(settingService, middleware.GoogleErrorWriter)
isOpenAICompatibleGatewayPlatform := func(c *gin.Context) bool {
isOpenAIResponsesCompatibleGatewayPlatform := func(c *gin.Context) bool {
switch getGroupPlatform(c) {
case service.PlatformOpenAI, service.PlatformGrok:
return true
@@ -39,6 +39,18 @@ func RegisterGatewayRoutes(
return false
}
}
isOpenAIGatewayPlatform := func(c *gin.Context) bool {
return getGroupPlatform(c) == service.PlatformOpenAI
}
rejectGrokUnsupportedEndpoint := func(c *gin.Context, endpoint string) {
service.MarkOpsClientBusinessLimited(c, service.OpsClientBusinessLimitedReasonLocalFeatureGate)
c.JSON(http.StatusNotFound, gin.H{
"error": gin.H{
"type": "not_found_error",
"message": endpoint + " is not supported for Grok groups",
},
})
}
// API网关(Claude API兼容)
gateway := r.Group("/v1")
@@ -51,7 +63,11 @@ func RegisterGatewayRoutes(
{
// /v1/messages: auto-route based on group platform
gateway.POST("/messages", func(c *gin.Context) {
if isOpenAICompatibleGatewayPlatform(c) {
if getGroupPlatform(c) == service.PlatformGrok {
rejectGrokUnsupportedEndpoint(c, "Messages API")
return
}
if isOpenAIGatewayPlatform(c) {
h.OpenAIGateway.Messages(c)
return
}
@@ -59,7 +75,7 @@ func RegisterGatewayRoutes(
})
// /v1/messages/count_tokens: OpenAI groups get 404
gateway.POST("/messages/count_tokens", func(c *gin.Context) {
if isOpenAICompatibleGatewayPlatform(c) {
if isOpenAIResponsesCompatibleGatewayPlatform(c) {
service.MarkOpsClientBusinessLimited(c, service.OpsClientBusinessLimitedReasonLocalFeatureGate)
c.JSON(http.StatusNotFound, gin.H{
"type": "error",
@@ -76,23 +92,33 @@ func RegisterGatewayRoutes(
gateway.GET("/usage", h.Gateway.Usage)
// OpenAI Responses API: auto-route based on group platform
gateway.POST("/responses", func(c *gin.Context) {
if isOpenAICompatibleGatewayPlatform(c) {
if isOpenAIResponsesCompatibleGatewayPlatform(c) {
h.OpenAIGateway.Responses(c)
return
}
h.Gateway.Responses(c)
})
gateway.POST("/responses/*subpath", func(c *gin.Context) {
if isOpenAICompatibleGatewayPlatform(c) {
if isOpenAIResponsesCompatibleGatewayPlatform(c) {
h.OpenAIGateway.Responses(c)
return
}
h.Gateway.Responses(c)
})
gateway.GET("/responses", h.OpenAIGateway.ResponsesWebSocket)
gateway.GET("/responses", func(c *gin.Context) {
if getGroupPlatform(c) == service.PlatformGrok {
rejectGrokUnsupportedEndpoint(c, "Responses WebSocket API")
return
}
h.OpenAIGateway.ResponsesWebSocket(c)
})
// OpenAI Chat Completions API: auto-route based on group platform
gateway.POST("/chat/completions", func(c *gin.Context) {
if isOpenAICompatibleGatewayPlatform(c) {
if getGroupPlatform(c) == service.PlatformGrok {
rejectGrokUnsupportedEndpoint(c, "Chat Completions API")
return
}
if isOpenAIGatewayPlatform(c) {
h.OpenAIGateway.ChatCompletions(c)
return
}
@@ -156,7 +182,7 @@ func RegisterGatewayRoutes(
// OpenAI Responses API(不带v1前缀的别名)— auto-route based on group platform
responsesHandler := func(c *gin.Context) {
if isOpenAICompatibleGatewayPlatform(c) {
if isOpenAIResponsesCompatibleGatewayPlatform(c) {
h.OpenAIGateway.Responses(c)
return
}
@@ -164,17 +190,33 @@ func RegisterGatewayRoutes(
}
r.POST("/responses", bodyLimit, clientRequestID, opsErrorLogger, endpointNorm, gin.HandlerFunc(apiKeyAuth), requireGroupAnthropic, responsesHandler)
r.POST("/responses/*subpath", bodyLimit, clientRequestID, opsErrorLogger, endpointNorm, gin.HandlerFunc(apiKeyAuth), requireGroupAnthropic, responsesHandler)
r.GET("/responses", bodyLimit, clientRequestID, opsErrorLogger, endpointNorm, gin.HandlerFunc(apiKeyAuth), requireGroupAnthropic, h.OpenAIGateway.ResponsesWebSocket)
r.GET("/responses", bodyLimit, clientRequestID, opsErrorLogger, endpointNorm, gin.HandlerFunc(apiKeyAuth), requireGroupAnthropic, func(c *gin.Context) {
if getGroupPlatform(c) == service.PlatformGrok {
rejectGrokUnsupportedEndpoint(c, "Responses WebSocket API")
return
}
h.OpenAIGateway.ResponsesWebSocket(c)
})
codexDirect := r.Group("/backend-api/codex")
codexDirect.Use(bodyLimit, clientRequestID, opsErrorLogger, endpointNorm, gin.HandlerFunc(apiKeyAuth), requireGroupAnthropic)
{
codexDirect.POST("/responses", responsesHandler)
codexDirect.POST("/responses/*subpath", responsesHandler)
codexDirect.GET("/responses", h.OpenAIGateway.ResponsesWebSocket)
codexDirect.GET("/responses", func(c *gin.Context) {
if getGroupPlatform(c) == service.PlatformGrok {
rejectGrokUnsupportedEndpoint(c, "Responses WebSocket API")
return
}
h.OpenAIGateway.ResponsesWebSocket(c)
})
}
// OpenAI Chat Completions API(不带v1前缀的别名)— auto-route based on group platform
r.POST("/chat/completions", bodyLimit, clientRequestID, opsErrorLogger, endpointNorm, gin.HandlerFunc(apiKeyAuth), requireGroupAnthropic, func(c *gin.Context) {
if isOpenAICompatibleGatewayPlatform(c) {
if getGroupPlatform(c) == service.PlatformGrok {
rejectGrokUnsupportedEndpoint(c, "Chat Completions API")
return
}
if isOpenAIGatewayPlatform(c) {
h.OpenAIGateway.ChatCompletions(c)
return
}
+44 -2
View File
@@ -14,10 +14,15 @@ import (
"github.com/stretchr/testify/require"
)
func newGatewayRoutesTestRouter() *gin.Engine {
func newGatewayRoutesTestRouter(platform ...string) *gin.Engine {
gin.SetMode(gin.TestMode)
router := gin.New()
groupPlatform := service.PlatformOpenAI
if len(platform) > 0 && platform[0] != "" {
groupPlatform = platform[0]
}
RegisterGatewayRoutes(
router,
&handler.Handlers{
@@ -28,7 +33,7 @@ func newGatewayRoutesTestRouter() *gin.Engine {
groupID := int64(1)
c.Set(string(servermiddleware.ContextKeyAPIKey), &service.APIKey{
GroupID: &groupID,
Group: &service.Group{Platform: service.PlatformOpenAI},
Group: &service.Group{Platform: groupPlatform},
})
c.Next()
}),
@@ -77,3 +82,40 @@ func TestGatewayRoutesOpenAIImagesPathsAreRegistered(t *testing.T) {
require.NotEqual(t, http.StatusNotFound, w.Code, "path=%s should hit OpenAI images handler", path)
}
}
func TestGatewayRoutesGrokOnlyAllowsResponsesHTTP(t *testing.T) {
router := newGatewayRoutesTestRouter(service.PlatformGrok)
for _, tc := range []struct {
method string
path string
}{
{http.MethodPost, "/v1/messages"},
{http.MethodPost, "/v1/chat/completions"},
{http.MethodPost, "/chat/completions"},
{http.MethodGet, "/v1/responses"},
{http.MethodGet, "/responses"},
{http.MethodGet, "/backend-api/codex/responses"},
} {
req := httptest.NewRequest(tc.method, tc.path, strings.NewReader(`{"model":"grok"}`))
req.Header.Set("Content-Type", "application/json")
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
require.Equal(t, http.StatusNotFound, w.Code, "method=%s path=%s", tc.method, tc.path)
require.Contains(t, w.Body.String(), "not supported for Grok groups")
}
for _, path := range []string{
"/v1/responses",
"/responses",
"/backend-api/codex/responses",
} {
req := httptest.NewRequest(http.MethodPost, path, strings.NewReader(`{"model":"grok","input":"hi"}`))
req.Header.Set("Content-Type", "application/json")
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
require.NotEqual(t, http.StatusNotFound, w.Code, "path=%s should still reach Responses handler", path)
}
}
@@ -152,7 +152,11 @@ func (s *GrokOAuthService) RefreshToken(ctx context.Context, refreshToken, proxy
if err != nil {
return nil, err
}
return s.tokenInfoFromResponse(tokenResp, clientID, nil), nil
tokenInfo := s.tokenInfoFromResponse(tokenResp, clientID, nil)
if tokenInfo.RefreshToken == "" {
tokenInfo.RefreshToken = refreshToken
}
return tokenInfo, nil
}
func (s *GrokOAuthService) ValidateRefreshToken(ctx context.Context, refreshToken string, proxyID *int64) (*GrokTokenInfo, error) {
@@ -0,0 +1,40 @@
//go:build unit
package service
import (
"context"
"testing"
"github.com/Wei-Shaw/sub2api/internal/pkg/xai"
"github.com/stretchr/testify/require"
)
type grokOAuthClientStub struct {
refreshResponse *xai.TokenResponse
}
func (s *grokOAuthClientStub) ExchangeCode(context.Context, string, string, string, string, string) (*xai.TokenResponse, error) {
return &xai.TokenResponse{}, nil
}
func (s *grokOAuthClientStub) RefreshToken(context.Context, string, string, string) (*xai.TokenResponse, error) {
return s.refreshResponse, nil
}
func TestGrokOAuthServiceRefreshTokenPreservesOriginalRefreshTokenWhenNotRotated(t *testing.T) {
svc := NewGrokOAuthService(nil, &grokOAuthClientStub{
refreshResponse: &xai.TokenResponse{
AccessToken: "new-access-token",
TokenType: "Bearer",
ExpiresIn: 3600,
},
})
defer svc.Stop()
info, err := svc.RefreshToken(context.Background(), "original-refresh-token", "", "client-id")
require.NoError(t, err)
require.Equal(t, "new-access-token", info.AccessToken)
require.Equal(t, "original-refresh-token", info.RefreshToken)
require.Equal(t, "client-id", info.ClientID)
}
@@ -7,6 +7,8 @@ import (
"strconv"
"strings"
"time"
"github.com/Wei-Shaw/sub2api/internal/util/logredact"
)
const (
@@ -130,7 +132,7 @@ func (p *GrokTokenProvider) markTempUnschedulable(account *Account, refreshErr e
}
now := time.Now()
until := now.Add(tokenRefreshTempUnschedDuration)
reason := "grok token refresh failed on request path: " + refreshErr.Error()
reason := "grok token refresh failed on request path: " + logredact.RedactText(refreshErr.Error())
bgCtx := context.Background()
if err := p.accountRepo.SetTempUnschedulable(bgCtx, account.ID, until, reason); err != nil {
slog.Warn(grokTokenProviderLogComponent+".set_temp_unschedulable_failed", "account_id", account.ID, "error", err)
@@ -152,8 +154,5 @@ func GrokTokenCacheKey(account *Account) string {
if account == nil {
return "grok:account:0"
}
if email := strings.TrimSpace(account.GetCredential("email")); email != "" {
return "grok:" + email
}
return "grok:account:" + strconv.FormatInt(account.ID, 10)
}
@@ -199,6 +199,68 @@ func TestOpenAITokenCacheKey(t *testing.T) {
}
}
func TestGrokTokenCacheKey(t *testing.T) {
tests := []struct {
name string
account *Account
expected string
}{
{
name: "basic_account",
account: &Account{
ID: 350,
},
expected: "grok:account:350",
},
{
name: "account_with_email_uses_account_id",
account: &Account{
ID: 351,
Credentials: map[string]any{
"email": "same-user@example.com",
},
},
expected: "grok:account:351",
},
{
name: "account_id_zero",
account: &Account{
ID: 0,
},
expected: "grok:account:0",
},
{
name: "nil_account",
account: nil,
expected: "grok:account:0",
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
result := GrokTokenCacheKey(tt.account)
require.Equal(t, tt.expected, result)
})
}
}
func TestGrokTokenCacheKeySeparatesAccountsWithSameEmail(t *testing.T) {
first := &Account{
ID: 351,
Credentials: map[string]any{
"email": "same-user@example.com",
},
}
second := &Account{
ID: 352,
Credentials: map[string]any{
"email": "same-user@example.com",
},
}
require.NotEqual(t, GrokTokenCacheKey(first), GrokTokenCacheKey(second))
}
func TestClaudeTokenCacheKey(t *testing.T) {
tests := []struct {
name string