mirror of
https://github.com/Wei-Shaw/sub2api.git
synced 2026-09-24 16:05:44 +08:00
feat(grok): enable free OAuth prompt caching
This commit is contained in:
@@ -710,7 +710,7 @@ func (s *AccountTestService) testGrokAccountConnection(c *gin.Context, account *
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
req.Header.Set("Accept", "application/json, text/event-stream")
|
||||
req.Header.Set("Authorization", "Bearer "+authToken)
|
||||
req.Header.Set("User-Agent", "sub2api-grok/1.0")
|
||||
applyGrokCLIHeaders(req.Header)
|
||||
|
||||
proxyURL := ""
|
||||
if account.ProxyID != nil && account.Proxy != nil {
|
||||
|
||||
@@ -60,6 +60,7 @@ func TestAccountTestService_TestAccountConnection_GrokUsesXAIResponses(t *testin
|
||||
|
||||
require.Equal(t, "https://api.x.ai/v1/responses", upstream.lastReq.URL.String())
|
||||
require.Equal(t, "Bearer grok-access-token", upstream.lastReq.Header.Get("Authorization"))
|
||||
require.Equal(t, grokCLIVersion, upstream.lastReq.Header.Get("X-Grok-Client-Version"))
|
||||
require.Equal(t, "grok-4.3", gjson.GetBytes(upstream.lastBody, "model").String())
|
||||
require.NotContains(t, rec.Body.String(), "claude")
|
||||
require.Contains(t, rec.Body.String(), `"model":"grok-4.3"`)
|
||||
|
||||
@@ -333,7 +333,7 @@ func (s *OpenAIGatewayService) ForwardGrokMedia(
|
||||
}
|
||||
upstreamReq.Header.Set("Authorization", "Bearer "+token)
|
||||
upstreamReq.Header.Set("Accept", "application/json")
|
||||
upstreamReq.Header.Set("User-Agent", "sub2api-grok/1.0")
|
||||
applyGrokCLIHeaders(upstreamReq.Header)
|
||||
if endpoint.RequiresRequestBody() {
|
||||
contentType = strings.TrimSpace(contentType)
|
||||
if contentType == "" {
|
||||
|
||||
@@ -82,7 +82,7 @@ func (s *GrokQuotaService) ProbeUsage(ctx context.Context, accountID int64) (*Gr
|
||||
req.Header.Set("Authorization", "Bearer "+token)
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
req.Header.Set("Accept", "application/json")
|
||||
req.Header.Set("User-Agent", "sub2api-grok-quota-probe/1.0")
|
||||
applyGrokCLIHeaders(req.Header)
|
||||
|
||||
resp, err := s.httpUpstream.Do(req, proxyURL, account.ID, maxInt(account.Concurrency, 1))
|
||||
if err != nil {
|
||||
|
||||
@@ -98,6 +98,7 @@ func TestGrokQuotaServiceProbeUsageStoresHeaders(t *testing.T) {
|
||||
require.EqualValues(t, 7, *result.Snapshot.Requests.Remaining)
|
||||
require.Equal(t, "https://api.x.ai/v1/responses", upstream.lastReq.URL.String())
|
||||
require.Equal(t, "Bearer access-token", upstream.lastReq.Header.Get("Authorization"))
|
||||
require.Equal(t, grokCLIVersion, upstream.lastReq.Header.Get("X-Grok-Client-Version"))
|
||||
require.Equal(t, "grok-4.3", gjson.GetBytes(upstream.lastBody, "model").String())
|
||||
require.Contains(t, string(upstream.lastBody), `"max_output_tokens":1`)
|
||||
require.Contains(t, string(upstream.lastBody), `"store":false`)
|
||||
|
||||
@@ -159,6 +159,7 @@ func (s *OpenAIGatewayService) sendCCUpstreamRequest(
|
||||
stream bool,
|
||||
bearerToken string,
|
||||
userAgent string,
|
||||
grokCacheIdentity string,
|
||||
) (*http.Response, error) {
|
||||
upstreamCtx, releaseUpstreamCtx := detachUpstreamContext(ctx)
|
||||
upstreamReq, err := http.NewRequestWithContext(upstreamCtx, http.MethodPost, targetURL, bytes.NewReader(body))
|
||||
@@ -190,6 +191,10 @@ func (s *OpenAIGatewayService) sendCCUpstreamRequest(
|
||||
|
||||
// 账号级请求头覆写(仅 openai api_key 账号启用时生效)
|
||||
account.ApplyHeaderOverrides(upstreamReq.Header)
|
||||
if account.Platform == PlatformGrok {
|
||||
applyGrokCLIHeaders(upstreamReq.Header)
|
||||
applyGrokCacheHeaders(upstreamReq.Header, grokCacheIdentity)
|
||||
}
|
||||
|
||||
proxyURL := ""
|
||||
if account.Proxy != nil {
|
||||
|
||||
@@ -76,6 +76,12 @@ func (s *OpenAIGatewayService) forwardAsRawChatCompletions(
|
||||
// 2. Resolve model mapping (same as ForwardAsChatCompletions)
|
||||
billingModel := resolveOpenAIForwardModel(account, originalModel, defaultMappedModel)
|
||||
upstreamModel := normalizeOpenAIModelForUpstream(account, billingModel)
|
||||
grokCacheIdentity := ""
|
||||
if account.Platform == PlatformGrok {
|
||||
// Resolve before image bridging or other body rewrites so the fallback is
|
||||
// anchored to the client's stable conversation prefix.
|
||||
grokCacheIdentity = resolveGrokCacheIdentity(c, body, "", upstreamModel)
|
||||
}
|
||||
reasoningEffort := extractOpenAIReasoningEffortFromBody(body, upstreamModel, billingModel, originalModel)
|
||||
// 国产模型默认 effort 补充:需要 mappedModel 判定,推迟到 billingModel 算出之后。
|
||||
reasoningEffort = ApplyThinkingEnabledFallback(reasoningEffort, body, billingModel)
|
||||
@@ -134,6 +140,12 @@ func (s *OpenAIGatewayService) forwardAsRawChatCompletions(
|
||||
return nil, fmt.Errorf("enable stream usage: %w", usageErr)
|
||||
}
|
||||
}
|
||||
if account.Platform == PlatformGrok {
|
||||
upstreamBody, err = stripGrokChatPromptCacheKey(upstreamBody)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("remove Responses-only Grok prompt cache key: %w", err)
|
||||
}
|
||||
}
|
||||
|
||||
logger.L().Debug("openai chat_completions raw: forwarding without protocol conversion",
|
||||
zap.Int64("account_id", account.ID),
|
||||
@@ -152,7 +164,7 @@ func (s *OpenAIGatewayService) forwardAsRawChatCompletions(
|
||||
if customUA == "" && account.Platform == PlatformGrok {
|
||||
customUA = "sub2api-grok/1.0"
|
||||
}
|
||||
resp, err := s.sendCCUpstreamRequest(ctx, c, account, targetURL, upstreamBody, clientStream, token, customUA)
|
||||
resp, err := s.sendCCUpstreamRequest(ctx, c, account, targetURL, upstreamBody, clientStream, token, customUA, grokCacheIdentity)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
@@ -49,7 +49,6 @@ func (s *OpenAIGatewayService) Forward(ctx context.Context, c *gin.Context, acco
|
||||
originalModel := reqModel
|
||||
|
||||
if account.Platform == PlatformGrok {
|
||||
_ = promptCacheKey
|
||||
return s.forwardGrokResponses(ctx, c, account, body, originalModel, reqStream, startTime)
|
||||
}
|
||||
|
||||
|
||||
@@ -20,6 +20,8 @@ import (
|
||||
const (
|
||||
grokComposerImageBridgeVisionModel = "grok-build-0.1"
|
||||
grokComposerImageBridgeMaxOutputTokens = 512
|
||||
grokUpstreamUserAgent = "sub2api-grok/1.0"
|
||||
grokCLIVersion = "0.2.93"
|
||||
)
|
||||
|
||||
func (s *OpenAIGatewayService) forwardGrokResponses(
|
||||
@@ -39,10 +41,15 @@ func (s *OpenAIGatewayService) forwardGrokResponses(
|
||||
if strings.TrimSpace(upstreamModel) == "" {
|
||||
upstreamModel = "grok-4.3"
|
||||
}
|
||||
cacheIdentity := resolveGrokCacheIdentity(c, body, "", upstreamModel)
|
||||
patchedBody, err := patchGrokResponsesBody(body, upstreamModel)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
patchedBody, err = applyGrokResponsesCacheIdentity(patchedBody, body, cacheIdentity, account.IsGrokOAuth())
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("apply grok prompt cache identity: %w", err)
|
||||
}
|
||||
|
||||
token, _, err := s.GetAccessToken(ctx, account)
|
||||
if err != nil {
|
||||
@@ -51,7 +58,7 @@ func (s *OpenAIGatewayService) forwardGrokResponses(
|
||||
|
||||
upstreamCtx, releaseUpstreamCtx := detachUpstreamContext(ctx)
|
||||
defer releaseUpstreamCtx()
|
||||
upstreamReq, err := buildGrokResponsesRequest(upstreamCtx, c, account, patchedBody, token)
|
||||
upstreamReq, err := buildGrokResponsesRequest(upstreamCtx, c, account, patchedBody, token, cacheIdentity)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
@@ -457,7 +464,9 @@ func (s *OpenAIGatewayService) describeGrokComposerImage(
|
||||
}
|
||||
|
||||
upstreamCtx, releaseUpstreamCtx := detachUpstreamContext(ctx)
|
||||
upstreamReq, err := buildGrokResponsesRequest(upstreamCtx, c, account, body, token)
|
||||
// Image-description probes are auxiliary requests, not conversation turns.
|
||||
// Do not bind them to the caller's Grok prompt-cache identity.
|
||||
upstreamReq, err := buildGrokResponsesRequest(upstreamCtx, c, account, body, token, "")
|
||||
releaseUpstreamCtx()
|
||||
if err != nil {
|
||||
return "", OpenAIUsage{}, fmt.Errorf("build grok composer image bridge request: %w", err)
|
||||
@@ -623,7 +632,7 @@ func addOpenAIUsage(dst *OpenAIUsage, usage OpenAIUsage) {
|
||||
dst.ImageOutputTokens += usage.ImageOutputTokens
|
||||
}
|
||||
|
||||
func buildGrokResponsesRequest(ctx context.Context, c *gin.Context, account *Account, body []byte, token string) (*http.Request, error) {
|
||||
func buildGrokResponsesRequest(ctx context.Context, c *gin.Context, account *Account, body []byte, token, cacheIdentity string) (*http.Request, error) {
|
||||
targetURL, err := xai.BuildResponsesURL(account.GetGrokBaseURL())
|
||||
if err != nil {
|
||||
return nil, err
|
||||
@@ -635,7 +644,8 @@ func buildGrokResponsesRequest(ctx context.Context, c *gin.Context, account *Acc
|
||||
req.Header.Set("Authorization", "Bearer "+token)
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
req.Header.Set("Accept", "application/json, text/event-stream")
|
||||
req.Header.Set("User-Agent", "sub2api-grok/1.0")
|
||||
applyGrokCLIHeaders(req.Header)
|
||||
applyGrokCacheHeaders(req.Header, cacheIdentity)
|
||||
if c != nil {
|
||||
if v := c.GetHeader("OpenAI-Beta"); strings.TrimSpace(v) != "" {
|
||||
req.Header.Set("OpenAI-Beta", v)
|
||||
@@ -644,6 +654,16 @@ func buildGrokResponsesRequest(ctx context.Context, c *gin.Context, account *Acc
|
||||
return req, nil
|
||||
}
|
||||
|
||||
// applyGrokCLIHeaders identifies subscription traffic as a supported Grok CLI
|
||||
// version. The CLI gateway rejects otherwise valid OAuth requests without it.
|
||||
func applyGrokCLIHeaders(headers http.Header) {
|
||||
if headers == nil {
|
||||
return
|
||||
}
|
||||
headers.Set("User-Agent", grokUpstreamUserAgent)
|
||||
headers.Set("X-Grok-Client-Version", grokCLIVersion)
|
||||
}
|
||||
|
||||
func (s *OpenAIGatewayService) updateGrokUsageSnapshot(ctx context.Context, accountID int64, snapshot *xai.QuotaSnapshot) {
|
||||
if s == nil || s.accountRepo == nil || accountID <= 0 || snapshot == nil {
|
||||
return
|
||||
|
||||
@@ -0,0 +1,149 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"net/http"
|
||||
"strings"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/tidwall/gjson"
|
||||
"github.com/tidwall/sjson"
|
||||
)
|
||||
|
||||
const (
|
||||
grokConversationIDHeader = "X-Grok-Conv-Id"
|
||||
grokFreeCacheNativeToolsJSON = `[{"type":"web_search"},{"type":"x_search"}]`
|
||||
grokFreeCacheDisabledToolChoice = "none"
|
||||
)
|
||||
|
||||
// resolveGrokCacheIdentity derives one stable, tenant-isolated routing identity
|
||||
// for xAI's server-side prompt cache. The returned value is safe to expose to
|
||||
// the upstream: it never contains the client's raw session identifier.
|
||||
//
|
||||
// A valid downstream API key is required. This intentionally fails closed on
|
||||
// internal probes and incomplete request contexts instead of creating a cache
|
||||
// identity that could be shared by unrelated tenants.
|
||||
func resolveGrokCacheIdentity(c *gin.Context, body []byte, explicitKey, upstreamModel string) string {
|
||||
apiKeyID := getAPIKeyIDFromContext(c)
|
||||
if apiKeyID <= 0 {
|
||||
return ""
|
||||
}
|
||||
// /responses/compact rejects tool_choice and does not represent a normal
|
||||
// conversation turn. Keep both cache identity and Free-tier routing
|
||||
// augmentation out of this path.
|
||||
if isOpenAIResponsesCompactPath(c) {
|
||||
return ""
|
||||
}
|
||||
|
||||
model := strings.ToLower(strings.TrimSpace(upstreamModel))
|
||||
if model == "" {
|
||||
return ""
|
||||
}
|
||||
|
||||
seed := explicitGrokCacheSeed(c, body, explicitKey)
|
||||
if seed == "" {
|
||||
seed = deriveOpenAIContentSessionSeed(body)
|
||||
}
|
||||
if seed == "" {
|
||||
return ""
|
||||
}
|
||||
|
||||
// generateSessionUUID hashes the whole seed before formatting it as a UUID.
|
||||
// Include a versioned namespace so this identity cannot collide with other
|
||||
// upstream session identifiers derived by sub2api.
|
||||
isolatedSeed := fmt.Sprintf("grok-prompt-cache:v1:%d:%s:%s", apiKeyID, model, seed)
|
||||
return generateSessionUUID(isolatedSeed)
|
||||
}
|
||||
|
||||
func explicitGrokCacheSeed(c *gin.Context, body []byte, explicitKey string) string {
|
||||
seed := ""
|
||||
if c != nil {
|
||||
seed = strings.TrimSpace(c.GetHeader("session_id"))
|
||||
if seed == "" {
|
||||
seed = strings.TrimSpace(c.GetHeader("conversation_id"))
|
||||
}
|
||||
if seed == "" {
|
||||
seed = strings.TrimSpace(c.GetHeader(grokConversationIDHeader))
|
||||
}
|
||||
}
|
||||
if seed == "" && len(body) > 0 {
|
||||
seed = strings.TrimSpace(gjson.GetBytes(body, "prompt_cache_key").String())
|
||||
}
|
||||
if seed == "" {
|
||||
seed = strings.TrimSpace(explicitKey)
|
||||
}
|
||||
return seed
|
||||
}
|
||||
|
||||
func isGrokRequestContext(c *gin.Context) bool {
|
||||
if c == nil {
|
||||
return false
|
||||
}
|
||||
v, exists := c.Get("api_key")
|
||||
if !exists {
|
||||
return false
|
||||
}
|
||||
apiKey, ok := v.(*APIKey)
|
||||
return ok && apiKey != nil && apiKey.Group != nil && apiKey.Group.Platform == PlatformGrok
|
||||
}
|
||||
|
||||
// applyGrokResponsesCacheIdentity writes the cache routing identity into an
|
||||
// xAI Responses request. Existing client values are deliberately replaced by
|
||||
// the tenant-isolated value to prevent collisions on shared OAuth accounts.
|
||||
//
|
||||
// Free OAuth requests without native search tools are routed by xAI to the
|
||||
// non-cacheable build-free model. For otherwise tool-free requests, add the
|
||||
// native tools with tool_choice=none: this selects the cache-capable tier
|
||||
// without allowing an actual search. Any explicit client tools or tool_choice
|
||||
// disable this augmentation so client function-calling semantics stay intact.
|
||||
func applyGrokResponsesCacheIdentity(body, intentSourceBody []byte, identity string, injectFreeTierTools bool) ([]byte, error) {
|
||||
identity = strings.TrimSpace(identity)
|
||||
if identity == "" {
|
||||
if gjson.GetBytes(body, "prompt_cache_key").Exists() {
|
||||
return sjson.DeleteBytes(body, "prompt_cache_key")
|
||||
}
|
||||
return body, nil
|
||||
}
|
||||
out, err := sjson.SetBytes(body, "prompt_cache_key", identity)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if !injectFreeTierTools {
|
||||
return out, nil
|
||||
}
|
||||
// Inspect the pre-sanitization source. patchGrokResponsesBody may remove an
|
||||
// unsupported client tool and its tool_choice; that must not turn an
|
||||
// explicit client tool intent into an eligible native-tool request.
|
||||
if gjson.GetBytes(intentSourceBody, "tools").Exists() || gjson.GetBytes(intentSourceBody, "tool_choice").Exists() {
|
||||
return out, nil
|
||||
}
|
||||
out, err = sjson.SetRawBytes(out, "tools", []byte(grokFreeCacheNativeToolsJSON))
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return sjson.SetBytes(out, "tool_choice", grokFreeCacheDisabledToolChoice)
|
||||
}
|
||||
|
||||
// applyGrokCacheHeaders applies the documented Chat Completions conversation
|
||||
// routing header. The request is built from a fresh header map, so client
|
||||
// supplied x-grok headers cannot override this server-derived value.
|
||||
func applyGrokCacheHeaders(headers http.Header, identity string) {
|
||||
if headers == nil {
|
||||
return
|
||||
}
|
||||
identity = strings.TrimSpace(identity)
|
||||
if identity == "" {
|
||||
headers.Del(grokConversationIDHeader)
|
||||
return
|
||||
}
|
||||
headers.Set(grokConversationIDHeader, identity)
|
||||
}
|
||||
|
||||
// stripGrokChatPromptCacheKey removes the Responses-only body field after it
|
||||
// has been used as an identity seed. Chat Completions routes cache by header.
|
||||
func stripGrokChatPromptCacheKey(body []byte) ([]byte, error) {
|
||||
if !gjson.GetBytes(body, "prompt_cache_key").Exists() {
|
||||
return body, nil
|
||||
}
|
||||
return sjson.DeleteBytes(body, "prompt_cache_key")
|
||||
}
|
||||
@@ -0,0 +1,273 @@
|
||||
//go:build unit
|
||||
|
||||
package service
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"sync"
|
||||
"testing"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/stretchr/testify/require"
|
||||
"github.com/tidwall/gjson"
|
||||
)
|
||||
|
||||
func newGrokCacheTestContext(apiKeyID int64) *gin.Context {
|
||||
recorder := httptest.NewRecorder()
|
||||
c, _ := gin.CreateTestContext(recorder)
|
||||
c.Request = httptest.NewRequest(http.MethodPost, "/v1/responses", nil)
|
||||
if apiKeyID > 0 {
|
||||
c.Set("api_key", &APIKey{ID: apiKeyID, Group: &Group{Platform: PlatformGrok}})
|
||||
}
|
||||
return c
|
||||
}
|
||||
|
||||
func TestResolveGrokCacheIdentityStableAcrossAppendOnlyTurns(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
c := newGrokCacheTestContext(101)
|
||||
round1 := []byte(`{"model":"grok","instructions":"be concise","tools":[{"type":"function","name":"lookup","parameters":{"type":"object"}}],"input":[{"role":"user","content":"first question"}]}`)
|
||||
round2 := []byte(`{"model":"grok","instructions":"be concise","tools":[{"type":"function","name":"lookup","parameters":{"type":"object"}}],"input":[{"role":"user","content":"first question"},{"role":"assistant","content":"first answer"},{"role":"user","content":"second question"}]}`)
|
||||
|
||||
first := resolveGrokCacheIdentity(c, round1, "", "grok-4.5")
|
||||
second := resolveGrokCacheIdentity(c, round2, "", "grok-4.5")
|
||||
|
||||
require.NotEmpty(t, first)
|
||||
require.Len(t, first, 36)
|
||||
require.Equal(t, first, second)
|
||||
}
|
||||
|
||||
func TestResolveGrokCacheIdentityIsolatesAPIKeyAndMappedModel(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
body := []byte(`{"model":"grok","input":"same prompt"}`)
|
||||
|
||||
base := resolveGrokCacheIdentity(newGrokCacheTestContext(201), body, "", "grok-4.5")
|
||||
otherTenant := resolveGrokCacheIdentity(newGrokCacheTestContext(202), body, "", "grok-4.5")
|
||||
otherModel := resolveGrokCacheIdentity(newGrokCacheTestContext(201), body, "", "grok-4.3")
|
||||
|
||||
require.NotEmpty(t, base)
|
||||
require.NotEqual(t, base, otherTenant)
|
||||
require.NotEqual(t, base, otherModel)
|
||||
}
|
||||
|
||||
func TestResolveGrokCacheIdentityUsesAndIsolatesNativeConversationHeader(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
c := newGrokCacheTestContext(301)
|
||||
c.Request.Header.Set(grokConversationIDHeader, "raw-native-conversation")
|
||||
body1 := []byte(`{"model":"grok","input":"one"}`)
|
||||
body2 := []byte(`{"model":"grok","input":"different body that must not replace the explicit session"}`)
|
||||
|
||||
first := resolveGrokCacheIdentity(c, body1, "body-cache-key", "grok-4.5")
|
||||
second := resolveGrokCacheIdentity(c, body2, "another-body-cache-key", "grok-4.5")
|
||||
|
||||
require.Equal(t, "raw-native-conversation", (&OpenAIGatewayService{}).ExtractSessionID(c, body1))
|
||||
require.Equal(t, first, second)
|
||||
require.NotEqual(t, "raw-native-conversation", first)
|
||||
require.NotContains(t, first, "raw-native-conversation")
|
||||
}
|
||||
|
||||
func TestResolveGrokCacheIdentityExplicitHeaderPriority(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
body := []byte(`{"model":"grok","prompt_cache_key":"body-key","input":"hi"}`)
|
||||
c := newGrokCacheTestContext(401)
|
||||
c.Request.Header.Set(grokConversationIDHeader, "grok-key")
|
||||
c.Request.Header.Set("conversation_id", "conversation-key")
|
||||
c.Request.Header.Set("session_id", "session-key")
|
||||
|
||||
got := resolveGrokCacheIdentity(c, body, "explicit-argument", "grok-4.5")
|
||||
onlySession := newGrokCacheTestContext(401)
|
||||
onlySession.Request.Header.Set("session_id", "session-key")
|
||||
want := resolveGrokCacheIdentity(onlySession, []byte(`{"model":"grok","input":"unrelated"}`), "", "grok-4.5")
|
||||
|
||||
require.Equal(t, want, got)
|
||||
}
|
||||
|
||||
func TestResolveGrokCacheIdentityFailsClosedWithoutAPIKeyContext(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
c := newGrokCacheTestContext(0)
|
||||
c.Request.Header.Set(grokConversationIDHeader, "native-session")
|
||||
|
||||
require.Empty(t, resolveGrokCacheIdentity(c, []byte(`{"model":"grok","input":"hi"}`), "", "grok-4.5"))
|
||||
require.Empty(t, resolveGrokCacheIdentity(nil, []byte(`{"model":"grok","prompt_cache_key":"key"}`), "key", "grok-4.5"))
|
||||
}
|
||||
|
||||
func TestGrokConversationHeaderIsScopedToGrokRequestScheduling(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
body := []byte(`{"model":"grok","prompt_cache_key":"body-session","input":"hi"}`)
|
||||
|
||||
grokContext := newGrokCacheTestContext(601)
|
||||
grokContext.Request.Header.Set(grokConversationIDHeader, "native-grok-session")
|
||||
require.Equal(t, "native-grok-session", (&OpenAIGatewayService{}).ExtractSessionID(grokContext, body))
|
||||
|
||||
openAIContext := newGrokCacheTestContext(601)
|
||||
openAIContext.Set("api_key", &APIKey{ID: 601, Group: &Group{Platform: PlatformOpenAI}})
|
||||
openAIContext.Request.Header.Set(grokConversationIDHeader, "must-be-ignored")
|
||||
require.Equal(t, "body-session", (&OpenAIGatewayService{}).ExtractSessionID(openAIContext, body))
|
||||
|
||||
withoutGrokHeader := newGrokCacheTestContext(601)
|
||||
withoutGrokHeader.Set("api_key", &APIKey{ID: 601, Group: &Group{Platform: PlatformOpenAI}})
|
||||
require.Equal(t,
|
||||
(&OpenAIGatewayService{}).GenerateSessionHash(withoutGrokHeader, body),
|
||||
(&OpenAIGatewayService{}).GenerateSessionHash(openAIContext, body),
|
||||
)
|
||||
}
|
||||
|
||||
func TestApplyGrokCacheIdentityWritesResponsesBodyAndHeader(t *testing.T) {
|
||||
sourceBody := []byte(`{"model":"grok-4.5","prompt_cache_key":"raw-client-key"}`)
|
||||
body, err := applyGrokResponsesCacheIdentity(sourceBody, sourceBody, "isolated-id", true)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, "isolated-id", gjson.GetBytes(body, "prompt_cache_key").String())
|
||||
require.Equal(t, "web_search", gjson.GetBytes(body, "tools.0.type").String())
|
||||
require.Equal(t, "x_search", gjson.GetBytes(body, "tools.1.type").String())
|
||||
require.Equal(t, grokFreeCacheDisabledToolChoice, gjson.GetBytes(body, "tool_choice").String())
|
||||
|
||||
headers := make(http.Header)
|
||||
headers.Set(grokConversationIDHeader, "spoofed-client-value")
|
||||
applyGrokCacheHeaders(headers, "isolated-id")
|
||||
require.Equal(t, "isolated-id", headers.Get(grokConversationIDHeader))
|
||||
applyGrokCacheHeaders(headers, "")
|
||||
require.Empty(t, headers.Get(grokConversationIDHeader))
|
||||
|
||||
chatBody, err := stripGrokChatPromptCacheKey(body)
|
||||
require.NoError(t, err)
|
||||
require.False(t, gjson.GetBytes(chatBody, "prompt_cache_key").Exists())
|
||||
|
||||
unscopedSourceBody := []byte(`{"model":"grok","prompt_cache_key":"raw-client-key"}`)
|
||||
unscopedBody, err := applyGrokResponsesCacheIdentity(unscopedSourceBody, unscopedSourceBody, "", true)
|
||||
require.NoError(t, err)
|
||||
require.False(t, gjson.GetBytes(unscopedBody, "prompt_cache_key").Exists())
|
||||
require.False(t, gjson.GetBytes(unscopedBody, "tools").Exists())
|
||||
require.False(t, gjson.GetBytes(unscopedBody, "tool_choice").Exists())
|
||||
}
|
||||
|
||||
func TestApplyGrokCacheIdentityPreservesExplicitClientToolFields(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
body string
|
||||
}{
|
||||
{
|
||||
name: "tools only",
|
||||
body: `{"model":"grok","tools":[{"type":"function","name":"lookup","parameters":{"type":"object"}}]}`,
|
||||
},
|
||||
{
|
||||
name: "empty tools array",
|
||||
body: `{"model":"grok","tools":[]}`,
|
||||
},
|
||||
{
|
||||
name: "null tools",
|
||||
body: `{"model":"grok","tools":null}`,
|
||||
},
|
||||
{
|
||||
name: "tool choice only",
|
||||
body: `{"model":"grok","tool_choice":{"type":"function","name":"lookup"}}`,
|
||||
},
|
||||
{
|
||||
name: "null tool choice",
|
||||
body: `{"model":"grok","tool_choice":null}`,
|
||||
},
|
||||
{
|
||||
name: "both fields",
|
||||
body: `{"model":"grok","tools":[{"type":"web_search"}],"tool_choice":"auto"}`,
|
||||
},
|
||||
{
|
||||
name: "unsupported tool",
|
||||
body: `{"model":"grok","tools":[{"type":"namespace","name":"client_tools"}]}`,
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
beforeTools := gjson.Get(tt.body, "tools")
|
||||
beforeChoice := gjson.Get(tt.body, "tool_choice")
|
||||
body, err := applyGrokResponsesCacheIdentity([]byte(tt.body), []byte(tt.body), "isolated-id", true)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, "isolated-id", gjson.GetBytes(body, "prompt_cache_key").String())
|
||||
require.Equal(t, beforeTools.Exists(), gjson.GetBytes(body, "tools").Exists())
|
||||
require.Equal(t, beforeTools.Raw, gjson.GetBytes(body, "tools").Raw)
|
||||
require.Equal(t, beforeChoice.Exists(), gjson.GetBytes(body, "tool_choice").Exists())
|
||||
require.Equal(t, beforeChoice.Raw, gjson.GetBytes(body, "tool_choice").Raw)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestApplyGrokCacheIdentityUsesPreSanitizationToolIntent(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
intentBody string
|
||||
}{
|
||||
{
|
||||
name: "unsupported tools removed by sanitizer",
|
||||
intentBody: `{"model":"grok","tools":[{"type":"namespace","name":"client_tools"}]}`,
|
||||
},
|
||||
{
|
||||
name: "tool choice removed with unsupported tool",
|
||||
intentBody: `{"model":"grok","tools":[{"type":"namespace","name":"client_tools"}],"tool_choice":{"type":"namespace","name":"client_tools"}}`,
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
// This is the shape apply receives after patchGrokResponsesBody has
|
||||
// removed unsupported tools and their associated tool_choice.
|
||||
patchedBody := []byte(`{"model":"grok-4.5","input":"hello"}`)
|
||||
body, err := applyGrokResponsesCacheIdentity(patchedBody, []byte(tt.intentBody), "isolated-id", true)
|
||||
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, "isolated-id", gjson.GetBytes(body, "prompt_cache_key").String())
|
||||
require.False(t, gjson.GetBytes(body, "tools").Exists())
|
||||
require.False(t, gjson.GetBytes(body, "tool_choice").Exists())
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestApplyGrokCacheIdentityWithoutFreeTierRoutingOnlyWritesIdentity(t *testing.T) {
|
||||
sourceBody := []byte(`{"model":"grok-4.5","input":"hello"}`)
|
||||
body, err := applyGrokResponsesCacheIdentity(sourceBody, sourceBody, "isolated-id", false)
|
||||
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, "isolated-id", gjson.GetBytes(body, "prompt_cache_key").String())
|
||||
require.False(t, gjson.GetBytes(body, "tools").Exists())
|
||||
require.False(t, gjson.GetBytes(body, "tool_choice").Exists())
|
||||
}
|
||||
|
||||
func TestGrokCompactRequestSkipsCacheIdentityAndNativeTools(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
c := newGrokCacheTestContext(701)
|
||||
c.Request = httptest.NewRequest(http.MethodPost, "/v1/responses/compact", nil)
|
||||
body := []byte(`{"model":"grok","input":"compact this","prompt_cache_key":"raw-client-key"}`)
|
||||
|
||||
identity := resolveGrokCacheIdentity(c, body, "", "grok-4.5")
|
||||
patched, err := applyGrokResponsesCacheIdentity(body, body, identity, true)
|
||||
|
||||
require.NoError(t, err)
|
||||
require.Empty(t, identity)
|
||||
require.False(t, gjson.GetBytes(patched, "prompt_cache_key").Exists())
|
||||
require.False(t, gjson.GetBytes(patched, "tools").Exists())
|
||||
require.False(t, gjson.GetBytes(patched, "tool_choice").Exists())
|
||||
}
|
||||
|
||||
func TestResolveGrokCacheIdentityConcurrentDeterminism(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
const workers = 50
|
||||
body := []byte(`{"model":"grok","messages":[{"role":"system","content":"stable"},{"role":"user","content":"hello"}]}`)
|
||||
identities := make(chan string, workers)
|
||||
var wg sync.WaitGroup
|
||||
for i := 0; i < workers; i++ {
|
||||
wg.Add(1)
|
||||
go func() {
|
||||
defer wg.Done()
|
||||
identities <- resolveGrokCacheIdentity(newGrokCacheTestContext(501), body, "", "grok-4.5")
|
||||
}()
|
||||
}
|
||||
wg.Wait()
|
||||
close(identities)
|
||||
|
||||
var first string
|
||||
for identity := range identities {
|
||||
if first == "" {
|
||||
first = identity
|
||||
}
|
||||
require.Equal(t, first, identity)
|
||||
}
|
||||
require.NotEmpty(t, first)
|
||||
}
|
||||
@@ -6,6 +6,7 @@ import (
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"io"
|
||||
"mime/multipart"
|
||||
"net/http"
|
||||
@@ -172,13 +173,15 @@ func TestBuildGrokResponsesRequestUsesAccountBaseURLAndBearerToken(t *testing.T)
|
||||
},
|
||||
}
|
||||
|
||||
req, err := buildGrokResponsesRequest(context.Background(), nil, account, []byte(`{"model":"grok-4.3"}`), "access-token")
|
||||
req, err := buildGrokResponsesRequest(context.Background(), nil, account, []byte(`{"model":"grok-4.3"}`), "access-token", "isolated-cache-id")
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, http.MethodPost, req.Method)
|
||||
require.Equal(t, "https://xai.test/v1/responses", req.URL.String())
|
||||
require.Equal(t, "Bearer access-token", req.Header.Get("Authorization"))
|
||||
require.Equal(t, "application/json", req.Header.Get("Content-Type"))
|
||||
require.Contains(t, req.Header.Get("Accept"), "text/event-stream")
|
||||
require.Equal(t, grokCLIVersion, req.Header.Get("X-Grok-Client-Version"))
|
||||
require.Equal(t, "isolated-cache-id", req.Header.Get(grokConversationIDHeader))
|
||||
|
||||
data, err := io.ReadAll(req.Body)
|
||||
require.NoError(t, err)
|
||||
@@ -196,7 +199,7 @@ func TestBuildGrokResponsesRequestRejectsUnsafeAccountBaseURL(t *testing.T) {
|
||||
},
|
||||
}
|
||||
|
||||
_, err := buildGrokResponsesRequest(context.Background(), nil, account, []byte(`{"model":"grok-4.3"}`), "access-token")
|
||||
_, err := buildGrokResponsesRequest(context.Background(), nil, account, []byte(`{"model":"grok-4.3"}`), "access-token", "")
|
||||
require.Error(t, err)
|
||||
require.Contains(t, err.Error(), "invalid base url")
|
||||
}
|
||||
@@ -324,6 +327,7 @@ func TestForwardGrokMediaImagesGenerationNormalizesImagineAlias(t *testing.T) {
|
||||
require.Equal(t, http.MethodPost, upstream.lastReq.Method)
|
||||
require.Equal(t, "Bearer api-key", upstream.lastReq.Header.Get("Authorization"))
|
||||
require.Equal(t, "application/json", upstream.lastReq.Header.Get("Content-Type"))
|
||||
require.Equal(t, grokCLIVersion, upstream.lastReq.Header.Get("X-Grok-Client-Version"))
|
||||
require.JSONEq(t, `{"model":"grok-imagine-image-quality","prompt":"draw a cat"}`, string(upstream.lastBody))
|
||||
require.Equal(t, http.StatusOK, recorder.Code)
|
||||
require.JSONEq(t, `{"data":[]}`, recorder.Body.String())
|
||||
@@ -611,8 +615,9 @@ func TestForwardAsChatCompletionsForGrokUsesXAIChatCompletionsAndSnapshots(t *te
|
||||
|
||||
recorder := httptest.NewRecorder()
|
||||
c, _ := gin.CreateTestContext(recorder)
|
||||
body := []byte(`{"model":"grok","messages":[{"role":"user","content":"hi"}],"stream":false}`)
|
||||
body := []byte(`{"model":"grok","messages":[{"role":"user","content":"hi"}],"stream":false,"prompt_cache_key":"raw-client-cache-key"}`)
|
||||
c.Request = httptest.NewRequest(http.MethodPost, "/v1/chat/completions", bytes.NewReader(body))
|
||||
c.Set("api_key", &APIKey{ID: 5101})
|
||||
|
||||
account := &Account{
|
||||
ID: 51,
|
||||
@@ -641,7 +646,7 @@ func TestForwardAsChatCompletionsForGrokUsesXAIChatCompletionsAndSnapshots(t *te
|
||||
"X-Ratelimit-Limit-Tokens": []string{"1000"},
|
||||
"X-Ratelimit-Remaining-Tokens": []string{"990"},
|
||||
},
|
||||
Body: io.NopCloser(strings.NewReader(`{"id":"chatcmpl","object":"chat.completion","model":"grok-4.3","choices":[],"usage":{"prompt_tokens":1,"completion_tokens":2}}`)),
|
||||
Body: io.NopCloser(strings.NewReader(`{"id":"chatcmpl","object":"chat.completion","model":"grok-4.3","choices":[],"usage":{"prompt_tokens":1,"completion_tokens":2,"prompt_tokens_details":{"cached_tokens":1}}}`)),
|
||||
}}
|
||||
svc := &OpenAIGatewayService{
|
||||
httpUpstream: upstream,
|
||||
@@ -653,11 +658,15 @@ func TestForwardAsChatCompletionsForGrokUsesXAIChatCompletionsAndSnapshots(t *te
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, xai.DefaultCLIBaseURL+"/chat/completions", upstream.lastReq.URL.String())
|
||||
require.Equal(t, "Bearer access-token", upstream.lastReq.Header.Get("Authorization"))
|
||||
require.NotEmpty(t, upstream.lastReq.Header.Get(grokConversationIDHeader))
|
||||
require.NotEqual(t, "raw-client-cache-key", upstream.lastReq.Header.Get(grokConversationIDHeader))
|
||||
require.Equal(t, "grok-4.5", gjson.GetBytes(upstream.lastBody, "model").String())
|
||||
require.False(t, gjson.GetBytes(upstream.lastBody, "prompt_cache_key").Exists())
|
||||
require.Equal(t, "grok", result.Model)
|
||||
require.Equal(t, "grok-4.5", result.UpstreamModel)
|
||||
require.Equal(t, 1, result.Usage.InputTokens)
|
||||
require.Equal(t, 2, result.Usage.OutputTokens)
|
||||
require.Equal(t, 1, result.Usage.CacheReadInputTokens)
|
||||
require.NotNil(t, repo.updates[51][grokQuotaSnapshotExtraKey])
|
||||
require.Equal(t, http.StatusOK, recorder.Code)
|
||||
}
|
||||
@@ -671,6 +680,7 @@ func TestForwardGrokResponsesStreamingUsesXAIResponsesAndSnapshots(t *testing.T)
|
||||
c.Request = httptest.NewRequest(http.MethodPost, "/v1/responses", bytes.NewReader(body))
|
||||
c.Request.Header.Set("Content-Type", "application/json")
|
||||
c.Request.Header.Set("OpenAI-Beta", "responses=experimental")
|
||||
c.Set("api_key", &APIKey{ID: 5201})
|
||||
|
||||
account := &Account{
|
||||
ID: 52,
|
||||
@@ -719,6 +729,11 @@ func TestForwardGrokResponsesStreamingUsesXAIResponsesAndSnapshots(t *testing.T)
|
||||
require.Equal(t, "Bearer access-token", upstream.lastReq.Header.Get("Authorization"))
|
||||
require.Equal(t, "responses=experimental", upstream.lastReq.Header.Get("OpenAI-Beta"))
|
||||
require.Equal(t, "grok-4.5", gjson.GetBytes(upstream.lastBody, "model").String())
|
||||
require.NotEmpty(t, gjson.GetBytes(upstream.lastBody, "prompt_cache_key").String())
|
||||
require.Equal(t, gjson.GetBytes(upstream.lastBody, "prompt_cache_key").String(), upstream.lastReq.Header.Get(grokConversationIDHeader))
|
||||
require.Equal(t, "web_search", gjson.GetBytes(upstream.lastBody, "tools.0.type").String())
|
||||
require.Equal(t, "x_search", gjson.GetBytes(upstream.lastBody, "tools.1.type").String())
|
||||
require.Equal(t, "none", gjson.GetBytes(upstream.lastBody, "tool_choice").String())
|
||||
require.Equal(t, "high", gjson.GetBytes(upstream.lastBody, "reasoning_effort").String())
|
||||
require.True(t, gjson.GetBytes(upstream.lastBody, "stream").Bool())
|
||||
require.True(t, result.Stream)
|
||||
@@ -734,6 +749,133 @@ func TestForwardGrokResponsesStreamingUsesXAIResponsesAndSnapshots(t *testing.T)
|
||||
require.NotNil(t, repo.updates[52][grokQuotaSnapshotExtraKey])
|
||||
}
|
||||
|
||||
func TestForwardGrokResponsesNonStreamingUsesCacheIdentityAndCachedUsage(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
|
||||
recorder := httptest.NewRecorder()
|
||||
c, _ := gin.CreateTestContext(recorder)
|
||||
body := []byte(`{"model":"grok","input":"hi","stream":false,"tools":[{"type":"namespace","name":"client_tools"}],"tool_choice":{"type":"namespace","name":"client_tools"}}`)
|
||||
c.Request = httptest.NewRequest(http.MethodPost, "/v1/responses", bytes.NewReader(body))
|
||||
c.Request.Header.Set("Content-Type", "application/json")
|
||||
c.Set("api_key", &APIKey{ID: 5202})
|
||||
|
||||
account := &Account{
|
||||
ID: 56,
|
||||
Name: "grok",
|
||||
Platform: PlatformGrok,
|
||||
Type: AccountTypeOAuth,
|
||||
Concurrency: 1,
|
||||
Credentials: map[string]any{
|
||||
"access_token": "access-token",
|
||||
"expires_at": time.Now().Add(time.Hour).UTC().Format(time.RFC3339),
|
||||
"base_url": xai.DefaultCLIBaseURL,
|
||||
},
|
||||
}
|
||||
repo := &grokQuotaAccountRepo{
|
||||
mockAccountRepoForPlatform: &mockAccountRepoForPlatform{
|
||||
accountsByID: map[int64]*Account{56: account},
|
||||
},
|
||||
}
|
||||
upstream := &httpUpstreamRecorder{resp: &http.Response{
|
||||
StatusCode: http.StatusOK,
|
||||
Header: http.Header{
|
||||
"Content-Type": []string{"application/json"},
|
||||
"Xai-Request-Id": []string{"xai-non-stream-req"},
|
||||
},
|
||||
Body: io.NopCloser(strings.NewReader(`{"id":"resp_grok_non_stream","object":"response","model":"grok-4.3","status":"completed","output":[{"type":"message","role":"assistant","content":[{"type":"output_text","text":"ok"}]}],"usage":{"input_tokens":7,"output_tokens":2,"total_tokens":9,"input_tokens_details":{"cached_tokens":4}}}`)),
|
||||
}}
|
||||
svc := &OpenAIGatewayService{
|
||||
httpUpstream: upstream,
|
||||
grokTokenProvider: NewGrokTokenProvider(repo, nil),
|
||||
accountRepo: repo,
|
||||
}
|
||||
|
||||
result, err := svc.forwardGrokResponses(context.Background(), c, account, body, "grok", false, time.Now())
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, result)
|
||||
require.False(t, result.Stream)
|
||||
require.Equal(t, "resp_grok_non_stream", result.ResponseID)
|
||||
require.Equal(t, 7, result.Usage.InputTokens)
|
||||
require.Equal(t, 2, result.Usage.OutputTokens)
|
||||
require.Equal(t, 4, result.Usage.CacheReadInputTokens)
|
||||
require.Equal(t, "grok-4.5", gjson.GetBytes(upstream.lastBody, "model").String())
|
||||
identity := gjson.GetBytes(upstream.lastBody, "prompt_cache_key").String()
|
||||
require.NotEmpty(t, identity)
|
||||
require.Equal(t, identity, upstream.lastReq.Header.Get(grokConversationIDHeader))
|
||||
// The sanitizer drops this unsupported client tool, but its explicit intent
|
||||
// must still prevent native cache-routing tools from being injected.
|
||||
require.False(t, gjson.GetBytes(upstream.lastBody, "tools").Exists())
|
||||
require.False(t, gjson.GetBytes(upstream.lastBody, "tool_choice").Exists())
|
||||
require.Equal(t, "resp_grok_non_stream", gjson.Get(recorder.Body.String(), "id").String())
|
||||
}
|
||||
|
||||
func TestForwardGrokResponsesFailoverKeepsCacheIdentityAcrossAccounts(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
|
||||
recorder := httptest.NewRecorder()
|
||||
c, _ := gin.CreateTestContext(recorder)
|
||||
body := []byte(`{"model":"grok","input":[{"role":"user","content":"stable prefix"}],"stream":false}`)
|
||||
c.Request = httptest.NewRequest(http.MethodPost, "/v1/responses", bytes.NewReader(body))
|
||||
c.Set("api_key", &APIKey{ID: 5203})
|
||||
|
||||
newAccount := func(id int64, token string) *Account {
|
||||
return &Account{
|
||||
ID: id,
|
||||
Name: fmt.Sprintf("grok-%d", id),
|
||||
Platform: PlatformGrok,
|
||||
Type: AccountTypeOAuth,
|
||||
Concurrency: 1,
|
||||
Credentials: map[string]any{
|
||||
"access_token": token,
|
||||
"expires_at": time.Now().Add(time.Hour).UTC().Format(time.RFC3339),
|
||||
"base_url": xai.DefaultCLIBaseURL,
|
||||
},
|
||||
}
|
||||
}
|
||||
firstAccount := newAccount(58, "access-token-a")
|
||||
secondAccount := newAccount(59, "access-token-b")
|
||||
repo := &grokQuotaAccountRepo{
|
||||
mockAccountRepoForPlatform: &mockAccountRepoForPlatform{
|
||||
accountsByID: map[int64]*Account{58: firstAccount, 59: secondAccount},
|
||||
},
|
||||
}
|
||||
upstream := &httpUpstreamRecorder{responses: []*http.Response{
|
||||
{
|
||||
StatusCode: http.StatusServiceUnavailable,
|
||||
Header: http.Header{"Content-Type": []string{"application/json"}},
|
||||
Body: io.NopCloser(strings.NewReader(`{"error":{"message":"temporary"}}`)),
|
||||
},
|
||||
{
|
||||
StatusCode: http.StatusOK,
|
||||
Header: http.Header{"Content-Type": []string{"application/json"}},
|
||||
Body: io.NopCloser(strings.NewReader(`{"id":"resp_after_failover","object":"response","model":"grok-4.3","status":"completed","output":[],"usage":{"input_tokens":5,"output_tokens":1}}`)),
|
||||
},
|
||||
}}
|
||||
svc := &OpenAIGatewayService{
|
||||
httpUpstream: upstream,
|
||||
grokTokenProvider: NewGrokTokenProvider(repo, nil),
|
||||
accountRepo: repo,
|
||||
}
|
||||
|
||||
_, err := svc.forwardGrokResponses(context.Background(), c, firstAccount, body, "grok", false, time.Now())
|
||||
var failoverErr *UpstreamFailoverError
|
||||
require.ErrorAs(t, err, &failoverErr)
|
||||
|
||||
result, err := svc.forwardGrokResponses(context.Background(), c, secondAccount, body, "grok", false, time.Now())
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, result)
|
||||
require.Len(t, upstream.requests, 2)
|
||||
require.Len(t, upstream.bodies, 2)
|
||||
firstIdentity := gjson.GetBytes(upstream.bodies[0], "prompt_cache_key").String()
|
||||
secondIdentity := gjson.GetBytes(upstream.bodies[1], "prompt_cache_key").String()
|
||||
require.NotEmpty(t, firstIdentity)
|
||||
require.Equal(t, firstIdentity, secondIdentity)
|
||||
require.Equal(t, firstIdentity, upstream.requests[0].Header.Get(grokConversationIDHeader))
|
||||
require.Equal(t, secondIdentity, upstream.requests[1].Header.Get(grokConversationIDHeader))
|
||||
require.Equal(t, "Bearer access-token-a", upstream.requests[0].Header.Get("Authorization"))
|
||||
require.Equal(t, "Bearer access-token-b", upstream.requests[1].Header.Get("Authorization"))
|
||||
}
|
||||
|
||||
func TestForwardAsChatCompletionsForGrokStreamingUsesRawXAIChatCompletions(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
|
||||
@@ -742,6 +884,8 @@ func TestForwardAsChatCompletionsForGrokStreamingUsesRawXAIChatCompletions(t *te
|
||||
body := []byte(`{"model":"grok","messages":[{"role":"user","content":"hi"}],"stream":true}`)
|
||||
c.Request = httptest.NewRequest(http.MethodPost, "/v1/chat/completions", bytes.NewReader(body))
|
||||
c.Request.Header.Set("Content-Type", "application/json")
|
||||
c.Request.Header.Set(grokConversationIDHeader, "native-client-conversation")
|
||||
c.Set("api_key", &APIKey{ID: 5301})
|
||||
|
||||
account := &Account{
|
||||
ID: 53,
|
||||
@@ -791,6 +935,9 @@ func TestForwardAsChatCompletionsForGrokStreamingUsesRawXAIChatCompletions(t *te
|
||||
require.Equal(t, "Bearer access-token", upstream.lastReq.Header.Get("Authorization"))
|
||||
require.Equal(t, "text/event-stream", upstream.lastReq.Header.Get("Accept"))
|
||||
require.Equal(t, "sub2api-grok/1.0", upstream.lastReq.Header.Get("User-Agent"))
|
||||
require.Equal(t, grokCLIVersion, upstream.lastReq.Header.Get("X-Grok-Client-Version"))
|
||||
require.NotEmpty(t, upstream.lastReq.Header.Get(grokConversationIDHeader))
|
||||
require.NotEqual(t, "native-client-conversation", upstream.lastReq.Header.Get(grokConversationIDHeader))
|
||||
require.Equal(t, "grok-4.5", gjson.GetBytes(upstream.lastBody, "model").String())
|
||||
require.True(t, gjson.GetBytes(upstream.lastBody, "stream_options.include_usage").Bool())
|
||||
require.True(t, result.Stream)
|
||||
@@ -809,6 +956,7 @@ func TestForwardAsChatCompletionsForGrokComposerBridgesImageInput(t *testing.T)
|
||||
body := []byte(`{"model":"grok-composer-2.5-fast","messages":[{"role":"system","content":"You are concise."},{"role":"user","content":[{"type":"text","text":"What is shown?"},{"type":"image_url","image_url":{"url":"data:image/png;base64,QUJD"}}]}],"stream":false}`)
|
||||
c.Request = httptest.NewRequest(http.MethodPost, "/v1/chat/completions", bytes.NewReader(body))
|
||||
c.Request.Header.Set("Content-Type", "application/json")
|
||||
c.Set("api_key", &APIKey{ID: 5501})
|
||||
|
||||
account := &Account{
|
||||
ID: 55,
|
||||
@@ -858,9 +1006,11 @@ func TestForwardAsChatCompletionsForGrokComposerBridgesImageInput(t *testing.T)
|
||||
require.NotNil(t, result)
|
||||
require.Len(t, upstream.requests, 2)
|
||||
require.Equal(t, xai.DefaultCLIBaseURL+"/responses", upstream.requests[0].URL.String())
|
||||
require.Empty(t, upstream.requests[0].Header.Get(grokConversationIDHeader))
|
||||
require.Equal(t, "grok-build-0.1", gjson.GetBytes(upstream.bodies[0], "model").String())
|
||||
require.Equal(t, "input_image", gjson.GetBytes(upstream.bodies[0], "input.0.content.1.type").String())
|
||||
require.Equal(t, xai.DefaultCLIBaseURL+"/chat/completions", upstream.requests[1].URL.String())
|
||||
require.NotEmpty(t, upstream.requests[1].Header.Get(grokConversationIDHeader))
|
||||
require.Equal(t, "grok-composer-2.5-fast", gjson.GetBytes(upstream.bodies[1], "model").String())
|
||||
require.False(t, strings.Contains(string(upstream.bodies[1]), "image_url"))
|
||||
require.Contains(t, gjson.GetBytes(upstream.bodies[1], "messages.1.content").String(), "Image 1 description")
|
||||
@@ -878,6 +1028,7 @@ func TestForwardAsAnthropicForGrokUsesXAIResponses(t *testing.T) {
|
||||
c, _ := gin.CreateTestContext(recorder)
|
||||
body := []byte(`{"model":"grok","max_tokens":32,"stream":false,"messages":[{"role":"user","content":"hi"}]}`)
|
||||
c.Request = httptest.NewRequest(http.MethodPost, "/v1/messages", bytes.NewReader(body))
|
||||
c.Set("api_key", &APIKey{ID: 5401})
|
||||
|
||||
account := &Account{
|
||||
ID: 54,
|
||||
@@ -896,7 +1047,7 @@ func TestForwardAsAnthropicForGrokUsesXAIResponses(t *testing.T) {
|
||||
accountsByID: map[int64]*Account{54: account},
|
||||
},
|
||||
}
|
||||
upstream := &httpUpstreamRecorder{resp: openAICompatSSECompletedResponse("resp_grok_messages", "grok-4.3")}
|
||||
upstream := &httpUpstreamRecorder{resp: grokMessagesSSECompletedResponse("resp_grok_messages", 3)}
|
||||
svc := &OpenAIGatewayService{
|
||||
httpUpstream: upstream,
|
||||
grokTokenProvider: NewGrokTokenProvider(repo, nil),
|
||||
@@ -908,17 +1059,84 @@ func TestForwardAsAnthropicForGrokUsesXAIResponses(t *testing.T) {
|
||||
require.Equal(t, xai.DefaultCLIBaseURL+"/responses", upstream.lastReq.URL.String())
|
||||
require.Equal(t, "Bearer access-token", upstream.lastReq.Header.Get("Authorization"))
|
||||
require.Equal(t, "sub2api-grok/1.0", upstream.lastReq.Header.Get("User-Agent"))
|
||||
require.Equal(t, grokCLIVersion, upstream.lastReq.Header.Get("X-Grok-Client-Version"))
|
||||
require.Equal(t, "grok-4.5", gjson.GetBytes(upstream.lastBody, "model").String())
|
||||
require.NotEmpty(t, gjson.GetBytes(upstream.lastBody, "prompt_cache_key").String())
|
||||
require.Equal(t, gjson.GetBytes(upstream.lastBody, "prompt_cache_key").String(), upstream.lastReq.Header.Get(grokConversationIDHeader))
|
||||
require.Equal(t, "web_search", gjson.GetBytes(upstream.lastBody, "tools.0.type").String())
|
||||
require.Equal(t, "x_search", gjson.GetBytes(upstream.lastBody, "tools.1.type").String())
|
||||
require.Equal(t, "none", gjson.GetBytes(upstream.lastBody, "tool_choice").String())
|
||||
require.Empty(t, upstream.lastReq.Header.Get("session_id"))
|
||||
require.True(t, gjson.GetBytes(upstream.lastBody, "stream").Bool())
|
||||
require.NotContains(t, string(upstream.lastBody), "chatgpt.com")
|
||||
require.Equal(t, "grok", result.Model)
|
||||
require.Equal(t, "grok-4.5", result.UpstreamModel)
|
||||
require.Equal(t, 5, result.Usage.InputTokens)
|
||||
require.Equal(t, 2, result.Usage.OutputTokens)
|
||||
require.Equal(t, 3, result.Usage.CacheReadInputTokens)
|
||||
require.Contains(t, recorder.Body.String(), `"type":"message"`)
|
||||
require.Equal(t, int64(3), gjson.Get(recorder.Body.String(), "usage.cache_read_input_tokens").Int())
|
||||
require.Contains(t, recorder.Body.String(), "ok")
|
||||
}
|
||||
|
||||
func TestForwardAsAnthropicForGrokStreamingPreservesCacheUsage(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
|
||||
recorder := httptest.NewRecorder()
|
||||
c, _ := gin.CreateTestContext(recorder)
|
||||
body := []byte(`{"model":"grok","max_tokens":32,"stream":true,"messages":[{"role":"user","content":"hi"}]}`)
|
||||
c.Request = httptest.NewRequest(http.MethodPost, "/v1/messages", bytes.NewReader(body))
|
||||
c.Set("api_key", &APIKey{ID: 5402})
|
||||
|
||||
account := &Account{
|
||||
ID: 57,
|
||||
Name: "grok",
|
||||
Platform: PlatformGrok,
|
||||
Type: AccountTypeOAuth,
|
||||
Concurrency: 1,
|
||||
Credentials: map[string]any{
|
||||
"access_token": "access-token",
|
||||
"expires_at": time.Now().Add(time.Hour).UTC().Format(time.RFC3339),
|
||||
"base_url": xai.DefaultCLIBaseURL,
|
||||
},
|
||||
}
|
||||
repo := &grokQuotaAccountRepo{
|
||||
mockAccountRepoForPlatform: &mockAccountRepoForPlatform{
|
||||
accountsByID: map[int64]*Account{57: account},
|
||||
},
|
||||
}
|
||||
upstream := &httpUpstreamRecorder{resp: grokMessagesSSECompletedResponse("resp_grok_messages_stream", 2)}
|
||||
svc := &OpenAIGatewayService{
|
||||
httpUpstream: upstream,
|
||||
grokTokenProvider: NewGrokTokenProvider(repo, nil),
|
||||
accountRepo: repo,
|
||||
}
|
||||
|
||||
result, err := svc.ForwardAsAnthropic(context.Background(), c, account, body, "", "")
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, result)
|
||||
require.Equal(t, 2, result.Usage.CacheReadInputTokens)
|
||||
identity := gjson.GetBytes(upstream.lastBody, "prompt_cache_key").String()
|
||||
require.NotEmpty(t, identity)
|
||||
require.Equal(t, identity, upstream.lastReq.Header.Get(grokConversationIDHeader))
|
||||
require.Contains(t, recorder.Header().Get("Content-Type"), "text/event-stream")
|
||||
require.Contains(t, recorder.Body.String(), `"cache_read_input_tokens":2`)
|
||||
}
|
||||
|
||||
func grokMessagesSSECompletedResponse(responseID string, cachedTokens int) *http.Response {
|
||||
body := strings.Join([]string{
|
||||
fmt.Sprintf(`data: {"type":"response.completed","response":{"id":%q,"object":"response","model":"grok-4.3","status":"completed","output":[{"type":"message","id":"msg_1","role":"assistant","status":"completed","content":[{"type":"output_text","text":"ok"}]}],"usage":{"input_tokens":5,"output_tokens":2,"total_tokens":7,"input_tokens_details":{"cached_tokens":%d}}}}`, responseID, cachedTokens),
|
||||
"",
|
||||
"data: [DONE]",
|
||||
"",
|
||||
}, "\n")
|
||||
return &http.Response{
|
||||
StatusCode: http.StatusOK,
|
||||
Header: http.Header{"Content-Type": []string{"text/event-stream"}},
|
||||
Body: io.NopCloser(strings.NewReader(body)),
|
||||
}
|
||||
}
|
||||
|
||||
func TestHandleGrokAccountUpstreamErrorTempUnschedulesReadinessStates(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
|
||||
@@ -244,12 +244,17 @@ func (s *OpenAIGatewayService) ForwardAsAnthropic(
|
||||
return nil, policyErr
|
||||
}
|
||||
responsesBody = updatedBody
|
||||
grokCacheIdentity := ""
|
||||
if account.Platform == PlatformGrok {
|
||||
grokCacheIdentity = resolveGrokCacheIdentity(c, responsesBody, promptCacheKey, upstreamModel)
|
||||
patchedBody, patchErr := patchGrokResponsesBody(responsesBody, upstreamModel)
|
||||
if patchErr != nil {
|
||||
return nil, patchErr
|
||||
}
|
||||
responsesBody = patchedBody
|
||||
responsesBody, patchErr = applyGrokResponsesCacheIdentity(patchedBody, responsesBody, grokCacheIdentity, account.IsGrokOAuth())
|
||||
if patchErr != nil {
|
||||
return nil, fmt.Errorf("apply grok prompt cache identity: %w", patchErr)
|
||||
}
|
||||
}
|
||||
|
||||
// 5. Get access token
|
||||
@@ -269,7 +274,7 @@ func (s *OpenAIGatewayService) ForwardAsAnthropic(
|
||||
upstreamCtx, releaseUpstreamCtx := detachUpstreamContext(ctx)
|
||||
var upstreamReq *http.Request
|
||||
if account.Platform == PlatformGrok {
|
||||
upstreamReq, err = buildGrokResponsesRequest(upstreamCtx, c, account, responsesBody, token)
|
||||
upstreamReq, err = buildGrokResponsesRequest(upstreamCtx, c, account, responsesBody, token, grokCacheIdentity)
|
||||
} else {
|
||||
upstreamReq, err = s.buildUpstreamRequest(upstreamCtx, c, account, responsesBody, token, isStream, promptCacheKey, false)
|
||||
}
|
||||
@@ -280,7 +285,7 @@ func (s *OpenAIGatewayService) ForwardAsAnthropic(
|
||||
|
||||
// Override session_id with a deterministic UUID derived from the isolated
|
||||
// session key, ensuring different API keys produce different upstream sessions.
|
||||
if promptCacheKey != "" {
|
||||
if account.Platform != PlatformGrok && promptCacheKey != "" {
|
||||
isolatedSessionID := generateSessionUUID(isolateOpenAISessionID(apiKeyID, promptCacheKey))
|
||||
upstreamReq.Header.Set("session_id", isolatedSessionID)
|
||||
if upstreamReq.Header.Get("conversation_id") != "" {
|
||||
|
||||
@@ -101,7 +101,7 @@ func (s *OpenAIGatewayService) forwardAnthropicViaRawChatCompletions(
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
resp, err := s.sendCCUpstreamRequest(ctx, c, account, targetURL, chatBody, clientStream, apiKey, account.GetOpenAIUserAgent())
|
||||
resp, err := s.sendCCUpstreamRequest(ctx, c, account, targetURL, chatBody, clientStream, apiKey, account.GetOpenAIUserAgent(), "")
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
@@ -92,7 +92,7 @@ func (s *OpenAIGatewayService) forwardResponsesViaRawChatCompletions(
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
resp, err := s.sendCCUpstreamRequest(ctx, c, account, targetURL, chatBody, clientStream, apiKey, account.GetOpenAIUserAgent())
|
||||
resp, err := s.sendCCUpstreamRequest(ctx, c, account, targetURL, chatBody, clientStream, apiKey, account.GetOpenAIUserAgent(), "")
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
@@ -23,17 +23,7 @@ import (
|
||||
// ExtractSessionID extracts the raw session ID from headers or body without hashing.
|
||||
// Used by ForwardAsAnthropic to pass as prompt_cache_key for upstream cache.
|
||||
func (s *OpenAIGatewayService) ExtractSessionID(c *gin.Context, body []byte) string {
|
||||
if c == nil {
|
||||
return ""
|
||||
}
|
||||
sessionID := strings.TrimSpace(c.GetHeader("session_id"))
|
||||
if sessionID == "" {
|
||||
sessionID = strings.TrimSpace(c.GetHeader("conversation_id"))
|
||||
}
|
||||
if sessionID == "" && len(body) > 0 {
|
||||
sessionID = strings.TrimSpace(gjson.GetBytes(body, "prompt_cache_key").String())
|
||||
}
|
||||
return sessionID
|
||||
return explicitOpenAIRequestSessionID(c, body)
|
||||
}
|
||||
|
||||
func explicitOpenAISessionID(c *gin.Context, body []byte) string {
|
||||
@@ -51,11 +41,33 @@ func explicitOpenAISessionID(c *gin.Context, body []byte) string {
|
||||
return sessionID
|
||||
}
|
||||
|
||||
// explicitOpenAIRequestSessionID extends the common OpenAI session signals
|
||||
// with Grok's native conversation header only for requests authenticated to a
|
||||
// Grok group. This keeps an unrelated x-grok-conv-id header from changing
|
||||
// scheduling or upstream session behavior for non-Grok groups.
|
||||
func explicitOpenAIRequestSessionID(c *gin.Context, body []byte) string {
|
||||
if c == nil {
|
||||
return ""
|
||||
}
|
||||
|
||||
sessionID := strings.TrimSpace(c.GetHeader("session_id"))
|
||||
if sessionID == "" {
|
||||
sessionID = strings.TrimSpace(c.GetHeader("conversation_id"))
|
||||
}
|
||||
if sessionID == "" && isGrokRequestContext(c) {
|
||||
sessionID = strings.TrimSpace(c.GetHeader(grokConversationIDHeader))
|
||||
}
|
||||
if sessionID == "" && len(body) > 0 {
|
||||
sessionID = strings.TrimSpace(gjson.GetBytes(body, "prompt_cache_key").String())
|
||||
}
|
||||
return sessionID
|
||||
}
|
||||
|
||||
// GenerateExplicitSessionHash generates a sticky-session hash only from explicit
|
||||
// client session signals. It intentionally skips content-derived fallback and is
|
||||
// used by stateless endpoints such as /v1/images.
|
||||
func (s *OpenAIGatewayService) GenerateExplicitSessionHash(c *gin.Context, body []byte) string {
|
||||
sessionID := explicitOpenAISessionID(c, body)
|
||||
sessionID := explicitOpenAIRequestSessionID(c, body)
|
||||
if sessionID == "" {
|
||||
return ""
|
||||
}
|
||||
@@ -70,14 +82,15 @@ func (s *OpenAIGatewayService) GenerateExplicitSessionHash(c *gin.Context, body
|
||||
// Priority:
|
||||
// 1. Header: session_id
|
||||
// 2. Header: conversation_id
|
||||
// 3. Body: prompt_cache_key (opencode)
|
||||
// 4. Body: content-based fallback (model + system + tools + first user message)
|
||||
// 3. Header: x-grok-conv-id (Grok groups only)
|
||||
// 4. Body: prompt_cache_key (opencode)
|
||||
// 5. Body: content-based fallback (model + system + tools + first user message)
|
||||
func (s *OpenAIGatewayService) GenerateSessionHash(c *gin.Context, body []byte) string {
|
||||
if c == nil {
|
||||
return ""
|
||||
}
|
||||
|
||||
sessionID := explicitOpenAISessionID(c, body)
|
||||
sessionID := explicitOpenAIRequestSessionID(c, body)
|
||||
if sessionID == "" && len(body) > 0 {
|
||||
sessionID = deriveOpenAIContentSessionSeed(body)
|
||||
}
|
||||
|
||||
@@ -421,6 +421,10 @@ func (s *OpenAIGatewayService) ProxyResponsesWebSocketFromClient(
|
||||
storeDisabled,
|
||||
)
|
||||
currentBridgePayload := firstPayload
|
||||
// Keep the first turn as the stable conversation seed. The mapped model
|
||||
// is resolved again for each turn below so an in-connection model switch
|
||||
// cannot reuse another model's upstream cache identity.
|
||||
grokCacheSeedPayload := firstPayload.payloadRaw
|
||||
var bridgeReplayInput []json.RawMessage
|
||||
bridgeReplayInputExists := false
|
||||
for turn := 1; ; turn++ {
|
||||
@@ -469,6 +473,13 @@ func (s *OpenAIGatewayService) ProxyResponsesWebSocketFromClient(
|
||||
openAIWSRawPayloadHasToolCallOutput(currentBridgePayload.payloadRaw),
|
||||
)
|
||||
}
|
||||
grokCacheIdentity := ""
|
||||
if account.Platform == PlatformGrok {
|
||||
grokCacheIdentity, err = resolveGrokWSCacheIdentity(c, account, grokCacheSeedPayload, currentBridgePayload.originalModel)
|
||||
if err != nil {
|
||||
return fmt.Errorf("resolve Grok websocket cache identity: %w", err)
|
||||
}
|
||||
}
|
||||
result, bridgeErr := s.proxyOpenAIWSHTTPBridgeTurn(
|
||||
ctx,
|
||||
c,
|
||||
@@ -480,6 +491,7 @@ func (s *OpenAIGatewayService) ProxyResponsesWebSocketFromClient(
|
||||
currentBridgePayload.imageBillingModel,
|
||||
currentBridgePayload.imageSizeTier,
|
||||
currentBridgePayload.imageInputSize,
|
||||
grokCacheIdentity,
|
||||
turn,
|
||||
writeClientMessage,
|
||||
)
|
||||
|
||||
@@ -155,6 +155,7 @@ func (s *OpenAIGatewayService) proxyOpenAIWSHTTPBridgeTurn(
|
||||
imageBillingModel string,
|
||||
imageSizeTier string,
|
||||
imageInputSize string,
|
||||
grokCacheIdentity string,
|
||||
turn int,
|
||||
writeClientMessage func([]byte) error,
|
||||
) (*OpenAIForwardResult, error) {
|
||||
@@ -179,21 +180,19 @@ func (s *OpenAIGatewayService) proxyOpenAIWSHTTPBridgeTurn(
|
||||
upstreamCtx, releaseUpstreamCtx := detachUpstreamContext(ctx)
|
||||
var upstreamReq *http.Request
|
||||
if account.Platform == PlatformGrok {
|
||||
upstreamModel := strings.TrimSpace(gjson.GetBytes(body, "model").String())
|
||||
if originalModel != "" {
|
||||
if mappedModel := normalizeOpenAIModelForUpstream(account, account.GetMappedModel(originalModel)); mappedModel != "" {
|
||||
upstreamModel = mappedModel
|
||||
}
|
||||
}
|
||||
if upstreamModel == "" {
|
||||
upstreamModel = "grok-4.3"
|
||||
}
|
||||
upstreamModel := resolveGrokWSUpstreamModel(account, body, originalModel)
|
||||
grokIntentSourceBody := body
|
||||
body, err = patchGrokResponsesBody(body, upstreamModel)
|
||||
if err != nil {
|
||||
releaseUpstreamCtx()
|
||||
return nil, err
|
||||
}
|
||||
upstreamReq, err = buildGrokResponsesRequest(upstreamCtx, c, account, body, token)
|
||||
body, err = applyGrokResponsesCacheIdentity(body, grokIntentSourceBody, grokCacheIdentity, account.IsGrokOAuth())
|
||||
if err != nil {
|
||||
releaseUpstreamCtx()
|
||||
return nil, fmt.Errorf("apply grok prompt cache identity: %w", err)
|
||||
}
|
||||
upstreamReq, err = buildGrokResponsesRequest(upstreamCtx, c, account, body, token, grokCacheIdentity)
|
||||
} else {
|
||||
upstreamReq, err = s.buildUpstreamRequestOpenAIPassthrough(upstreamCtx, c, account, body, token)
|
||||
}
|
||||
@@ -407,3 +406,25 @@ func (s *OpenAIGatewayService) proxyOpenAIWSHTTPBridgeTurn(
|
||||
}
|
||||
return resultWithUsage(), errors.New("upstream http bridge stream ended before terminal event")
|
||||
}
|
||||
|
||||
func resolveGrokWSCacheIdentity(c *gin.Context, account *Account, payload []byte, originalModel string) (string, error) {
|
||||
body, err := prepareOpenAIWSHTTPBridgeBody(payload)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
upstreamModel := resolveGrokWSUpstreamModel(account, body, originalModel)
|
||||
return resolveGrokCacheIdentity(c, body, "", upstreamModel), nil
|
||||
}
|
||||
|
||||
func resolveGrokWSUpstreamModel(account *Account, body []byte, originalModel string) string {
|
||||
upstreamModel := strings.TrimSpace(gjson.GetBytes(body, "model").String())
|
||||
if account != nil && originalModel != "" {
|
||||
if mappedModel := normalizeOpenAIModelForUpstream(account, account.GetMappedModel(originalModel)); mappedModel != "" {
|
||||
upstreamModel = mappedModel
|
||||
}
|
||||
}
|
||||
if upstreamModel == "" {
|
||||
upstreamModel = "grok-4.3"
|
||||
}
|
||||
return upstreamModel
|
||||
}
|
||||
|
||||
@@ -3,6 +3,7 @@ package service
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
@@ -127,6 +128,7 @@ func TestOpenAIWSHTTPBridgeRelaysSSEFramesAsWebSocketMessages(t *testing.T) {
|
||||
"",
|
||||
"",
|
||||
"",
|
||||
"",
|
||||
1,
|
||||
writeClient,
|
||||
)
|
||||
@@ -179,21 +181,28 @@ func TestOpenAIWSHTTPBridgeRelaysSSEFramesAsWebSocketMessages(t *testing.T) {
|
||||
func TestProxyResponsesWebSocketFromClientForGrokUsesXAIHTTPBridge(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
|
||||
sseBody := strings.Join([]string{
|
||||
`data: {"type":"response.created","response":{"id":"resp_grok_ws","model":"grok-4.3"}}`,
|
||||
"",
|
||||
`data: {"type":"response.output_text.delta","response":{"id":"resp_grok_ws"},"delta":"ok"}`,
|
||||
"",
|
||||
`data: {"type":"response.completed","response":{"id":"resp_grok_ws","model":"grok-4.3","usage":{"input_tokens":4,"output_tokens":2}}}`,
|
||||
"",
|
||||
}, "\n")
|
||||
upstream := &httpUpstreamRecorder{resp: &http.Response{
|
||||
StatusCode: http.StatusOK,
|
||||
Header: http.Header{
|
||||
"Content-Type": []string{"text/event-stream"},
|
||||
"Xai-Request-Id": []string{"xai-ws-req"},
|
||||
},
|
||||
Body: io.NopCloser(strings.NewReader(sseBody)),
|
||||
bridgeResponse := func(responseID, requestID string, cachedTokens int) *http.Response {
|
||||
sseBody := strings.Join([]string{
|
||||
`data: {"type":"response.created","response":{"id":"` + responseID + `","model":"grok-4.3"}}`,
|
||||
"",
|
||||
`data: {"type":"response.output_text.delta","response":{"id":"` + responseID + `"},"delta":"ok"}`,
|
||||
"",
|
||||
`data: {"type":"response.completed","response":{"id":"` + responseID + `","model":"grok-4.3","usage":{"input_tokens":4,"output_tokens":2,"input_tokens_details":{"cached_tokens":` + fmt.Sprintf("%d", cachedTokens) + `}}}}`,
|
||||
"",
|
||||
}, "\n")
|
||||
return &http.Response{
|
||||
StatusCode: http.StatusOK,
|
||||
Header: http.Header{
|
||||
"Content-Type": []string{"text/event-stream"},
|
||||
"Xai-Request-Id": []string{requestID},
|
||||
},
|
||||
Body: io.NopCloser(strings.NewReader(sseBody)),
|
||||
}
|
||||
}
|
||||
upstream := &httpUpstreamRecorder{responses: []*http.Response{
|
||||
bridgeResponse("resp_grok_ws_1", "xai-ws-req-1", 0),
|
||||
bridgeResponse("resp_grok_ws_2", "xai-ws-req-2", 3),
|
||||
bridgeResponse("resp_grok_ws_3", "xai-ws-req-3", 0),
|
||||
}}
|
||||
svc := &OpenAIGatewayService{
|
||||
cfg: &config.Config{
|
||||
@@ -241,6 +250,7 @@ func TestProxyResponsesWebSocketFromClientForGrokUsesXAIHTTPBridge(t *testing.T)
|
||||
req := r.Clone(r.Context())
|
||||
req.Header = req.Header.Clone()
|
||||
ginCtx.Request = req
|
||||
ginCtx.Set("api_key", &APIKey{ID: 7101})
|
||||
|
||||
errCh <- svc.ProxyResponsesWebSocketFromClient(r.Context(), ginCtx, conn, account, "access-token", firstMessage, nil)
|
||||
}))
|
||||
@@ -271,6 +281,33 @@ func TestProxyResponsesWebSocketFromClientForGrokUsesXAIHTTPBridge(t *testing.T)
|
||||
require.Equal(t, "response.created", gjson.GetBytes(created, "type").String())
|
||||
require.Equal(t, "response.output_text.delta", gjson.GetBytes(delta, "type").String())
|
||||
require.Equal(t, "response.completed", gjson.GetBytes(completed, "type").String())
|
||||
require.Equal(t, "resp_grok_ws_1", gjson.GetBytes(completed, "response.id").String())
|
||||
|
||||
writeCtx, cancelWrite = context.WithTimeout(context.Background(), 3*time.Second)
|
||||
err = clientConn.Write(writeCtx, coderws.MessageText, []byte(`{"type":"response.create","generate":true,"model":"grok","stream":true,"previous_response_id":"resp_grok_ws_1","input":"second turn"}`))
|
||||
cancelWrite()
|
||||
require.NoError(t, err)
|
||||
|
||||
created = readEvent()
|
||||
delta = readEvent()
|
||||
completed = readEvent()
|
||||
require.Equal(t, "response.created", gjson.GetBytes(created, "type").String())
|
||||
require.Equal(t, "response.output_text.delta", gjson.GetBytes(delta, "type").String())
|
||||
require.Equal(t, "response.completed", gjson.GetBytes(completed, "type").String())
|
||||
require.Equal(t, "resp_grok_ws_2", gjson.GetBytes(completed, "response.id").String())
|
||||
|
||||
writeCtx, cancelWrite = context.WithTimeout(context.Background(), 3*time.Second)
|
||||
err = clientConn.Write(writeCtx, coderws.MessageText, []byte(`{"type":"response.create","generate":true,"model":"grok-4.3","stream":true,"previous_response_id":"resp_grok_ws_2","input":"third turn with a different model"}`))
|
||||
cancelWrite()
|
||||
require.NoError(t, err)
|
||||
|
||||
created = readEvent()
|
||||
delta = readEvent()
|
||||
completed = readEvent()
|
||||
require.Equal(t, "response.created", gjson.GetBytes(created, "type").String())
|
||||
require.Equal(t, "response.output_text.delta", gjson.GetBytes(delta, "type").String())
|
||||
require.Equal(t, "response.completed", gjson.GetBytes(completed, "type").String())
|
||||
require.Equal(t, "resp_grok_ws_3", gjson.GetBytes(completed, "response.id").String())
|
||||
|
||||
_ = clientConn.Close(coderws.StatusNormalClosure, "done")
|
||||
select {
|
||||
@@ -280,10 +317,30 @@ func TestProxyResponsesWebSocketFromClientForGrokUsesXAIHTTPBridge(t *testing.T)
|
||||
require.Fail(t, "proxy did not finish after client close")
|
||||
}
|
||||
|
||||
require.Len(t, upstream.requests, 3)
|
||||
require.Len(t, upstream.bodies, 3)
|
||||
require.Equal(t, xai.DefaultCLIBaseURL+"/responses", upstream.lastReq.URL.String())
|
||||
require.Equal(t, "Bearer access-token", upstream.lastReq.Header.Get("Authorization"))
|
||||
require.Equal(t, "sub2api-grok/1.0", upstream.lastReq.Header.Get("User-Agent"))
|
||||
require.Equal(t, "grok-4.5", gjson.GetBytes(upstream.lastBody, "model").String())
|
||||
require.Equal(t, grokCLIVersion, upstream.lastReq.Header.Get("X-Grok-Client-Version"))
|
||||
require.Equal(t, "grok-4.5", gjson.GetBytes(upstream.bodies[0], "model").String())
|
||||
require.Equal(t, "grok-4.5", gjson.GetBytes(upstream.bodies[1], "model").String())
|
||||
require.Equal(t, "grok-4.3", gjson.GetBytes(upstream.bodies[2], "model").String())
|
||||
require.NotEmpty(t, gjson.GetBytes(upstream.lastBody, "prompt_cache_key").String())
|
||||
require.Equal(t, gjson.GetBytes(upstream.lastBody, "prompt_cache_key").String(), upstream.lastReq.Header.Get(grokConversationIDHeader))
|
||||
require.Equal(t, "web_search", gjson.GetBytes(upstream.lastBody, "tools.0.type").String())
|
||||
require.Equal(t, "x_search", gjson.GetBytes(upstream.lastBody, "tools.1.type").String())
|
||||
require.Equal(t, "none", gjson.GetBytes(upstream.lastBody, "tool_choice").String())
|
||||
firstIdentity := gjson.GetBytes(upstream.bodies[0], "prompt_cache_key").String()
|
||||
secondIdentity := gjson.GetBytes(upstream.bodies[1], "prompt_cache_key").String()
|
||||
thirdIdentity := gjson.GetBytes(upstream.bodies[2], "prompt_cache_key").String()
|
||||
require.NotEmpty(t, firstIdentity)
|
||||
require.Equal(t, firstIdentity, secondIdentity)
|
||||
require.NotEmpty(t, thirdIdentity)
|
||||
require.NotEqual(t, firstIdentity, thirdIdentity)
|
||||
require.Equal(t, firstIdentity, upstream.requests[0].Header.Get(grokConversationIDHeader))
|
||||
require.Equal(t, secondIdentity, upstream.requests[1].Header.Get(grokConversationIDHeader))
|
||||
require.Equal(t, thirdIdentity, upstream.requests[2].Header.Get(grokConversationIDHeader))
|
||||
require.False(t, gjson.GetBytes(upstream.lastBody, "type").Exists())
|
||||
require.False(t, gjson.GetBytes(upstream.lastBody, "generate").Exists())
|
||||
require.False(t, gjson.GetBytes(upstream.lastBody, "prompt_cache_retention").Exists())
|
||||
|
||||
Reference in New Issue
Block a user