mirror of
https://github.com/Wei-Shaw/sub2api.git
synced 2026-10-07 17:08:33 +08:00
Merge pull request #4719 from superman2003/fix/grok-403-model-sync-4713-4715
fix(grok): sync OAuth models and isolate policy 403s
This commit is contained in:
@@ -850,6 +850,22 @@ func (s *OpenAIGatewayService) handleGrokMediaErrorResponse(
|
||||
upstreamDetail = truncateString(string(body), maxBytes)
|
||||
}
|
||||
setOpsUpstreamError(c, resp.StatusCode, upstreamMsg, upstreamDetail)
|
||||
if isGrokContentPolicyRejection(resp.StatusCode, body) {
|
||||
clientMsg := grokContentPolicyClientMessage(body)
|
||||
appendOpsUpstreamError(c, OpsUpstreamErrorEvent{
|
||||
Platform: account.Platform,
|
||||
AccountID: account.ID,
|
||||
AccountName: account.Name,
|
||||
UpstreamStatusCode: resp.StatusCode,
|
||||
UpstreamRequestID: requestIDHeader,
|
||||
Kind: "http_error",
|
||||
Message: clientMsg,
|
||||
Detail: upstreamDetail,
|
||||
})
|
||||
MarkResponseCommitted(c)
|
||||
writeGrokMediaErrorResponse(c, http.StatusForbidden, "invalid_request_error", clientMsg)
|
||||
return nil, fmt.Errorf("grok content policy rejection: %s", clientMsg)
|
||||
}
|
||||
|
||||
if status, errType, errMsg, matched := applyErrorPassthroughRule(
|
||||
c,
|
||||
@@ -882,7 +898,7 @@ func (s *OpenAIGatewayService) handleGrokMediaErrorResponse(
|
||||
}
|
||||
|
||||
kind := "http_error"
|
||||
if s.shouldFailoverUpstreamError(resp.StatusCode) {
|
||||
if s.shouldFailoverGrokUpstreamError(resp.StatusCode, body) {
|
||||
kind = "failover"
|
||||
}
|
||||
appendOpsUpstreamError(c, OpsUpstreamErrorEvent{
|
||||
|
||||
@@ -0,0 +1,238 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"strings"
|
||||
"time"
|
||||
)
|
||||
|
||||
// isGrokContentPolicyRejection identifies request-scoped safety refusals from
|
||||
// xAI. These failures are caused by the prompt or media, so retrying another
|
||||
// OAuth account cannot change the outcome and would incorrectly drain a pool.
|
||||
// Keep this matcher deliberately narrow: account entitlement and suspension
|
||||
// messages may mention policy but must retain the normal account failover path.
|
||||
func isGrokContentPolicyRejection(statusCode int, responseBody []byte) bool {
|
||||
if statusCode != http.StatusForbidden || len(responseBody) == 0 {
|
||||
return false
|
||||
}
|
||||
if grokAccountAccessMessage(string(responseBody)) {
|
||||
return false
|
||||
}
|
||||
|
||||
var payload any
|
||||
if json.Unmarshal(responseBody, &payload) == nil {
|
||||
if grokStructuredAccountAccessMarker(payload) {
|
||||
return false
|
||||
}
|
||||
if grokStructuredContentPolicyMarker(payload) {
|
||||
return true
|
||||
}
|
||||
}
|
||||
|
||||
return grokContentPolicyMessage(string(responseBody))
|
||||
}
|
||||
|
||||
func grokStructuredAccountAccessMarker(value any) bool {
|
||||
switch node := value.(type) {
|
||||
case map[string]any:
|
||||
for key, child := range node {
|
||||
normalizedKey := normalizeGrokErrorMarker(key)
|
||||
switch normalizedKey {
|
||||
case "code", "error_code", "type", "category", "reason":
|
||||
if marker, ok := child.(string); ok && isGrokAccountAccessCode(marker) {
|
||||
return true
|
||||
}
|
||||
}
|
||||
if grokStructuredAccountAccessMarker(child) {
|
||||
return true
|
||||
}
|
||||
}
|
||||
case []any:
|
||||
for _, child := range node {
|
||||
if grokStructuredAccountAccessMarker(child) {
|
||||
return true
|
||||
}
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
func grokStructuredContentPolicyMarker(value any) bool {
|
||||
switch node := value.(type) {
|
||||
case map[string]any:
|
||||
for key, child := range node {
|
||||
normalizedKey := normalizeGrokErrorMarker(key)
|
||||
switch normalizedKey {
|
||||
case "code", "error_code", "type", "category", "reason":
|
||||
if marker, ok := child.(string); ok && isGrokContentPolicyCode(marker) {
|
||||
return true
|
||||
}
|
||||
}
|
||||
if grokStructuredContentPolicyMarker(child) {
|
||||
return true
|
||||
}
|
||||
}
|
||||
case []any:
|
||||
for _, child := range node {
|
||||
if grokStructuredContentPolicyMarker(child) {
|
||||
return true
|
||||
}
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
func normalizeGrokErrorMarker(value string) string {
|
||||
value = strings.ToLower(strings.TrimSpace(value))
|
||||
value = strings.ReplaceAll(value, "-", "_")
|
||||
value = strings.ReplaceAll(value, " ", "_")
|
||||
return value
|
||||
}
|
||||
|
||||
func isGrokContentPolicyCode(value string) bool {
|
||||
switch normalizeGrokErrorMarker(value) {
|
||||
case "content_filter",
|
||||
"content_policy",
|
||||
"content_policy_violation",
|
||||
"content_moderation",
|
||||
"cyber_policy",
|
||||
"new_sensitive":
|
||||
return true
|
||||
default:
|
||||
return false
|
||||
}
|
||||
}
|
||||
|
||||
func isGrokAccountAccessCode(value string) bool {
|
||||
switch normalizeGrokErrorMarker(value) {
|
||||
case "account_suspended",
|
||||
"account_disabled",
|
||||
"user_suspended",
|
||||
"user_disabled",
|
||||
"subscription_required",
|
||||
"entitlement_required",
|
||||
"not_entitled",
|
||||
"plan_required",
|
||||
"permission_denied":
|
||||
return true
|
||||
default:
|
||||
return false
|
||||
}
|
||||
}
|
||||
|
||||
func grokAccountAccessMessage(value string) bool {
|
||||
lower := strings.ToLower(strings.TrimSpace(value))
|
||||
for _, phrase := range []string{
|
||||
"account suspended",
|
||||
"account has been suspended",
|
||||
"account disabled",
|
||||
"account has been disabled",
|
||||
"user suspended",
|
||||
"user has been suspended",
|
||||
"subscription required",
|
||||
"entitlement required",
|
||||
"not entitled",
|
||||
} {
|
||||
if strings.Contains(lower, phrase) {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
func grokContentPolicyMessage(value string) bool {
|
||||
lower := strings.ToLower(strings.TrimSpace(value))
|
||||
if lower == "" {
|
||||
return false
|
||||
}
|
||||
|
||||
// xAI's media safety responses use these exact phrases. They are specific
|
||||
// enough not to classify a generic account-policy or entitlement message.
|
||||
for _, phrase := range []string{
|
||||
"the moderation feature is not available",
|
||||
"image is sensitive",
|
||||
"text is sensitive",
|
||||
"prohibited content",
|
||||
"forbidden content",
|
||||
"content policy violation",
|
||||
"content policy rejection",
|
||||
"content policy rejected",
|
||||
"content moderation rejection",
|
||||
"content moderation rejected",
|
||||
"content moderation blocked",
|
||||
"request blocked by content moderation",
|
||||
"request rejected by content moderation",
|
||||
"request blocked by policy",
|
||||
"request rejected by policy",
|
||||
"request violates policy",
|
||||
"prompt violates content policy",
|
||||
"prompt violates policy",
|
||||
"input violates content policy",
|
||||
"input violates policy",
|
||||
} {
|
||||
if strings.Contains(lower, phrase) {
|
||||
return true
|
||||
}
|
||||
}
|
||||
|
||||
return false
|
||||
}
|
||||
|
||||
func grokContentPolicyClientMessage(responseBody []byte) string {
|
||||
message := sanitizeUpstreamErrorMessage(strings.TrimSpace(extractUpstreamErrorMessage(responseBody)))
|
||||
if message == "" {
|
||||
return "Request blocked by upstream content policy"
|
||||
}
|
||||
return message
|
||||
}
|
||||
|
||||
// shouldFailoverGrokUpstreamError is the body-aware counterpart of the
|
||||
// status-only failover helper. Grok content refusals must stay on the current
|
||||
// account and be returned to the caller instead of consuming the account pool.
|
||||
func (s *OpenAIGatewayService) shouldFailoverGrokUpstreamError(statusCode int, responseBody []byte) bool {
|
||||
if isGrokContentPolicyRejection(statusCode, responseBody) {
|
||||
return false
|
||||
}
|
||||
return s.shouldFailoverUpstreamError(statusCode)
|
||||
}
|
||||
|
||||
// applyGrokForbiddenPolicy applies an administrator's existing temporary
|
||||
// unschedulable rules to a non-content 403. It reports true only when a rule
|
||||
// matched; unmatched responses retain the legacy entitlement cooldown.
|
||||
func (s *OpenAIGatewayService) applyGrokForbiddenPolicy(ctx context.Context, account *Account, responseBody []byte) bool {
|
||||
if account == nil || !account.IsTempUnschedulableEnabled() {
|
||||
return false
|
||||
}
|
||||
|
||||
matches := matchTempUnschedulableRules(account, http.StatusForbidden, responseBody)
|
||||
if len(matches) == 0 {
|
||||
return false
|
||||
}
|
||||
|
||||
match := matches[0]
|
||||
// Reuse the central policy implementation when it has a repository. This
|
||||
// preserves the existing reason/cache format and avoids duplicating writes.
|
||||
if s != nil && s.rateLimitService != nil && s.rateLimitService.accountRepo != nil {
|
||||
stateCtx, cancel := openAIAccountStateContext(ctx)
|
||||
handled := s.rateLimitService.tryTempUnschedulable(
|
||||
stateCtx,
|
||||
account,
|
||||
http.StatusForbidden,
|
||||
responseBody,
|
||||
)
|
||||
cancel()
|
||||
if handled {
|
||||
return true
|
||||
}
|
||||
}
|
||||
|
||||
// A partially constructed service (for example a unit-test gateway) still
|
||||
// honors the configured duration instead of silently falling back to 30m.
|
||||
cooldown := time.Duration(match.rule.DurationMinutes) * time.Minute
|
||||
if cooldown > 0 {
|
||||
s.tempUnscheduleGrok(ctx, account, cooldown, "grok configured forbidden rule")
|
||||
}
|
||||
return true
|
||||
}
|
||||
@@ -0,0 +1,319 @@
|
||||
//go:build unit
|
||||
|
||||
package service
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"io"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
func TestIsGrokContentPolicyRejection(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
status int
|
||||
body string
|
||||
want bool
|
||||
}{
|
||||
{
|
||||
name: "new sensitive code",
|
||||
status: http.StatusForbidden,
|
||||
body: `{"error":{"code":"new_sensitive","message":"image is sensitive"}}`,
|
||||
want: true,
|
||||
},
|
||||
{
|
||||
name: "content policy violation code",
|
||||
status: http.StatusForbidden,
|
||||
body: `{"response":{"error":{"code":"content_policy_violation"}}}`,
|
||||
want: true,
|
||||
},
|
||||
{
|
||||
name: "cyber policy code",
|
||||
status: http.StatusForbidden,
|
||||
body: `{"error":{"code":"cyber_policy","message":"request rejected"}}`,
|
||||
want: true,
|
||||
},
|
||||
{
|
||||
name: "moderation feature unavailable",
|
||||
status: http.StatusForbidden,
|
||||
body: `{"error":{"message":"The moderation feature is not available for this request"}}`,
|
||||
want: true,
|
||||
},
|
||||
{
|
||||
name: "explicit prompt moderation rejection",
|
||||
status: http.StatusForbidden,
|
||||
body: `{"error":{"message":"request rejected by content moderation"}}`,
|
||||
want: true,
|
||||
},
|
||||
{
|
||||
name: "entitlement forbidden",
|
||||
status: http.StatusForbidden,
|
||||
body: `{"error":{"message":"subscription required"}}`,
|
||||
want: false,
|
||||
},
|
||||
{
|
||||
name: "account policy suspension is not request policy",
|
||||
status: http.StatusForbidden,
|
||||
body: `{"error":{"message":"account suspended due to policy violation"}}`,
|
||||
want: false,
|
||||
},
|
||||
{
|
||||
name: "structured account suspension overrides policy reason",
|
||||
status: http.StatusForbidden,
|
||||
body: `{"error":{"code":"account_suspended","reason":"policy_violation","message":"account suspended due to policy violation"}}`,
|
||||
want: false,
|
||||
},
|
||||
{
|
||||
name: "ambiguous policy violation code is not enough",
|
||||
status: http.StatusForbidden,
|
||||
body: `{"error":{"code":"policy_violation","message":"policy violation"}}`,
|
||||
want: false,
|
||||
},
|
||||
{
|
||||
name: "policy violation with request scoped message",
|
||||
status: http.StatusForbidden,
|
||||
body: `{"error":{"code":"policy_violation","message":"request blocked by policy"}}`,
|
||||
want: true,
|
||||
},
|
||||
{
|
||||
name: "wrong status",
|
||||
status: http.StatusBadRequest,
|
||||
body: `{"error":{"code":"new_sensitive"}}`,
|
||||
want: false,
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
require.Equal(t, tt.want, isGrokContentPolicyRejection(tt.status, []byte(tt.body)))
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestGrokContentPolicy403DoesNotMutateOrFailover(t *testing.T) {
|
||||
repo := &grokQuotaAccountRepo{}
|
||||
svc := &OpenAIGatewayService{accountRepo: repo}
|
||||
account := &Account{ID: 4715, Platform: PlatformGrok, Type: AccountTypeOAuth}
|
||||
body := []byte(`{"error":{"code":"new_sensitive","message":"text is sensitive"}}`)
|
||||
|
||||
svc.handleGrokAccountUpstreamError(context.Background(), account, http.StatusForbidden, nil, body)
|
||||
|
||||
require.Zero(t, repo.tempUnschedCalls)
|
||||
require.Zero(t, repo.rateLimitedCalls)
|
||||
require.Zero(t, repo.updateCalls)
|
||||
require.False(t, svc.isOpenAIAccountRuntimeBlocked(account))
|
||||
require.False(t, svc.shouldFailoverGrokUpstreamError(http.StatusForbidden, body))
|
||||
|
||||
gin.SetMode(gin.TestMode)
|
||||
recorder := httptest.NewRecorder()
|
||||
c, _ := gin.CreateTestContext(recorder)
|
||||
c.Request = httptest.NewRequest(http.MethodPost, "/v1/messages", nil)
|
||||
resp := &http.Response{StatusCode: http.StatusForbidden, Header: http.Header{}}
|
||||
got := svc.failoverOpenAIUpstreamHTTPError(context.Background(), c, account, resp, body, "text is sensitive", "grok-4.5")
|
||||
require.Nil(t, got)
|
||||
require.Zero(t, repo.tempUnschedCalls)
|
||||
}
|
||||
|
||||
func TestGrokContentPolicy403SharedErrorFallbackDoesNotMutate(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
body := []byte(`{"error":{"code":"content_filter","message":"prohibited content"}}`)
|
||||
repo := &grokQuotaAccountRepo{}
|
||||
svc := &OpenAIGatewayService{accountRepo: repo}
|
||||
account := &Account{
|
||||
ID: 4719,
|
||||
Platform: PlatformGrok,
|
||||
Type: AccountTypeOAuth,
|
||||
Credentials: map[string]any{
|
||||
"custom_error_codes_enabled": true,
|
||||
"custom_error_codes": []any{float64(http.StatusTooManyRequests)},
|
||||
},
|
||||
}
|
||||
|
||||
newContext := func() (*gin.Context, *httptest.ResponseRecorder) {
|
||||
recorder := httptest.NewRecorder()
|
||||
c, _ := gin.CreateTestContext(recorder)
|
||||
c.Request = httptest.NewRequest(http.MethodPost, "/v1/responses", nil)
|
||||
return c, recorder
|
||||
}
|
||||
|
||||
c, recorder := newContext()
|
||||
resp := &http.Response{
|
||||
StatusCode: http.StatusForbidden,
|
||||
Header: http.Header{"Content-Type": []string{"application/json"}},
|
||||
Body: io.NopCloser(strings.NewReader(string(body))),
|
||||
}
|
||||
_, err := svc.handleErrorResponse(context.Background(), resp, c, account, nil, "grok-4.5")
|
||||
require.Error(t, err)
|
||||
require.Equal(t, http.StatusForbidden, recorder.Code)
|
||||
require.Contains(t, recorder.Body.String(), "invalid_request_error")
|
||||
|
||||
c, recorder = newContext()
|
||||
resp = &http.Response{
|
||||
StatusCode: http.StatusForbidden,
|
||||
Header: http.Header{"Content-Type": []string{"application/json"}},
|
||||
Body: io.NopCloser(strings.NewReader(string(body))),
|
||||
}
|
||||
_, err = svc.handleCompatErrorResponse(resp, c, account, writeChatCompletionsError, "grok-4.5")
|
||||
require.Error(t, err)
|
||||
require.Equal(t, http.StatusForbidden, recorder.Code)
|
||||
require.Contains(t, recorder.Body.String(), "invalid_request_error")
|
||||
|
||||
require.Zero(t, repo.tempUnschedCalls)
|
||||
require.Zero(t, repo.rateLimitedCalls)
|
||||
require.Zero(t, repo.updateCalls)
|
||||
}
|
||||
|
||||
func TestGrokContentPolicy403MediaResponseBypassesCustomErrorCodes(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
body := `{"error":{"code":"new_sensitive","message":"image is sensitive"}}`
|
||||
repo := &grokQuotaAccountRepo{}
|
||||
svc := &OpenAIGatewayService{accountRepo: repo}
|
||||
account := &Account{
|
||||
ID: 4720,
|
||||
Platform: PlatformGrok,
|
||||
Type: AccountTypeOAuth,
|
||||
Credentials: map[string]any{
|
||||
"custom_error_codes_enabled": true,
|
||||
"custom_error_codes": []any{float64(http.StatusTooManyRequests)},
|
||||
},
|
||||
}
|
||||
recorder := httptest.NewRecorder()
|
||||
c, _ := gin.CreateTestContext(recorder)
|
||||
c.Request = httptest.NewRequest(http.MethodPost, "/v1/images/generations", nil)
|
||||
resp := &http.Response{
|
||||
StatusCode: http.StatusForbidden,
|
||||
Header: http.Header{"Content-Type": []string{"application/json"}},
|
||||
Body: io.NopCloser(strings.NewReader(body)),
|
||||
}
|
||||
|
||||
_, err := svc.handleGrokMediaErrorResponse(context.Background(), resp, c, account, "request-id", "grok-imagine")
|
||||
require.Error(t, err)
|
||||
require.Equal(t, http.StatusForbidden, recorder.Code)
|
||||
require.Contains(t, recorder.Body.String(), "invalid_request_error")
|
||||
require.Zero(t, repo.tempUnschedCalls)
|
||||
require.Zero(t, repo.rateLimitedCalls)
|
||||
require.Zero(t, repo.updateCalls)
|
||||
}
|
||||
|
||||
func TestGrokContentPolicySSEErrorDoesNotMutateOrFailover(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
repo := &grokQuotaAccountRepo{}
|
||||
upstream := &httpUpstreamRecorder{resp: &http.Response{
|
||||
StatusCode: http.StatusOK,
|
||||
Header: http.Header{"Content-Type": []string{"text/event-stream"}},
|
||||
Body: io.NopCloser(strings.NewReader(
|
||||
"data: {\"type\":\"error\",\"error\":{\"code\":\"new_sensitive\",\"message\":\"text is sensitive\"}}\n\n",
|
||||
)),
|
||||
}}
|
||||
svc := &OpenAIGatewayService{accountRepo: repo, httpUpstream: upstream}
|
||||
account := &Account{ID: 4721, Platform: PlatformGrok, Type: AccountTypeOAuth, Concurrency: 1}
|
||||
recorder := httptest.NewRecorder()
|
||||
c, _ := gin.CreateTestContext(recorder)
|
||||
c.Request = httptest.NewRequest(http.MethodGet, "/v1/responses", nil)
|
||||
payload := []byte(`{"type":"response.create","model":"grok-4.5","input":"hi"}`)
|
||||
var writes [][]byte
|
||||
|
||||
result, err := svc.proxyOpenAIWSHTTPBridgeTurn(
|
||||
context.Background(), c, account, "access-token", payload, len(payload),
|
||||
"grok-4.5", "", "", "", "cache-id", 1,
|
||||
func(message []byte) error {
|
||||
writes = append(writes, append([]byte(nil), message...))
|
||||
return nil
|
||||
},
|
||||
)
|
||||
|
||||
require.Error(t, err)
|
||||
require.NotNil(t, result)
|
||||
var failoverErr *UpstreamFailoverError
|
||||
require.False(t, errors.As(err, &failoverErr))
|
||||
require.Len(t, writes, 1)
|
||||
require.Contains(t, string(writes[0]), "new_sensitive")
|
||||
require.Zero(t, repo.tempUnschedCalls)
|
||||
require.Zero(t, repo.rateLimitedCalls)
|
||||
require.Zero(t, repo.updateCalls)
|
||||
require.False(t, svc.isOpenAIAccountRuntimeBlocked(account))
|
||||
}
|
||||
|
||||
func TestHandleGrokAccountUpstreamErrorEntitlement403KeepsDefaultCooldown(t *testing.T) {
|
||||
repo := &grokQuotaAccountRepo{}
|
||||
svc := &OpenAIGatewayService{accountRepo: repo}
|
||||
account := &Account{ID: 4716, Platform: PlatformGrok, Type: AccountTypeOAuth}
|
||||
before := time.Now()
|
||||
|
||||
svc.handleGrokAccountUpstreamError(
|
||||
context.Background(), account, http.StatusForbidden, nil,
|
||||
[]byte(`{"error":{"message":"subscription required"}}`),
|
||||
)
|
||||
|
||||
require.Equal(t, 1, repo.tempUnschedCalls)
|
||||
require.Equal(t, "grok access or entitlement denied", repo.lastTempUnschedReason)
|
||||
require.Greater(t, repo.lastTempUnschedUntil, before.Add(29*time.Minute))
|
||||
require.Less(t, repo.lastTempUnschedUntil, before.Add(31*time.Minute))
|
||||
}
|
||||
|
||||
func TestHandleGrokAccountUpstreamError403UsesConfiguredRule(t *testing.T) {
|
||||
repo := &grokQuotaAccountRepo{}
|
||||
svc := &OpenAIGatewayService{accountRepo: repo}
|
||||
account := &Account{
|
||||
ID: 4717,
|
||||
Platform: PlatformGrok,
|
||||
Type: AccountTypeOAuth,
|
||||
Credentials: map[string]any{
|
||||
"temp_unschedulable_enabled": true,
|
||||
"temp_unschedulable_rules": []any{
|
||||
map[string]any{
|
||||
"error_code": float64(http.StatusForbidden),
|
||||
"keywords": []any{"subscription"},
|
||||
"duration_minutes": float64(7),
|
||||
},
|
||||
},
|
||||
},
|
||||
}
|
||||
before := time.Now()
|
||||
|
||||
svc.handleGrokAccountUpstreamError(
|
||||
context.Background(), account, http.StatusForbidden, nil,
|
||||
[]byte(`{"error":{"message":"subscription required"}}`),
|
||||
)
|
||||
|
||||
require.Equal(t, 1, repo.tempUnschedCalls)
|
||||
require.Greater(t, repo.lastTempUnschedUntil, before.Add(6*time.Minute))
|
||||
require.Less(t, repo.lastTempUnschedUntil, before.Add(8*time.Minute))
|
||||
}
|
||||
|
||||
func TestHandleGrokAccountUpstreamError403ConfiguredUnmatchedKeepsDefaultCooldown(t *testing.T) {
|
||||
repo := &grokQuotaAccountRepo{}
|
||||
svc := &OpenAIGatewayService{accountRepo: repo}
|
||||
account := &Account{
|
||||
ID: 4718,
|
||||
Platform: PlatformGrok,
|
||||
Type: AccountTypeOAuth,
|
||||
Credentials: map[string]any{
|
||||
"temp_unschedulable_enabled": true,
|
||||
"temp_unschedulable_rules": []any{
|
||||
map[string]any{
|
||||
"error_code": float64(http.StatusForbidden),
|
||||
"keywords": []any{"different failure"},
|
||||
"duration_minutes": float64(7),
|
||||
},
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
svc.handleGrokAccountUpstreamError(
|
||||
context.Background(), account, http.StatusForbidden, nil,
|
||||
[]byte(`{"error":{"message":"subscription required"}}`),
|
||||
)
|
||||
|
||||
require.Equal(t, 1, repo.tempUnschedCalls)
|
||||
require.Equal(t, "grok access or entitlement denied", repo.lastTempUnschedReason)
|
||||
require.True(t, svc.isOpenAIAccountRuntimeBlocked(account))
|
||||
}
|
||||
@@ -48,6 +48,9 @@ func isOpenAIAccount(account *Account) bool {
|
||||
// handleOpenAIAccountUpstreamError expects canonicalModel to be the model used
|
||||
// for scheduling after applying account mapping exactly once.
|
||||
func (s *OpenAIGatewayService) handleOpenAIAccountUpstreamError(ctx context.Context, account *Account, statusCode int, headers http.Header, responseBody []byte, canonicalModel ...string) bool {
|
||||
if account != nil && account.Platform == PlatformGrok && isGrokContentPolicyRejection(statusCode, responseBody) {
|
||||
return false
|
||||
}
|
||||
stateCtx, cancel := openAIAccountStateContext(ctx)
|
||||
defer cancel()
|
||||
|
||||
|
||||
@@ -88,10 +88,14 @@ func (s *OpenAIGatewayService) failoverOpenAIUpstreamHTTPError(
|
||||
upstreamMsg string,
|
||||
upstreamModel string,
|
||||
) *UpstreamFailoverError {
|
||||
shouldFailover := s.shouldFailoverOpenAIUpstreamResponse(resp.StatusCode, upstreamMsg, respBody)
|
||||
if account != nil && account.Platform == PlatformGrok {
|
||||
shouldFailover = s.shouldFailoverGrokUpstreamError(resp.StatusCode, respBody)
|
||||
}
|
||||
if account != nil && account.Platform == PlatformGrok {
|
||||
s.handleGrokAccountUpstreamError(ctx, account, resp.StatusCode, resp.Header, respBody)
|
||||
}
|
||||
if !s.shouldFailoverOpenAIUpstreamResponse(resp.StatusCode, upstreamMsg, respBody) {
|
||||
if !shouldFailover {
|
||||
return nil
|
||||
}
|
||||
upstreamDetail := ""
|
||||
|
||||
@@ -174,17 +174,21 @@ func (s *OpenAIGatewayService) forwardAsRawChatCompletions(
|
||||
if resp.StatusCode >= 400 {
|
||||
respBody, upstreamMsg := s.readOpenAIUpstreamError(resp)
|
||||
if account.Platform == PlatformGrok {
|
||||
kind := "http_error"
|
||||
if s.shouldFailoverGrokUpstreamError(resp.StatusCode, respBody) {
|
||||
kind = "failover"
|
||||
}
|
||||
appendOpsUpstreamError(c, OpsUpstreamErrorEvent{
|
||||
Platform: account.Platform,
|
||||
AccountID: account.ID,
|
||||
AccountName: account.Name,
|
||||
UpstreamStatusCode: resp.StatusCode,
|
||||
UpstreamRequestID: firstNonEmpty(resp.Header.Get("x-request-id"), resp.Header.Get("xai-request-id")),
|
||||
Kind: "failover",
|
||||
Kind: kind,
|
||||
Message: upstreamMsg,
|
||||
})
|
||||
s.handleGrokAccountUpstreamError(ctx, account, resp.StatusCode, resp.Header, respBody)
|
||||
if s.shouldFailoverUpstreamError(resp.StatusCode) {
|
||||
if s.shouldFailoverGrokUpstreamError(resp.StatusCode, respBody) {
|
||||
return nil, &UpstreamFailoverError{
|
||||
StatusCode: resp.StatusCode,
|
||||
ResponseBody: respBody,
|
||||
|
||||
@@ -148,17 +148,21 @@ func (s *OpenAIGatewayService) forwardGrokResponses(
|
||||
if upstreamMsg == "" {
|
||||
upstreamMsg = fmt.Sprintf("xAI upstream returned status %d", resp.StatusCode)
|
||||
}
|
||||
kind := "http_error"
|
||||
if s.shouldFailoverGrokUpstreamError(resp.StatusCode, respBody) {
|
||||
kind = "failover"
|
||||
}
|
||||
appendOpsUpstreamError(c, OpsUpstreamErrorEvent{
|
||||
Platform: account.Platform,
|
||||
AccountID: account.ID,
|
||||
AccountName: account.Name,
|
||||
UpstreamStatusCode: resp.StatusCode,
|
||||
UpstreamRequestID: firstNonEmpty(resp.Header.Get("x-request-id"), resp.Header.Get("xai-request-id")),
|
||||
Kind: "failover",
|
||||
Kind: kind,
|
||||
Message: upstreamMsg,
|
||||
})
|
||||
s.handleGrokAccountUpstreamError(ctx, account, resp.StatusCode, resp.Header, respBody)
|
||||
if s.shouldFailoverUpstreamError(resp.StatusCode) {
|
||||
if s.shouldFailoverGrokUpstreamError(resp.StatusCode, respBody) {
|
||||
return nil, &UpstreamFailoverError{
|
||||
StatusCode: resp.StatusCode,
|
||||
ResponseBody: respBody,
|
||||
@@ -904,17 +908,21 @@ func (s *OpenAIGatewayService) describeGrokComposerImage(
|
||||
if upstreamMsg == "" {
|
||||
upstreamMsg = fmt.Sprintf("xAI image bridge upstream returned status %d", resp.StatusCode)
|
||||
}
|
||||
kind := "http_error"
|
||||
if s.shouldFailoverGrokUpstreamError(resp.StatusCode, respBody) {
|
||||
kind = "failover"
|
||||
}
|
||||
appendOpsUpstreamError(c, OpsUpstreamErrorEvent{
|
||||
Platform: account.Platform,
|
||||
AccountID: account.ID,
|
||||
AccountName: account.Name,
|
||||
UpstreamStatusCode: resp.StatusCode,
|
||||
UpstreamRequestID: firstNonEmpty(resp.Header.Get("x-request-id"), resp.Header.Get("xai-request-id")),
|
||||
Kind: "failover",
|
||||
Kind: kind,
|
||||
Message: upstreamMsg,
|
||||
})
|
||||
s.handleGrokAccountUpstreamError(ctx, account, resp.StatusCode, resp.Header, respBody)
|
||||
if s.shouldFailoverUpstreamError(resp.StatusCode) {
|
||||
if s.shouldFailoverGrokUpstreamError(resp.StatusCode, respBody) {
|
||||
return "", OpenAIUsage{}, &UpstreamFailoverError{
|
||||
StatusCode: resp.StatusCode,
|
||||
ResponseBody: respBody,
|
||||
@@ -1082,6 +1090,7 @@ func applyGrokCLIHeaders(headers http.Header) {
|
||||
}
|
||||
headers.Set("User-Agent", grokUpstreamUserAgent)
|
||||
headers.Set("X-Grok-Client-Version", grokCLIVersion)
|
||||
headers.Set("X-Grok-Client-Mode", "interactive")
|
||||
}
|
||||
|
||||
func (s *OpenAIGatewayService) updateGrokUsageSnapshot(ctx context.Context, account *Account, snapshot *xai.QuotaSnapshot) {
|
||||
@@ -1336,12 +1345,18 @@ func (s *OpenAIGatewayService) handleGrokAccountUpstreamError(ctx context.Contex
|
||||
if s == nil || account == nil {
|
||||
return
|
||||
}
|
||||
if isGrokContentPolicyRejection(statusCode, responseBody) {
|
||||
return
|
||||
}
|
||||
now := time.Now()
|
||||
s.updateGrokUsageSnapshot(ctx, account, parseGrokQuotaSnapshot(headers, statusCode, now))
|
||||
switch statusCode {
|
||||
case http.StatusUnauthorized:
|
||||
s.tempUnscheduleGrok(ctx, account, 10*time.Minute, "grok credentials unauthorized")
|
||||
case http.StatusForbidden:
|
||||
if s.applyGrokForbiddenPolicy(ctx, account, responseBody) {
|
||||
return
|
||||
}
|
||||
s.tempUnscheduleGrok(ctx, account, 30*time.Minute, "grok access or entitlement denied")
|
||||
case http.StatusTooManyRequests:
|
||||
// updateGrokUsageSnapshot installs both runtime and durable rate-limit state.
|
||||
|
||||
@@ -608,17 +608,21 @@ func (s *OpenAIGatewayService) forwardGrokChatCompletionsViaResponses(
|
||||
if upstreamMsg == "" {
|
||||
upstreamMsg = fmt.Sprintf("xAI upstream returned status %d", resp.StatusCode)
|
||||
}
|
||||
kind := "http_error"
|
||||
if s.shouldFailoverGrokUpstreamError(resp.StatusCode, respBody) {
|
||||
kind = "failover"
|
||||
}
|
||||
appendOpsUpstreamError(c, OpsUpstreamErrorEvent{
|
||||
Platform: account.Platform,
|
||||
AccountID: account.ID,
|
||||
AccountName: account.Name,
|
||||
UpstreamStatusCode: resp.StatusCode,
|
||||
UpstreamRequestID: firstNonEmpty(resp.Header.Get("x-request-id"), resp.Header.Get("xai-request-id")),
|
||||
Kind: "failover",
|
||||
Kind: kind,
|
||||
Message: upstreamMsg,
|
||||
})
|
||||
s.handleGrokAccountUpstreamError(ctx, account, resp.StatusCode, resp.Header, respBody)
|
||||
if s.shouldFailoverUpstreamError(resp.StatusCode) {
|
||||
if s.shouldFailoverGrokUpstreamError(resp.StatusCode, respBody) {
|
||||
return nil, &UpstreamFailoverError{
|
||||
StatusCode: resp.StatusCode,
|
||||
ResponseBody: respBody,
|
||||
|
||||
@@ -345,6 +345,19 @@ func (s *OpenAIGatewayService) handleErrorResponse(
|
||||
}
|
||||
return nil, fmt.Errorf("openai cyber_policy: %s", cyberMsg)
|
||||
}
|
||||
if account != nil && account.Platform == PlatformGrok && isGrokContentPolicyRejection(resp.StatusCode, body) {
|
||||
clientMsg := grokContentPolicyClientMessage(body)
|
||||
setOpsUpstreamError(c, resp.StatusCode, clientMsg, truncateString(string(body), 2048))
|
||||
writeOpenAIPassthroughResponseHeaders(c.Writer.Header(), resp.Header, s.responseHeaderFilter)
|
||||
MarkResponseCommitted(c)
|
||||
c.JSON(http.StatusForbidden, gin.H{
|
||||
"error": gin.H{
|
||||
"type": "invalid_request_error",
|
||||
"message": clientMsg,
|
||||
},
|
||||
})
|
||||
return nil, fmt.Errorf("grok content policy rejection: %s", clientMsg)
|
||||
}
|
||||
|
||||
upstreamMsg := strings.TrimSpace(extractUpstreamErrorMessage(body))
|
||||
upstreamMsg = sanitizeUpstreamErrorMessage(upstreamMsg)
|
||||
@@ -559,6 +572,13 @@ func (s *OpenAIGatewayService) handleCompatErrorResponse(
|
||||
}
|
||||
return nil, fmt.Errorf("openai cyber_policy: %s", cyberMsg)
|
||||
}
|
||||
if account != nil && account.Platform == PlatformGrok && isGrokContentPolicyRejection(resp.StatusCode, body) {
|
||||
clientMsg := grokContentPolicyClientMessage(body)
|
||||
setOpsUpstreamError(c, resp.StatusCode, clientMsg, truncateString(string(body), 2048))
|
||||
MarkResponseCommitted(c)
|
||||
writeError(c, http.StatusForbidden, "invalid_request_error", clientMsg)
|
||||
return nil, fmt.Errorf("grok content policy rejection: %s", clientMsg)
|
||||
}
|
||||
|
||||
upstreamMsg := strings.TrimSpace(extractUpstreamErrorMessage(body))
|
||||
if upstreamMsg == "" {
|
||||
|
||||
@@ -248,6 +248,7 @@ func (s *OpenAIGatewayService) proxyOpenAIWSHTTPBridgeTurn(
|
||||
}
|
||||
shouldFailover := s.shouldFailoverOpenAIUpstreamResponse(resp.StatusCode, upstreamMsg, respBody)
|
||||
if account.Platform == PlatformGrok {
|
||||
shouldFailover = s.shouldFailoverGrokUpstreamError(resp.StatusCode, respBody)
|
||||
s.handleGrokAccountUpstreamError(ctx, account, resp.StatusCode, resp.Header, respBody)
|
||||
if turn == 1 && shouldFailover {
|
||||
return nil, newOpenAIUpstreamFailoverError(resp.StatusCode, resp.Header, respBody, upstreamMsg, false)
|
||||
@@ -393,7 +394,16 @@ func (s *OpenAIGatewayService) proxyOpenAIWSHTTPBridgeTurn(
|
||||
statusCode := openAIWSErrorHTTPStatusFromRaw(errCodeRaw, errTypeRaw)
|
||||
shouldFailover := s.shouldFailoverOpenAIUpstreamResponse(statusCode, errMessage, upstreamMessage)
|
||||
if account.Platform == PlatformGrok {
|
||||
s.handleGrokAccountUpstreamError(ctx, account, statusCode, resp.Header, upstreamMessage)
|
||||
// SSE error events do not carry an HTTP status. The local status
|
||||
// mapper therefore defaults unknown xAI codes (for example
|
||||
// new_sensitive) to 502; classify the body as a request-scoped
|
||||
// 403 before applying status-based failover or account state.
|
||||
if isGrokContentPolicyRejection(http.StatusForbidden, upstreamMessage) {
|
||||
shouldFailover = false
|
||||
} else {
|
||||
shouldFailover = s.shouldFailoverGrokUpstreamError(statusCode, upstreamMessage)
|
||||
s.handleGrokAccountUpstreamError(ctx, account, statusCode, resp.Header, upstreamMessage)
|
||||
}
|
||||
} else if shouldFailover {
|
||||
accountStatus := statusCode
|
||||
if transientStatus := openAIWSPayloadTransientStatus(upstreamMessage); transientStatus != 0 {
|
||||
|
||||
@@ -116,7 +116,11 @@ func (s *AccountTestService) FetchUpstreamSupportedModels(ctx context.Context, a
|
||||
)
|
||||
}
|
||||
|
||||
models, err := extractUpstreamModelIDs(body)
|
||||
extractModels := extractUpstreamModelIDs
|
||||
if account.IsGrok() {
|
||||
extractModels = extractGrokUpstreamModelIDs
|
||||
}
|
||||
models, err := extractModels(body)
|
||||
if err != nil {
|
||||
return nil, newUpstreamModelSyncUpstreamError("Upstream model list response was not valid JSON", err)
|
||||
}
|
||||
@@ -147,31 +151,79 @@ func (s *AccountTestService) buildUpstreamModelsRequest(ctx context.Context, acc
|
||||
}
|
||||
|
||||
func (s *AccountTestService) buildGrokUpstreamModelsRequest(ctx context.Context, account *Account) (*http.Request, error) {
|
||||
if account.Type != AccountTypeAPIKey {
|
||||
if account == nil {
|
||||
return nil, newUpstreamModelSyncConfigError("Account is required", nil)
|
||||
}
|
||||
|
||||
var (
|
||||
authToken string
|
||||
normalizedBaseURL string
|
||||
isOAuth = account.IsGrokOAuth()
|
||||
)
|
||||
switch account.Type {
|
||||
case AccountTypeAPIKey:
|
||||
authToken = strings.TrimSpace(account.GetCredential("api_key"))
|
||||
if authToken == "" {
|
||||
return nil, newUpstreamModelSyncConfigError("No Grok API key is available", nil)
|
||||
}
|
||||
|
||||
baseURL := strings.TrimSpace(account.GetCredential("base_url"))
|
||||
if baseURL == "" {
|
||||
baseURL = "https://api.x.ai"
|
||||
}
|
||||
validatedBaseURL, err := s.validateUpstreamBaseURL(baseURL)
|
||||
if err != nil {
|
||||
return nil, newUpstreamModelSyncConfigError("Invalid Grok base URL", err)
|
||||
}
|
||||
normalizedBaseURL = validatedBaseURL
|
||||
case AccountTypeOAuth:
|
||||
if s.grokTokenProvider == nil {
|
||||
return nil, newUpstreamModelSyncConfigError("Grok token provider is not configured", nil)
|
||||
}
|
||||
accessToken, err := s.grokTokenProvider.GetAccessTokenForManualTest(ctx, account)
|
||||
if err != nil {
|
||||
return nil, newUpstreamModelSyncUpstreamError("Failed to get Grok access token", err)
|
||||
}
|
||||
authToken = strings.TrimSpace(accessToken)
|
||||
if authToken == "" {
|
||||
return nil, newUpstreamModelSyncConfigError("No Grok access token is available", nil)
|
||||
}
|
||||
|
||||
validator, err := grokBaseURLValidator(account, s.cfg)
|
||||
if err != nil {
|
||||
return nil, newUpstreamModelSyncConfigError("Invalid Grok base URL", err)
|
||||
}
|
||||
validatedBaseURL, err := validator(account.GetGrokBaseURL())
|
||||
if err != nil {
|
||||
return nil, newUpstreamModelSyncConfigError("Invalid Grok base URL", err)
|
||||
}
|
||||
normalizedBaseURL = validatedBaseURL
|
||||
default:
|
||||
return nil, newUpstreamModelSyncUnsupportedError(
|
||||
fmt.Sprintf("Unsupported Grok account type for upstream model sync: %s", account.Type), nil,
|
||||
)
|
||||
}
|
||||
apiKey := strings.TrimSpace(account.GetCredential("api_key"))
|
||||
if apiKey == "" {
|
||||
return nil, newUpstreamModelSyncConfigError("No Grok API key is available", nil)
|
||||
}
|
||||
|
||||
baseURL := strings.TrimSpace(account.GetCredential("base_url"))
|
||||
if baseURL == "" {
|
||||
baseURL = "https://api.x.ai"
|
||||
}
|
||||
normalizedBaseURL, err := s.validateUpstreamBaseURL(baseURL)
|
||||
if err != nil {
|
||||
return nil, newUpstreamModelSyncConfigError("Invalid Grok base URL", err)
|
||||
}
|
||||
|
||||
req, err := http.NewRequestWithContext(ctx, http.MethodGet, buildOpenAIModelsURL(normalizedBaseURL), nil)
|
||||
if err != nil {
|
||||
return nil, newUpstreamModelSyncConfigError("Invalid Grok model list URL", err)
|
||||
}
|
||||
req.Header.Set("Accept", "application/json")
|
||||
req.Header.Set("Authorization", "Bearer "+apiKey)
|
||||
req.Header.Set("Authorization", "Bearer "+authToken)
|
||||
if isOAuth {
|
||||
// The shared HTTP transport adds the official CLI marker/version for the
|
||||
// exact proxy host. Keep the request builder aligned with the other Grok
|
||||
// probes and only forward account identity headers to that trusted host.
|
||||
applyGrokCLIHeaders(req.Header)
|
||||
if isGrokCLIProxyTarget(req.URL.String()) {
|
||||
if userID := strings.TrimSpace(account.GetCredential("sub")); userID != "" {
|
||||
req.Header.Set("X-UserID", userID)
|
||||
}
|
||||
if email := strings.TrimSpace(account.GetCredential("email")); email != "" {
|
||||
req.Header.Set("X-Email", email)
|
||||
}
|
||||
}
|
||||
}
|
||||
account.ApplyHeaderOverrides(req.Header)
|
||||
return req, nil
|
||||
}
|
||||
@@ -438,11 +490,31 @@ func buildGeminiModelsURL(base string) string {
|
||||
}
|
||||
|
||||
type upstreamModelEntry struct {
|
||||
ID string `json:"id"`
|
||||
Name string `json:"name"`
|
||||
ID string `json:"id"`
|
||||
Model string `json:"model"`
|
||||
ModelID string `json:"modelId"`
|
||||
ModelIDSnake string `json:"model_id"`
|
||||
Name string `json:"name"`
|
||||
Meta json.RawMessage `json:"_meta"`
|
||||
}
|
||||
|
||||
type upstreamModelEntryMetadata struct {
|
||||
ID string `json:"id"`
|
||||
Model string `json:"model"`
|
||||
ModelID string `json:"modelId"`
|
||||
ModelIDSnake string `json:"model_id"`
|
||||
Name string `json:"name"`
|
||||
}
|
||||
|
||||
func extractUpstreamModelIDs(body []byte) ([]string, error) {
|
||||
return extractUpstreamModelIDsWithSelector(body, upstreamModelEntryID)
|
||||
}
|
||||
|
||||
func extractGrokUpstreamModelIDs(body []byte) ([]string, error) {
|
||||
return extractUpstreamModelIDsWithSelector(body, grokUpstreamModelEntryID)
|
||||
}
|
||||
|
||||
func extractUpstreamModelIDsWithSelector(body []byte, selectID func(upstreamModelEntry) string) ([]string, error) {
|
||||
var response struct {
|
||||
Data []upstreamModelEntry `json:"data"`
|
||||
Models []upstreamModelEntry `json:"models"`
|
||||
@@ -455,24 +527,24 @@ func extractUpstreamModelIDs(body []byte) ([]string, error) {
|
||||
|
||||
models := make([]string, 0, len(arrayResponse))
|
||||
for _, entry := range arrayResponse {
|
||||
models = append(models, upstreamModelEntryID(entry))
|
||||
models = append(models, selectID(entry))
|
||||
}
|
||||
return dedupeAndSortModelIDs(models), nil
|
||||
}
|
||||
|
||||
models := make([]string, 0, len(response.Data)+len(response.Models))
|
||||
for _, entry := range response.Data {
|
||||
models = append(models, upstreamModelEntryID(entry))
|
||||
models = append(models, selectID(entry))
|
||||
}
|
||||
for _, entry := range response.Models {
|
||||
models = append(models, upstreamModelEntryID(entry))
|
||||
models = append(models, selectID(entry))
|
||||
}
|
||||
|
||||
if len(models) == 0 {
|
||||
var arrayResponse []upstreamModelEntry
|
||||
if err := json.Unmarshal(body, &arrayResponse); err == nil {
|
||||
for _, entry := range arrayResponse {
|
||||
models = append(models, upstreamModelEntryID(entry))
|
||||
models = append(models, selectID(entry))
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -488,6 +560,37 @@ func upstreamModelEntryID(entry upstreamModelEntry) string {
|
||||
return strings.TrimPrefix(modelID, "models/")
|
||||
}
|
||||
|
||||
func grokUpstreamModelEntryID(entry upstreamModelEntry) string {
|
||||
candidates := []string{
|
||||
entry.Model,
|
||||
entry.ModelID,
|
||||
entry.ModelIDSnake,
|
||||
entry.ID,
|
||||
}
|
||||
if len(entry.Meta) > 0 {
|
||||
var meta upstreamModelEntryMetadata
|
||||
if err := json.Unmarshal(entry.Meta, &meta); err == nil {
|
||||
candidates = append(candidates,
|
||||
meta.Model,
|
||||
meta.ModelID,
|
||||
meta.ModelIDSnake,
|
||||
meta.ID,
|
||||
meta.Name,
|
||||
)
|
||||
}
|
||||
}
|
||||
// `name` is a display label in the Grok catalog, so keep it as the final
|
||||
// compatibility fallback rather than preferring it over protocol model IDs.
|
||||
candidates = append(candidates, entry.Name)
|
||||
for _, candidate := range candidates {
|
||||
modelID := strings.TrimSpace(candidate)
|
||||
if modelID != "" {
|
||||
return strings.TrimPrefix(modelID, "models/")
|
||||
}
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
func dedupeAndSortModelIDs(models []string) []string {
|
||||
seen := make(map[string]struct{}, len(models))
|
||||
result := make([]string, 0, len(models))
|
||||
|
||||
@@ -7,6 +7,7 @@ import (
|
||||
"net/http"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/Wei-Shaw/sub2api/internal/config"
|
||||
"github.com/stretchr/testify/require"
|
||||
@@ -20,6 +21,25 @@ func upstreamModelSyncTestConfig() *config.Config {
|
||||
}
|
||||
}
|
||||
|
||||
func grokOAuthModelSyncTestAccount(baseURL string) *Account {
|
||||
credentials := map[string]any{
|
||||
"access_token": "oauth-access-token",
|
||||
"refresh_token": "oauth-refresh-token",
|
||||
"expires_at": time.Now().Add(time.Hour).UTC().Format(time.RFC3339),
|
||||
"sub": "grok-user-id",
|
||||
"email": "grok-user@example.com",
|
||||
}
|
||||
if strings.TrimSpace(baseURL) != "" {
|
||||
credentials["base_url"] = baseURL
|
||||
}
|
||||
return &Account{
|
||||
ID: 10,
|
||||
Platform: PlatformGrok,
|
||||
Type: AccountTypeOAuth,
|
||||
Credentials: credentials,
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildV1ModelsURL(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
@@ -115,6 +135,11 @@ func TestExtractUpstreamModelIDs(t *testing.T) {
|
||||
body: `[{"id":"z-model"},{"name":"models/a-model"}]`,
|
||||
want: []string{"a-model", "z-model"},
|
||||
},
|
||||
{
|
||||
name: "standard id wins over provider-specific model field",
|
||||
body: `{"data":[{"id":"canonical-id","model":"display-model"}]}`,
|
||||
want: []string{"canonical-id"},
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
@@ -129,6 +154,14 @@ func TestExtractUpstreamModelIDs(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestExtractGrokUpstreamModelIDs(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
models, err := extractGrokUpstreamModelIDs([]byte(`{"data":[{"id":"display-id","model":"grok-4.5"},{"modelId":"grok-build-0.1"},{"model_id":"grok-composer-2.5-fast"},{"name":"Grok Meta Display Name","_meta":{"model":"grok-meta"}},{"name":"grok-name"},{"id":"grok-safe","_meta":"not-an-object"}]}`))
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, []string{"grok-4.5", "grok-build-0.1", "grok-composer-2.5-fast", "grok-meta", "grok-name", "grok-safe"}, models)
|
||||
}
|
||||
|
||||
func TestBuildUpstreamModelsRequestsForAPIKeyAccounts(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
@@ -214,20 +247,36 @@ func TestBuildUpstreamModelsRequestsForAPIKeyAccounts(t *testing.T) {
|
||||
require.Equal(t, "antigravity-key", antigravityReq.Header.Get("x-api-key"))
|
||||
}
|
||||
|
||||
func TestBuildUpstreamModelsRequestRejectsGrokOAuth(t *testing.T) {
|
||||
func TestBuildUpstreamModelsRequestSupportsGrokOAuth(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
svc := &AccountTestService{
|
||||
cfg: upstreamModelSyncTestConfig(),
|
||||
grokTokenProvider: NewGrokTokenProvider(nil, nil),
|
||||
}
|
||||
req, err := svc.buildUpstreamModelsRequest(context.Background(), grokOAuthModelSyncTestAccount(""))
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, "https://cli-chat-proxy.grok.com/v1/models", req.URL.String())
|
||||
require.Equal(t, "Bearer oauth-access-token", req.Header.Get("Authorization"))
|
||||
require.Equal(t, grokCLIVersion, req.Header.Get("X-Grok-Client-Version"))
|
||||
require.Equal(t, "interactive", req.Header.Get("X-Grok-Client-Mode"))
|
||||
require.Equal(t, grokUpstreamUserAgent, req.Header.Get("User-Agent"))
|
||||
require.Equal(t, "grok-user-id", req.Header.Get("X-UserID"))
|
||||
require.Equal(t, "grok-user@example.com", req.Header.Get("X-Email"))
|
||||
require.NotContains(t, req.Header.Get("Authorization"), "oauth-refresh-token")
|
||||
}
|
||||
|
||||
func TestBuildUpstreamModelsRequestGrokOAuthRequiresTokenProvider(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
svc := &AccountTestService{cfg: upstreamModelSyncTestConfig()}
|
||||
_, err := svc.buildUpstreamModelsRequest(context.Background(), &Account{
|
||||
Platform: PlatformGrok,
|
||||
Type: AccountTypeOAuth,
|
||||
})
|
||||
_, err := svc.buildUpstreamModelsRequest(context.Background(), grokOAuthModelSyncTestAccount(""))
|
||||
require.Error(t, err)
|
||||
|
||||
var syncErr *UpstreamModelSyncError
|
||||
require.True(t, errors.As(err, &syncErr))
|
||||
require.Equal(t, UpstreamModelSyncErrorUnsupported, syncErr.Kind)
|
||||
require.Contains(t, syncErr.SafeMessage(), "Unsupported Grok account type")
|
||||
require.Equal(t, UpstreamModelSyncErrorConfiguration, syncErr.Kind)
|
||||
require.Contains(t, syncErr.SafeMessage(), "token provider")
|
||||
}
|
||||
|
||||
func TestBuildAntigravityAPIKeyModelsRequestRejectsOfficialCloudCodeBase(t *testing.T) {
|
||||
@@ -321,6 +370,45 @@ func TestFetchUpstreamSupportedModelsParsesGrokAPIKeyResponse(t *testing.T) {
|
||||
require.Equal(t, "Bearer xai-key", upstream.lastReq.Header.Get("Authorization"))
|
||||
}
|
||||
|
||||
func TestFetchUpstreamSupportedModelsParsesGrokOAuthResponse(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
upstream := &httpUpstreamRecorder{resp: &http.Response{
|
||||
StatusCode: http.StatusOK,
|
||||
Header: http.Header{"Content-Type": []string{"application/json"}},
|
||||
Body: io.NopCloser(strings.NewReader(`{"data":[{"model":"grok-4.5"},{"model":"grok-4.5"},{"modelId":"grok-build-0.1"}]}`)),
|
||||
}}
|
||||
svc := &AccountTestService{
|
||||
httpUpstream: upstream,
|
||||
cfg: upstreamModelSyncTestConfig(),
|
||||
grokTokenProvider: NewGrokTokenProvider(nil, nil),
|
||||
}
|
||||
|
||||
models, err := svc.FetchUpstreamSupportedModels(context.Background(), grokOAuthModelSyncTestAccount(""))
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, []string{"grok-4.5", "grok-build-0.1"}, models)
|
||||
require.Equal(t, "https://cli-chat-proxy.grok.com/v1/models", upstream.lastReq.URL.String())
|
||||
require.Equal(t, "Bearer oauth-access-token", upstream.lastReq.Header.Get("Authorization"))
|
||||
require.Equal(t, grokCLIVersion, upstream.lastReq.Header.Get("X-Grok-Client-Version"))
|
||||
require.Equal(t, "interactive", upstream.lastReq.Header.Get("X-Grok-Client-Mode"))
|
||||
require.Equal(t, "grok-user-id", upstream.lastReq.Header.Get("X-UserID"))
|
||||
require.Equal(t, "grok-user@example.com", upstream.lastReq.Header.Get("X-Email"))
|
||||
}
|
||||
|
||||
func TestBuildUpstreamModelsRequestGrokOAuthDoesNotSendIdentityToCustomBase(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
svc := &AccountTestService{
|
||||
cfg: upstreamModelSyncTestConfig(),
|
||||
grokTokenProvider: NewGrokTokenProvider(nil, nil),
|
||||
}
|
||||
req, err := svc.buildUpstreamModelsRequest(context.Background(), grokOAuthModelSyncTestAccount("https://relay.example/v1"))
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, "https://relay.example/v1/models", req.URL.String())
|
||||
require.Empty(t, req.Header.Get("X-UserID"))
|
||||
require.Empty(t, req.Header.Get("X-Email"))
|
||||
}
|
||||
|
||||
func TestFetchUpstreamSupportedModelsDoesNotExposeUpstreamBody(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
|
||||
Reference in New Issue
Block a user