148 lines
4.7 KiB
Go
148 lines
4.7 KiB
Go
//go:build unit
|
|
|
|
package admin
|
|
|
|
import (
|
|
"context"
|
|
"io"
|
|
"net/http"
|
|
"net/http/httptest"
|
|
"strings"
|
|
"testing"
|
|
"time"
|
|
|
|
"github.com/gin-gonic/gin"
|
|
"github.com/stretchr/testify/require"
|
|
|
|
"github.com/Wei-Shaw/sub2api/internal/pkg/tlsfingerprint"
|
|
"github.com/Wei-Shaw/sub2api/internal/pkg/xai"
|
|
"github.com/Wei-Shaw/sub2api/internal/service"
|
|
)
|
|
|
|
type grokQuotaHandlerAccountRepo struct {
|
|
service.AccountRepository
|
|
account *service.Account
|
|
updates map[int64]map[string]any
|
|
}
|
|
|
|
func (r *grokQuotaHandlerAccountRepo) GetByID(_ context.Context, id int64) (*service.Account, error) {
|
|
if r.account != nil && r.account.ID == id {
|
|
return r.account, nil
|
|
}
|
|
return nil, service.ErrAccountNotFound
|
|
}
|
|
|
|
func (r *grokQuotaHandlerAccountRepo) UpdateExtra(_ context.Context, id int64, updates map[string]any) error {
|
|
if r.updates == nil {
|
|
r.updates = make(map[int64]map[string]any)
|
|
}
|
|
r.updates[id] = updates
|
|
return nil
|
|
}
|
|
|
|
type grokQuotaHandlerUpstream struct {
|
|
resp *http.Response
|
|
lastReq *http.Request
|
|
lastBody []byte
|
|
}
|
|
|
|
func (u *grokQuotaHandlerUpstream) Do(req *http.Request, _ string, _ int64, _ int) (*http.Response, error) {
|
|
u.lastReq = req
|
|
if req.Body != nil {
|
|
u.lastBody, _ = io.ReadAll(req.Body)
|
|
}
|
|
return u.resp, nil
|
|
}
|
|
|
|
func (u *grokQuotaHandlerUpstream) DoWithTLS(
|
|
req *http.Request,
|
|
proxyURL string,
|
|
accountID int64,
|
|
accountConcurrency int,
|
|
_ *tlsfingerprint.Profile,
|
|
) (*http.Response, error) {
|
|
return u.Do(req, proxyURL, accountID, accountConcurrency)
|
|
}
|
|
|
|
func TestGrokOAuthHandlerQueryQuotaProbesUpstream(t *testing.T) {
|
|
gin.SetMode(gin.TestMode)
|
|
|
|
repo := &grokQuotaHandlerAccountRepo{account: &service.Account{
|
|
ID: 42,
|
|
Platform: service.PlatformGrok,
|
|
Type: service.AccountTypeOAuth,
|
|
Concurrency: 1,
|
|
Credentials: map[string]any{
|
|
"access_token": "access-token",
|
|
"expires_at": time.Now().Add(time.Hour).UTC().Format(time.RFC3339),
|
|
},
|
|
}}
|
|
upstream := &grokQuotaHandlerUpstream{resp: &http.Response{
|
|
StatusCode: http.StatusOK,
|
|
Header: http.Header{
|
|
"X-Ratelimit-Limit-Requests": []string{"10"},
|
|
"X-Ratelimit-Remaining-Requests": []string{"8"},
|
|
},
|
|
Body: io.NopCloser(strings.NewReader(`{"id":"resp_probe"}`)),
|
|
}}
|
|
quotaService := service.NewGrokQuotaService(repo, nil, service.NewGrokTokenProvider(repo, nil), upstream)
|
|
handler := NewGrokOAuthHandler(nil, nil, quotaService)
|
|
|
|
router := gin.New()
|
|
router.GET("/api/v1/admin/grok/accounts/:id/quota", handler.QueryQuota)
|
|
rec := httptest.NewRecorder()
|
|
req := httptest.NewRequest(http.MethodGet, "/api/v1/admin/grok/accounts/42/quota", nil)
|
|
router.ServeHTTP(rec, req)
|
|
|
|
require.Equal(t, http.StatusOK, rec.Code)
|
|
require.Contains(t, rec.Body.String(), `"source":"active_probe"`)
|
|
require.Contains(t, rec.Body.String(), `"headers_observed":true`)
|
|
require.NotContains(t, rec.Body.String(), "access-token")
|
|
require.Equal(t, xai.DefaultBaseURL+"/responses", upstream.lastReq.URL.String())
|
|
require.Equal(t, "Bearer access-token", upstream.lastReq.Header.Get("Authorization"))
|
|
require.Contains(t, string(upstream.lastBody), `"store":false`)
|
|
require.NotNil(t, repo.updates[42])
|
|
}
|
|
|
|
func TestGrokOAuthHandlerResetQuotaReturnsUnsupported(t *testing.T) {
|
|
gin.SetMode(gin.TestMode)
|
|
|
|
repo := &grokQuotaHandlerAccountRepo{account: &service.Account{
|
|
ID: 43,
|
|
Platform: service.PlatformGrok,
|
|
Type: service.AccountTypeOAuth,
|
|
}}
|
|
quotaService := service.NewGrokQuotaService(repo, nil, nil, nil)
|
|
handler := NewGrokOAuthHandler(nil, nil, quotaService)
|
|
|
|
router := gin.New()
|
|
router.POST("/api/v1/admin/grok/accounts/:id/reset-quota", handler.ResetQuota)
|
|
rec := httptest.NewRecorder()
|
|
req := httptest.NewRequest(http.MethodPost, "/api/v1/admin/grok/accounts/43/reset-quota", nil)
|
|
router.ServeHTTP(rec, req)
|
|
|
|
require.Equal(t, http.StatusNotImplemented, rec.Code)
|
|
require.Contains(t, rec.Body.String(), `"reason":"GROK_QUOTA_RESET_UNSUPPORTED"`)
|
|
require.NotContains(t, rec.Body.String(), "access-token")
|
|
}
|
|
|
|
func TestGrokOAuthHandlerRuntimeSanityDoesNotExposeSecrets(t *testing.T) {
|
|
gin.SetMode(gin.TestMode)
|
|
t.Setenv(xai.EnvBaseURL, "http://127.0.0.1:8080/v1?access_token=secret")
|
|
t.Setenv(xai.EnvClientID, "client-secret-like-value")
|
|
|
|
handler := NewGrokOAuthHandler(nil, nil, nil)
|
|
router := gin.New()
|
|
router.GET("/api/v1/admin/grok/runtime-sanity", handler.RuntimeSanity)
|
|
rec := httptest.NewRecorder()
|
|
req := httptest.NewRequest(http.MethodGet, "/api/v1/admin/grok/runtime-sanity", nil)
|
|
router.ServeHTTP(rec, req)
|
|
|
|
require.Equal(t, http.StatusOK, rec.Code)
|
|
require.Contains(t, rec.Body.String(), `"public_gateway_scope":"responses_only"`)
|
|
require.Contains(t, rec.Body.String(), `"valid":false`)
|
|
require.NotContains(t, rec.Body.String(), "access_token")
|
|
require.NotContains(t, rec.Body.String(), "secret")
|
|
require.NotContains(t, rec.Body.String(), "client-secret-like-value")
|
|
}
|