mirror of
https://github.com/Wei-Shaw/sub2api.git
synced 2026-09-24 16:05:44 +08:00
feat: add grok quota probe parity
This commit is contained in:
@@ -13,12 +13,18 @@ import (
|
||||
type GrokOAuthHandler struct {
|
||||
grokOAuthService *service.GrokOAuthService
|
||||
adminService service.AdminService
|
||||
quotaService *service.GrokQuotaService
|
||||
}
|
||||
|
||||
func NewGrokOAuthHandler(grokOAuthService *service.GrokOAuthService, adminService service.AdminService) *GrokOAuthHandler {
|
||||
func NewGrokOAuthHandler(
|
||||
grokOAuthService *service.GrokOAuthService,
|
||||
adminService service.AdminService,
|
||||
quotaService *service.GrokQuotaService,
|
||||
) *GrokOAuthHandler {
|
||||
return &GrokOAuthHandler{
|
||||
grokOAuthService: grokOAuthService,
|
||||
adminService: adminService,
|
||||
quotaService: quotaService,
|
||||
}
|
||||
}
|
||||
|
||||
@@ -197,3 +203,39 @@ func (h *GrokOAuthHandler) CreateAccountFromOAuth(c *gin.Context) {
|
||||
}
|
||||
response.Success(c, dto.AccountFromService(account))
|
||||
}
|
||||
|
||||
func (h *GrokOAuthHandler) QueryQuota(c *gin.Context) {
|
||||
accountID, err := strconv.ParseInt(c.Param("id"), 10, 64)
|
||||
if err != nil {
|
||||
response.BadRequest(c, "Invalid account ID")
|
||||
return
|
||||
}
|
||||
if h.quotaService == nil {
|
||||
response.BadRequest(c, "grok quota service is not enabled")
|
||||
return
|
||||
}
|
||||
result, err := h.quotaService.ProbeUsage(c.Request.Context(), accountID)
|
||||
if err != nil {
|
||||
response.ErrorFrom(c, err)
|
||||
return
|
||||
}
|
||||
response.Success(c, result)
|
||||
}
|
||||
|
||||
func (h *GrokOAuthHandler) ResetQuota(c *gin.Context) {
|
||||
accountID, err := strconv.ParseInt(c.Param("id"), 10, 64)
|
||||
if err != nil {
|
||||
response.BadRequest(c, "Invalid account ID")
|
||||
return
|
||||
}
|
||||
if h.quotaService == nil {
|
||||
response.BadRequest(c, "grok quota service is not enabled")
|
||||
return
|
||||
}
|
||||
result, err := h.quotaService.ResetQuota(c.Request.Context(), accountID)
|
||||
if err != nil {
|
||||
response.ErrorFrom(c, err)
|
||||
return
|
||||
}
|
||||
response.Success(c, result)
|
||||
}
|
||||
|
||||
@@ -0,0 +1,127 @@
|
||||
//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, service.NewGrokTokenProvider(repo, nil, 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)
|
||||
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")
|
||||
}
|
||||
Reference in New Issue
Block a user