mirror of
https://github.com/Wei-Shaw/sub2api.git
synced 2026-10-07 12:57:57 +08:00
fix: harden grok oauth gateway paths
This commit is contained in:
@@ -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")
|
||||
}
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user