From 8a87a658ad2b226bc371941c58429443912189ad Mon Sep 17 00:00:00 2001 From: Heatherm Huang Date: Wed, 17 Jun 2026 18:38:23 +0800 Subject: [PATCH] test: cover grok readiness paths --- .../service/grok_token_provider_test.go | 93 ++++++++++++ .../service/openai_gateway_grok_test.go | 136 ++++++++++++++++++ 2 files changed, 229 insertions(+) create mode 100644 backend/internal/service/grok_token_provider_test.go diff --git a/backend/internal/service/grok_token_provider_test.go b/backend/internal/service/grok_token_provider_test.go new file mode 100644 index 0000000000..d147602a43 --- /dev/null +++ b/backend/internal/service/grok_token_provider_test.go @@ -0,0 +1,93 @@ +//go:build unit + +package service + +import ( + "context" + "errors" + "testing" + "time" + + "github.com/Wei-Shaw/sub2api/internal/pkg/xai" + "github.com/stretchr/testify/require" +) + +type grokTokenCacheForProviderTest struct { + token string + setKey string + setToken string + setTTL time.Duration + lockResult bool + releaseCalls int +} + +func (c *grokTokenCacheForProviderTest) GetAccessToken(context.Context, string) (string, error) { + if c.token == "" { + return "", errors.New("not cached") + } + return c.token, nil +} + +func (c *grokTokenCacheForProviderTest) SetAccessToken(_ context.Context, key string, token string, ttl time.Duration) error { + c.setKey = key + c.setToken = token + c.setTTL = ttl + return nil +} + +func (c *grokTokenCacheForProviderTest) DeleteAccessToken(context.Context, string) error { + return nil +} + +func (c *grokTokenCacheForProviderTest) AcquireRefreshLock(context.Context, string, time.Duration) (bool, error) { + return c.lockResult, nil +} + +func (c *grokTokenCacheForProviderTest) ReleaseRefreshLock(context.Context, string) error { + c.releaseCalls++ + return nil +} + +func TestGrokTokenProviderRefreshesExpiredTokenOnRequestPath(t *testing.T) { + t.Setenv(xai.EnvBaseURL, xai.DefaultCLIBaseURL) + + expiredAt := time.Now().Add(-time.Minute).UTC().Format(time.RFC3339) + account := &Account{ + ID: 54, + Platform: PlatformGrok, + Type: AccountTypeOAuth, + Credentials: map[string]any{ + "access_token": "expired-access-token", + "refresh_token": "refresh-token", + "expires_at": expiredAt, + "base_url": xai.DefaultCLIBaseURL, + "client_id": "client-id", + }, + } + repo := &tokenRefreshAccountRepo{} + repo.accountsByID = map[int64]*Account{54: account} + cache := &grokTokenCacheForProviderTest{lockResult: true} + oauthSvc := NewGrokOAuthService(nil, &grokOAuthClientStub{ + refreshResponse: &xai.TokenResponse{ + AccessToken: "new-access-token", + TokenType: "Bearer", + ExpiresIn: 3600, + }, + }) + defer oauthSvc.Stop() + + provider := NewGrokTokenProvider(repo, cache, oauthSvc) + provider.SetRefreshAPI(NewOAuthRefreshAPI(repo, cache), NewGrokTokenRefresher(oauthSvc)) + + token, err := provider.GetAccessToken(context.Background(), account) + require.NoError(t, err) + require.Equal(t, "new-access-token", token) + require.Equal(t, 1, repo.updateCredentialsCalls) + require.Equal(t, "new-access-token", repo.accountsByID[54].GetGrokAccessToken()) + require.Equal(t, "refresh-token", repo.accountsByID[54].GetGrokRefreshToken()) + require.Equal(t, xai.DefaultCLIBaseURL, repo.accountsByID[54].GetGrokBaseURL()) + require.Equal(t, "grok:account:54", cache.setKey) + require.Equal(t, "new-access-token", cache.setToken) + require.Greater(t, cache.setTTL, time.Duration(0)) + require.Equal(t, 1, cache.releaseCalls) +} diff --git a/backend/internal/service/openai_gateway_grok_test.go b/backend/internal/service/openai_gateway_grok_test.go index 12fda57229..4615f762f9 100644 --- a/backend/internal/service/openai_gateway_grok_test.go +++ b/backend/internal/service/openai_gateway_grok_test.go @@ -134,3 +134,139 @@ func TestForwardAsChatCompletionsForGrokUsesXAIChatCompletionsAndSnapshots(t *te require.NotNil(t, repo.updates[51][grokQuotaSnapshotExtraKey]) require.Equal(t, http.StatusOK, recorder.Code) } + +func TestForwardGrokResponsesStreamingUsesXAIResponsesAndSnapshots(t *testing.T) { + gin.SetMode(gin.TestMode) + + recorder := httptest.NewRecorder() + c, _ := gin.CreateTestContext(recorder) + body := []byte(`{"model":"grok","input":"hi","stream":true}`) + 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") + + account := &Account{ + ID: 52, + 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{52: account}, + }, + } + upstreamBody := strings.Join([]string{ + `data: {"type":"response.output_text.delta","sequence_number":0,"delta":"ok"}`, + "", + `data: {"type":"response.completed","sequence_number":1,"response":{"id":"resp_grok","model":"grok-4.3","usage":{"input_tokens":5,"output_tokens":3,"input_tokens_details":{"cached_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-stream-req"}, + "X-Ratelimit-Limit-Requests": []string{"10"}, + "X-Ratelimit-Remaining-Requests": []string{"8"}, + "X-Ratelimit-Limit-Tokens": []string{"1000"}, + "X-Ratelimit-Remaining-Tokens": []string{"990"}, + }, + Body: io.NopCloser(strings.NewReader(upstreamBody)), + }} + svc := &OpenAIGatewayService{ + httpUpstream: upstream, + grokTokenProvider: NewGrokTokenProvider(repo, nil, nil), + accountRepo: repo, + } + + result, err := svc.forwardGrokResponses(context.Background(), c, account, body, "grok", true, time.Now()) + require.NoError(t, err) + require.Equal(t, xai.DefaultCLIBaseURL+"/responses", upstream.lastReq.URL.String()) + 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.3", gjson.GetBytes(upstream.lastBody, "model").String()) + require.True(t, gjson.GetBytes(upstream.lastBody, "stream").Bool()) + require.True(t, result.Stream) + require.Equal(t, "resp_grok", result.ResponseID) + require.Equal(t, "xai-stream-req", result.RequestID) + require.Equal(t, 5, result.Usage.InputTokens) + require.Equal(t, 3, result.Usage.OutputTokens) + require.Equal(t, 2, result.Usage.CacheReadInputTokens) + require.Contains(t, recorder.Header().Get("Content-Type"), "text/event-stream") + require.Contains(t, recorder.Body.String(), "response.output_text.delta") + require.NotNil(t, repo.updates[52][grokQuotaSnapshotExtraKey]) +} + +func TestForwardAsChatCompletionsForGrokStreamingUsesRawXAIChatCompletions(t *testing.T) { + gin.SetMode(gin.TestMode) + + recorder := httptest.NewRecorder() + c, _ := gin.CreateTestContext(recorder) + 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") + + account := &Account{ + ID: 53, + 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{53: account}, + }, + } + upstreamBody := strings.Join([]string{ + `data: {"id":"chatcmpl_grok","object":"chat.completion.chunk","model":"grok-4.3","choices":[{"index":0,"delta":{"content":"ok"}}]}`, + "", + `data: {"id":"chatcmpl_grok","object":"chat.completion.chunk","model":"grok-4.3","choices":[],"usage":{"prompt_tokens":6,"completion_tokens":4,"total_tokens":10,"prompt_tokens_details":{"cached_tokens":1}}}`, + "", + "data: [DONE]", + "", + }, "\n") + upstream := &httpUpstreamRecorder{resp: &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{ + "Content-Type": []string{"text/event-stream"}, + "X-Request-Id": []string{"chat-stream-req"}, + "X-Ratelimit-Limit-Requests": []string{"10"}, + "X-Ratelimit-Remaining-Requests": []string{"7"}, + }, + Body: io.NopCloser(strings.NewReader(upstreamBody)), + }} + svc := &OpenAIGatewayService{ + cfg: rawChatCompletionsTestConfig(), + httpUpstream: upstream, + grokTokenProvider: NewGrokTokenProvider(repo, nil, nil), + accountRepo: repo, + } + + result, err := svc.ForwardAsChatCompletions(context.Background(), c, account, body, "", "") + 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.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, "grok-4.3", gjson.GetBytes(upstream.lastBody, "model").String()) + require.True(t, gjson.GetBytes(upstream.lastBody, "stream_options.include_usage").Bool()) + require.True(t, result.Stream) + require.Equal(t, 6, result.Usage.InputTokens) + require.Equal(t, 4, result.Usage.OutputTokens) + require.Equal(t, 1, result.Usage.CacheReadInputTokens) + require.Contains(t, recorder.Body.String(), "data: [DONE]") + require.NotNil(t, repo.updates[53][grokQuotaSnapshotExtraKey]) +}