mirror of
https://github.com/Wei-Shaw/sub2api.git
synced 2026-09-24 16:05:44 +08:00
Merge remote-tracking branch 'origin/main' into feat/grok-sso-device-oauth
# Conflicts: # frontend/src/api/admin/grok.ts
This commit is contained in:
@@ -115,6 +115,10 @@ func (h *AccountHandler) ImportCodexSession(c *gin.Context) {
|
||||
response.BadRequest(c, "Invalid request: "+err.Error())
|
||||
return
|
||||
}
|
||||
if err := service.ValidateOpenAILongContextBillingExtra(service.PlatformOpenAI, req.Extra); err != nil {
|
||||
response.ErrorFrom(c, err)
|
||||
return
|
||||
}
|
||||
if req.Concurrency != nil && *req.Concurrency < 0 {
|
||||
response.BadRequest(c, "concurrency must be >= 0")
|
||||
return
|
||||
|
||||
@@ -630,6 +630,7 @@ func TestImportCodexSessionsAccessTokenOnlySameUserUpdatesExisting(t *testing.T)
|
||||
"chatgpt_user_id": "user-1",
|
||||
"access_token": existingToken,
|
||||
},
|
||||
Extra: map[string]any{"openai_long_context_billing_enabled": false},
|
||||
}})
|
||||
handler := NewAccountHandler(svc, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil)
|
||||
req := CodexSessionImportRequest{SkipDefaultGroupBind: boolPtr(true)}
|
||||
@@ -650,6 +651,9 @@ func TestImportCodexSessionsAccessTokenOnlySameUserUpdatesExisting(t *testing.T)
|
||||
if len(svc.updatedAccounts) != 1 || svc.updatedAccounts[0].id != 10 {
|
||||
t.Fatalf("updated accounts = %+v, want account 10", svc.updatedAccounts)
|
||||
}
|
||||
if got := svc.updatedAccounts[0].input.Extra["openai_long_context_billing_enabled"]; got != false {
|
||||
t.Fatalf("openai_long_context_billing_enabled = %v, want false", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestImportCodexSessionsUpgradesAccessTokenOnlyAccountWithRefreshToken(t *testing.T) {
|
||||
|
||||
@@ -784,6 +784,10 @@ func (h *AccountHandler) Create(c *gin.Context) {
|
||||
response.BadRequest(c, "Invalid request: "+err.Error())
|
||||
return
|
||||
}
|
||||
if err := service.ValidateOpenAILongContextBillingExtra(req.Platform, req.Extra); err != nil {
|
||||
response.ErrorFrom(c, err)
|
||||
return
|
||||
}
|
||||
if req.RateMultiplier != nil && *req.RateMultiplier < 0 {
|
||||
response.BadRequest(c, "rate_multiplier must be >= 0")
|
||||
return
|
||||
@@ -1299,6 +1303,10 @@ func (h *AccountHandler) ApplyOAuthCredentials(c *gin.Context) {
|
||||
response.ErrorFrom(c, infraerrors.BadRequest("NOT_OAUTH", "cannot apply oauth credentials to non-OAuth account"))
|
||||
return
|
||||
}
|
||||
if err := service.ValidateOpenAILongContextBillingExtra(existing.Platform, req.Extra); err != nil {
|
||||
response.ErrorFrom(c, err)
|
||||
return
|
||||
}
|
||||
|
||||
updatedAccount, err := h.adminService.UpdateAccount(ctx, accountID, &service.UpdateAccountInput{
|
||||
Type: req.Type,
|
||||
@@ -1592,6 +1600,12 @@ func (h *AccountHandler) BatchCreate(c *gin.Context) {
|
||||
response.BadRequest(c, "Invalid request: "+err.Error())
|
||||
return
|
||||
}
|
||||
for _, item := range req.Accounts {
|
||||
if err := service.ValidateOpenAILongContextBillingExtra(item.Platform, item.Extra); err != nil {
|
||||
response.ErrorFrom(c, err)
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
executeAdminIdempotentJSON(c, "admin.accounts.batch_create", req, service.DefaultWriteIdempotencyTTL(), func(ctx context.Context) (any, error) {
|
||||
success := 0
|
||||
|
||||
@@ -0,0 +1,165 @@
|
||||
package admin
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"testing"
|
||||
|
||||
infraerrors "github.com/Wei-Shaw/sub2api/internal/pkg/errors"
|
||||
"github.com/Wei-Shaw/sub2api/internal/service"
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
func TestAccountAdminBoundariesRejectMalformedOpenAILongContextBillingValue(t *testing.T) {
|
||||
const malformedExtra = `"extra":{"openai_long_context_billing_enabled":"true"}`
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
method string
|
||||
path string
|
||||
body string
|
||||
mount func(*gin.Engine, *AccountHandler)
|
||||
setup func(*stubAdminService)
|
||||
}{
|
||||
{
|
||||
name: "create",
|
||||
method: http.MethodPost,
|
||||
path: "/accounts",
|
||||
body: `{"name":"account","platform":"openai","type":"apikey","credentials":{"api_key":"test"},` + malformedExtra + `}`,
|
||||
mount: func(router *gin.Engine, handler *AccountHandler) { router.POST("/accounts", handler.Create) },
|
||||
},
|
||||
{
|
||||
name: "update",
|
||||
method: http.MethodPut,
|
||||
path: "/accounts/1",
|
||||
body: `{` + malformedExtra + `}`,
|
||||
mount: func(router *gin.Engine, handler *AccountHandler) { router.PUT("/accounts/:id", handler.Update) },
|
||||
setup: func(stub *stubAdminService) {
|
||||
stub.updateAccountErr = infraerrors.BadRequest("OPENAI_LONG_CONTEXT_BILLING_INVALID", "invalid")
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "bulk update",
|
||||
method: http.MethodPost,
|
||||
path: "/accounts/bulk-update",
|
||||
body: `{"account_ids":[1],` + malformedExtra + `}`,
|
||||
mount: func(router *gin.Engine, handler *AccountHandler) {
|
||||
router.POST("/accounts/bulk-update", handler.BulkUpdate)
|
||||
},
|
||||
setup: func(stub *stubAdminService) {
|
||||
stub.bulkUpdateAccountErr = infraerrors.BadRequest("OPENAI_LONG_CONTEXT_BILLING_INVALID", "invalid")
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "batch create",
|
||||
method: http.MethodPost,
|
||||
path: "/accounts/batch",
|
||||
body: `{"accounts":[{"name":"account","platform":"openai","type":"apikey","credentials":{"api_key":"test"},` + malformedExtra + `}]}`,
|
||||
mount: func(router *gin.Engine, handler *AccountHandler) { router.POST("/accounts/batch", handler.BatchCreate) },
|
||||
},
|
||||
{
|
||||
name: "Codex session import",
|
||||
method: http.MethodPost,
|
||||
path: "/accounts/import-codex-session",
|
||||
body: `{"content":"token",` + malformedExtra + `}`,
|
||||
mount: func(router *gin.Engine, handler *AccountHandler) {
|
||||
router.POST("/accounts/import-codex-session", handler.ImportCodexSession)
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
stub := newStubAdminService()
|
||||
if tt.setup != nil {
|
||||
tt.setup(stub)
|
||||
}
|
||||
handler := NewAccountHandler(stub, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil)
|
||||
router := gin.New()
|
||||
tt.mount(router, handler)
|
||||
recorder := httptest.NewRecorder()
|
||||
request := httptest.NewRequest(tt.method, tt.path, bytes.NewBufferString(tt.body))
|
||||
request.Header.Set("Content-Type", "application/json")
|
||||
|
||||
router.ServeHTTP(recorder, request)
|
||||
|
||||
require.Equal(t, http.StatusBadRequest, recorder.Code)
|
||||
var responseBody struct {
|
||||
Reason string `json:"reason"`
|
||||
}
|
||||
require.NoError(t, json.Unmarshal(recorder.Body.Bytes(), &responseBody))
|
||||
require.Equal(t, "OPENAI_LONG_CONTEXT_BILLING_INVALID", responseBody.Reason)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestAccountCreateBoundaryDoesNotApplyOpenAIValidationToOtherPlatforms(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
handler := NewAccountHandler(newStubAdminService(), nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil)
|
||||
router := gin.New()
|
||||
router.POST("/accounts", handler.Create)
|
||||
recorder := httptest.NewRecorder()
|
||||
request := httptest.NewRequest(http.MethodPost, "/accounts", bytes.NewBufferString(
|
||||
`{"name":"account","platform":"anthropic","type":"apikey","credentials":{"api_key":"test"},"extra":{"openai_long_context_billing_enabled":"provider-owned"}}`,
|
||||
))
|
||||
request.Header.Set("Content-Type", "application/json")
|
||||
|
||||
router.ServeHTTP(recorder, request)
|
||||
|
||||
require.Equal(t, http.StatusOK, recorder.Code)
|
||||
}
|
||||
|
||||
func TestApplyOAuthCredentialsRejectsMalformedOpenAILongContextBillingBeforeMutation(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
stub := newStubAdminService()
|
||||
stub.getAccountResult = &service.Account{
|
||||
ID: 1,
|
||||
Platform: service.PlatformOpenAI,
|
||||
Type: service.AccountTypeOAuth,
|
||||
}
|
||||
handler := NewAccountHandler(stub, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil)
|
||||
router := gin.New()
|
||||
router.POST("/accounts/:id/apply-oauth-credentials", handler.ApplyOAuthCredentials)
|
||||
recorder := httptest.NewRecorder()
|
||||
request := httptest.NewRequest(http.MethodPost, "/accounts/1/apply-oauth-credentials", bytes.NewBufferString(
|
||||
`{"type":"oauth","credentials":{"access_token":"new-token"},"extra":{"openai_long_context_billing_enabled":"true"}}`,
|
||||
))
|
||||
request.Header.Set("Content-Type", "application/json")
|
||||
|
||||
router.ServeHTTP(recorder, request)
|
||||
|
||||
require.Equal(t, http.StatusBadRequest, recorder.Code)
|
||||
var responseBody struct {
|
||||
Reason string `json:"reason"`
|
||||
}
|
||||
require.NoError(t, json.Unmarshal(recorder.Body.Bytes(), &responseBody))
|
||||
require.Equal(t, "OPENAI_LONG_CONTEXT_BILLING_INVALID", responseBody.Reason)
|
||||
require.Zero(t, stub.updateAccountCalls)
|
||||
require.Zero(t, stub.updateAccountExtraCalls)
|
||||
}
|
||||
|
||||
func TestOpenAIOAuthCodexPATBoundaryRejectsMalformedOpenAILongContextBillingValueBeforeTokenValidation(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
handler := NewOpenAIOAuthHandler(nil, newStubAdminService(), nil)
|
||||
router := gin.New()
|
||||
router.Use(gin.Recovery())
|
||||
router.POST("/openai/create-from-codex-pat", handler.CreateAccountFromCodexPAT)
|
||||
recorder := httptest.NewRecorder()
|
||||
request := httptest.NewRequest(http.MethodPost, "/openai/create-from-codex-pat", bytes.NewBufferString(
|
||||
`{"access_token":"token","extra":{"openai_long_context_billing_enabled":1}}`,
|
||||
))
|
||||
request.Header.Set("Content-Type", "application/json")
|
||||
|
||||
router.ServeHTTP(recorder, request)
|
||||
|
||||
require.Equal(t, http.StatusBadRequest, recorder.Code)
|
||||
var responseBody struct {
|
||||
Reason string `json:"reason"`
|
||||
}
|
||||
require.NoError(t, json.Unmarshal(recorder.Body.Bytes(), &responseBody))
|
||||
require.Equal(t, "OPENAI_LONG_CONTEXT_BILLING_INVALID", responseBody.Reason)
|
||||
}
|
||||
@@ -33,6 +33,9 @@ type stubAdminService struct {
|
||||
createSparkShadowErr error
|
||||
updateAccountErr error
|
||||
bulkUpdateAccountErr error
|
||||
getAccountResult *service.Account
|
||||
updateAccountCalls int
|
||||
updateAccountExtraCalls int
|
||||
checkMixedErr error
|
||||
lastMixedCheck struct {
|
||||
accountID int64
|
||||
@@ -388,6 +391,9 @@ func (s *stubAdminService) ListOpenAISchedulableAccountsForSchedulerScore(_ cont
|
||||
}
|
||||
|
||||
func (s *stubAdminService) GetAccount(ctx context.Context, id int64) (*service.Account, error) {
|
||||
if s.getAccountResult != nil {
|
||||
return s.getAccountResult, nil
|
||||
}
|
||||
account := service.Account{ID: id, Name: "account", Status: service.StatusActive}
|
||||
return &account, nil
|
||||
}
|
||||
@@ -413,6 +419,7 @@ func (s *stubAdminService) CreateAccount(ctx context.Context, input *service.Cre
|
||||
}
|
||||
|
||||
func (s *stubAdminService) UpdateAccount(ctx context.Context, id int64, input *service.UpdateAccountInput) (*service.Account, error) {
|
||||
s.updateAccountCalls++
|
||||
if s.updateAccountErr != nil {
|
||||
return nil, s.updateAccountErr
|
||||
}
|
||||
@@ -421,6 +428,7 @@ func (s *stubAdminService) UpdateAccount(ctx context.Context, id int64, input *s
|
||||
}
|
||||
|
||||
func (s *stubAdminService) UpdateAccountExtra(ctx context.Context, id int64, updates map[string]any) error {
|
||||
s.updateAccountExtraCalls++
|
||||
return nil
|
||||
}
|
||||
|
||||
|
||||
@@ -454,7 +454,7 @@ func (h *GrokOAuthHandler) QueryQuota(c *gin.Context) {
|
||||
response.BadRequest(c, "grok quota service is not enabled")
|
||||
return
|
||||
}
|
||||
result, err := h.quotaService.ProbeUsage(c.Request.Context(), accountID)
|
||||
result, err := h.quotaService.QueryQuota(c.Request.Context(), accountID)
|
||||
if err != nil {
|
||||
response.ErrorFrom(c, err)
|
||||
return
|
||||
|
||||
@@ -8,6 +8,7 @@ import (
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"strings"
|
||||
"sync"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
@@ -41,17 +42,35 @@ func (r *grokQuotaHandlerAccountRepo) UpdateExtra(_ context.Context, id int64, u
|
||||
}
|
||||
|
||||
type grokQuotaHandlerUpstream struct {
|
||||
resp *http.Response
|
||||
lastReq *http.Request
|
||||
lastBody []byte
|
||||
mu sync.Mutex
|
||||
requests []*http.Request
|
||||
bodies [][]byte
|
||||
}
|
||||
|
||||
func (u *grokQuotaHandlerUpstream) Do(req *http.Request, _ string, _ int64, _ int) (*http.Response, error) {
|
||||
u.lastReq = req
|
||||
var body []byte
|
||||
if req.Body != nil {
|
||||
u.lastBody, _ = io.ReadAll(req.Body)
|
||||
body, _ = io.ReadAll(req.Body)
|
||||
}
|
||||
return u.resp, nil
|
||||
u.mu.Lock()
|
||||
u.requests = append(u.requests, req)
|
||||
u.bodies = append(u.bodies, body)
|
||||
u.mu.Unlock()
|
||||
if req.URL.Path == "/v1/responses" {
|
||||
return &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"}`)),
|
||||
}, nil
|
||||
}
|
||||
payload := `{"config":{"billingPeriodStart":"2026-07-01T00:00:00Z","billingPeriodEnd":"2026-08-01T00:00:00Z"}}`
|
||||
if req.URL.RawQuery == "format=credits" {
|
||||
payload = `{"config":{"currentPeriod":{"type":"WEEKLY","start":"2026-07-09T03:25:00Z","end":"2026-07-16T03:25:00Z"}}}`
|
||||
}
|
||||
return &http.Response{StatusCode: http.StatusOK, Header: make(http.Header), Body: io.NopCloser(strings.NewReader(payload))}, nil
|
||||
}
|
||||
|
||||
func (u *grokQuotaHandlerUpstream) DoWithTLS(
|
||||
@@ -77,14 +96,7 @@ func TestGrokOAuthHandlerQueryQuotaProbesUpstream(t *testing.T) {
|
||||
"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"}`)),
|
||||
}}
|
||||
upstream := &grokQuotaHandlerUpstream{}
|
||||
quotaService := service.NewGrokQuotaService(repo, nil, service.NewGrokTokenProvider(repo, nil), upstream)
|
||||
handler := NewGrokOAuthHandler(nil, nil, quotaService)
|
||||
|
||||
@@ -95,12 +107,23 @@ func TestGrokOAuthHandlerQueryQuotaProbesUpstream(t *testing.T) {
|
||||
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(), `"source":"hybrid_probe"`)
|
||||
require.Contains(t, rec.Body.String(), `"billing":`)
|
||||
require.Contains(t, rec.Body.String(), `"snapshot":`)
|
||||
require.Contains(t, rec.Body.String(), `"headers_observed":true`)
|
||||
require.NotContains(t, rec.Body.String(), "access-token")
|
||||
require.Equal(t, xai.DefaultCLIBaseURL+"/responses", upstream.lastReq.URL.String())
|
||||
require.Equal(t, "Bearer access-token", upstream.lastReq.Header.Get("Authorization"))
|
||||
require.Contains(t, string(upstream.lastBody), `"store":false`)
|
||||
upstream.mu.Lock()
|
||||
requests := append([]*http.Request(nil), upstream.requests...)
|
||||
bodies := append([][]byte(nil), upstream.bodies...)
|
||||
upstream.mu.Unlock()
|
||||
require.Len(t, requests, 3)
|
||||
for i, upstreamReq := range requests {
|
||||
require.Equal(t, "Bearer access-token", upstreamReq.Header.Get("Authorization"))
|
||||
if upstreamReq.URL.String() == xai.DefaultCLIBaseURL+"/responses" {
|
||||
require.Contains(t, string(bodies[i]), `"model":"grok-4.5"`)
|
||||
require.Contains(t, string(bodies[i]), `"store":false`)
|
||||
}
|
||||
}
|
||||
require.NotNil(t, repo.updates[42])
|
||||
}
|
||||
|
||||
|
||||
@@ -304,6 +304,10 @@ func (h *OpenAIOAuthHandler) CreateAccountFromCodexPAT(c *gin.Context) {
|
||||
response.BadRequest(c, "Invalid request: "+err.Error())
|
||||
return
|
||||
}
|
||||
if err := service.ValidateOpenAILongContextBillingExtra(service.PlatformOpenAI, req.Extra); err != nil {
|
||||
response.ErrorFrom(c, err)
|
||||
return
|
||||
}
|
||||
if req.Concurrency != nil && *req.Concurrency < 0 {
|
||||
response.BadRequest(c, "concurrency must be >= 0")
|
||||
return
|
||||
|
||||
@@ -15,6 +15,7 @@ import (
|
||||
type opsSystemLogCleanupRequest struct {
|
||||
StartTime string `json:"start_time"`
|
||||
EndTime string `json:"end_time"`
|
||||
Host string `json:"host"`
|
||||
|
||||
Level string `json:"level"`
|
||||
Component string `json:"component"`
|
||||
@@ -56,6 +57,7 @@ func (h *OpsHandler) ListSystemLogs(c *gin.Context) {
|
||||
PageSize: pageSize,
|
||||
StartTime: &start,
|
||||
EndTime: &end,
|
||||
Host: strings.TrimSpace(c.Query("host")),
|
||||
Level: strings.TrimSpace(c.Query("level")),
|
||||
Component: strings.TrimSpace(c.Query("component")),
|
||||
RequestID: strings.TrimSpace(c.Query("request_id")),
|
||||
@@ -153,6 +155,7 @@ func (h *OpsHandler) CleanupSystemLogs(c *gin.Context) {
|
||||
filter := &service.OpsSystemLogCleanupFilter{
|
||||
StartTime: start,
|
||||
EndTime: end,
|
||||
Host: strings.TrimSpace(req.Host),
|
||||
Level: strings.TrimSpace(req.Level),
|
||||
Component: strings.TrimSpace(req.Component),
|
||||
RequestID: strings.TrimSpace(req.RequestID),
|
||||
|
||||
@@ -2,6 +2,7 @@ package admin
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
@@ -19,6 +20,26 @@ type responseEnvelope struct {
|
||||
Data json.RawMessage `json:"data"`
|
||||
}
|
||||
|
||||
type opsSystemLogCaptureRepo struct {
|
||||
service.OpsRepository
|
||||
listFilter *service.OpsSystemLogFilter
|
||||
cleanupFilter *service.OpsSystemLogCleanupFilter
|
||||
}
|
||||
|
||||
func (r *opsSystemLogCaptureRepo) ListSystemLogs(_ context.Context, filter *service.OpsSystemLogFilter) (*service.OpsSystemLogList, error) {
|
||||
r.listFilter = filter
|
||||
return &service.OpsSystemLogList{Logs: []*service.OpsSystemLog{}, Page: filter.Page, PageSize: filter.PageSize}, nil
|
||||
}
|
||||
|
||||
func (r *opsSystemLogCaptureRepo) DeleteSystemLogs(_ context.Context, filter *service.OpsSystemLogCleanupFilter) (int64, error) {
|
||||
r.cleanupFilter = filter
|
||||
return 1, nil
|
||||
}
|
||||
|
||||
func (r *opsSystemLogCaptureRepo) InsertSystemLogCleanupAudit(_ context.Context, _ *service.OpsSystemLogCleanupAudit) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
func newOpsSystemLogTestRouter(handler *OpsHandler, withUser bool) *gin.Engine {
|
||||
gin.SetMode(gin.TestMode)
|
||||
r := gin.New()
|
||||
@@ -121,6 +142,23 @@ func TestOpsSystemLogHandler_ListSuccess(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestOpsSystemLogHandler_ListAcceptsHost(t *testing.T) {
|
||||
repo := &opsSystemLogCaptureRepo{}
|
||||
svc := service.NewOpsService(repo, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil)
|
||||
h := NewOpsHandler(svc)
|
||||
r := newOpsSystemLogTestRouter(h, false)
|
||||
|
||||
w := httptest.NewRecorder()
|
||||
req := httptest.NewRequest(http.MethodGet, "/logs?host=api-node-1", nil)
|
||||
r.ServeHTTP(w, req)
|
||||
if w.Code != http.StatusOK {
|
||||
t.Fatalf("status=%d, want 200", w.Code)
|
||||
}
|
||||
if repo.listFilter == nil || repo.listFilter.Host != "api-node-1" {
|
||||
t.Fatalf("host filter = %+v, want api-node-1", repo.listFilter)
|
||||
}
|
||||
}
|
||||
|
||||
func TestOpsSystemLogHandler_CleanupUnauthorized(t *testing.T) {
|
||||
svc := service.NewOpsService(nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil)
|
||||
h := NewOpsHandler(svc)
|
||||
@@ -205,6 +243,24 @@ func TestOpsSystemLogHandler_CleanupAcceptsAPIKeyID(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestOpsSystemLogHandler_CleanupAcceptsHost(t *testing.T) {
|
||||
repo := &opsSystemLogCaptureRepo{}
|
||||
svc := service.NewOpsService(repo, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil)
|
||||
h := NewOpsHandler(svc)
|
||||
r := newOpsSystemLogTestRouter(h, true)
|
||||
|
||||
w := httptest.NewRecorder()
|
||||
req := httptest.NewRequest(http.MethodPost, "/logs/cleanup", bytes.NewBufferString(`{"host":"api-node-1"}`))
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
r.ServeHTTP(w, req)
|
||||
if w.Code != http.StatusOK {
|
||||
t.Fatalf("status=%d, want 200", w.Code)
|
||||
}
|
||||
if repo.cleanupFilter == nil || repo.cleanupFilter.Host != "api-node-1" {
|
||||
t.Fatalf("host filter = %+v, want api-node-1", repo.cleanupFilter)
|
||||
}
|
||||
}
|
||||
|
||||
func TestOpsSystemLogHandler_CleanupInvalidAPIKeyID(t *testing.T) {
|
||||
svc := service.NewOpsService(nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil)
|
||||
h := NewOpsHandler(svc)
|
||||
|
||||
@@ -16,6 +16,7 @@ import (
|
||||
"time"
|
||||
|
||||
"github.com/Wei-Shaw/sub2api/internal/pkg/logger"
|
||||
servermiddleware "github.com/Wei-Shaw/sub2api/internal/server/middleware"
|
||||
"github.com/Wei-Shaw/sub2api/internal/service"
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/gorilla/websocket"
|
||||
@@ -323,7 +324,7 @@ func (h *OpsHandler) QPSWSHandler(c *gin.Context) {
|
||||
// If realtime monitoring is disabled, prefer a successful WS upgrade followed by a clean close
|
||||
// with a deterministic close code. This prevents clients from spinning on 404/1006 reconnect loops.
|
||||
if !h.opsService.IsRealtimeMonitoringEnabled(c.Request.Context()) {
|
||||
conn, err := upgrader.Upgrade(c.Writer, c.Request, nil)
|
||||
conn, err := upgrader.Upgrade(c.Writer, c.Request, servermiddleware.ServerTimingResponseHeader(c))
|
||||
if err != nil {
|
||||
c.JSON(http.StatusNotFound, gin.H{"error": "ops realtime monitoring is disabled"})
|
||||
return
|
||||
@@ -358,7 +359,7 @@ func (h *OpsHandler) QPSWSHandler(c *gin.Context) {
|
||||
defer releaseOpsWSIPSlot(clientIP)
|
||||
}
|
||||
|
||||
conn, err := upgrader.Upgrade(c.Writer, c.Request, nil)
|
||||
conn, err := upgrader.Upgrade(c.Writer, c.Request, servermiddleware.ServerTimingResponseHeader(c))
|
||||
if err != nil {
|
||||
logger.LegacyPrintf("handler.admin.ops_ws", "[OpsWS] upgrade failed: %v", err)
|
||||
return
|
||||
|
||||
@@ -599,54 +599,55 @@ func usageLogFromServiceUser(l *service.UsageLog) UsageLog {
|
||||
requestedModel = l.Model
|
||||
}
|
||||
return UsageLog{
|
||||
ID: l.ID,
|
||||
UserID: l.UserID,
|
||||
APIKeyID: l.APIKeyID,
|
||||
AccountID: l.AccountID,
|
||||
RequestID: l.RequestID,
|
||||
Model: requestedModel,
|
||||
ServiceTier: l.ServiceTier,
|
||||
ReasoningEffort: l.ReasoningEffort,
|
||||
InboundEndpoint: l.InboundEndpoint,
|
||||
GroupID: l.GroupID,
|
||||
SubscriptionID: l.SubscriptionID,
|
||||
InputTokens: l.InputTokens,
|
||||
OutputTokens: l.OutputTokens,
|
||||
CacheCreationTokens: l.CacheCreationTokens,
|
||||
CacheReadTokens: l.CacheReadTokens,
|
||||
CacheCreation5mTokens: l.CacheCreation5mTokens,
|
||||
CacheCreation1hTokens: l.CacheCreation1hTokens,
|
||||
InputCost: l.InputCost,
|
||||
OutputCost: l.OutputCost,
|
||||
CacheCreationCost: l.CacheCreationCost,
|
||||
CacheReadCost: l.CacheReadCost,
|
||||
TotalCost: l.TotalCost,
|
||||
ActualCost: l.ActualCost,
|
||||
RateMultiplier: l.RateMultiplier,
|
||||
BillingType: l.BillingType,
|
||||
RequestType: requestType.String(),
|
||||
Stream: stream,
|
||||
OpenAIWSMode: openAIWSMode,
|
||||
DurationMs: l.DurationMs,
|
||||
FirstTokenMs: l.FirstTokenMs,
|
||||
ImageCount: l.ImageCount,
|
||||
ImageSize: l.ImageSize,
|
||||
ImageInputSize: l.ImageInputSize,
|
||||
ImageOutputSize: l.ImageOutputSize,
|
||||
ImageOutputTokens: l.ImageOutputTokens,
|
||||
ImageOutputCost: l.ImageOutputCost,
|
||||
ImageSizeSource: l.ImageSizeSource,
|
||||
ImageSizeBreakdown: l.ImageSizeBreakdown,
|
||||
MediaType: l.MediaType,
|
||||
UserAgent: l.UserAgent,
|
||||
IPAddress: l.IPAddress,
|
||||
CacheTTLOverridden: l.CacheTTLOverridden,
|
||||
BillingMode: l.BillingMode,
|
||||
CreatedAt: l.CreatedAt,
|
||||
User: UserFromServiceShallow(l.User),
|
||||
APIKey: APIKeyFromService(l.APIKey),
|
||||
Group: GroupFromServiceShallow(l.Group),
|
||||
Subscription: UserSubscriptionFromService(l.Subscription),
|
||||
ID: l.ID,
|
||||
UserID: l.UserID,
|
||||
APIKeyID: l.APIKeyID,
|
||||
AccountID: l.AccountID,
|
||||
RequestID: l.RequestID,
|
||||
Model: requestedModel,
|
||||
ServiceTier: l.ServiceTier,
|
||||
ReasoningEffort: l.ReasoningEffort,
|
||||
InboundEndpoint: l.InboundEndpoint,
|
||||
GroupID: l.GroupID,
|
||||
SubscriptionID: l.SubscriptionID,
|
||||
InputTokens: l.InputTokens,
|
||||
OutputTokens: l.OutputTokens,
|
||||
CacheCreationTokens: l.CacheCreationTokens,
|
||||
CacheReadTokens: l.CacheReadTokens,
|
||||
CacheCreation5mTokens: l.CacheCreation5mTokens,
|
||||
CacheCreation1hTokens: l.CacheCreation1hTokens,
|
||||
InputCost: l.InputCost,
|
||||
OutputCost: l.OutputCost,
|
||||
CacheCreationCost: l.CacheCreationCost,
|
||||
CacheReadCost: l.CacheReadCost,
|
||||
TotalCost: l.TotalCost,
|
||||
ActualCost: l.ActualCost,
|
||||
RateMultiplier: l.RateMultiplier,
|
||||
LongContextBillingApplied: l.LongContextBillingApplied,
|
||||
BillingType: l.BillingType,
|
||||
RequestType: requestType.String(),
|
||||
Stream: stream,
|
||||
OpenAIWSMode: openAIWSMode,
|
||||
DurationMs: l.DurationMs,
|
||||
FirstTokenMs: l.FirstTokenMs,
|
||||
ImageCount: l.ImageCount,
|
||||
ImageSize: l.ImageSize,
|
||||
ImageInputSize: l.ImageInputSize,
|
||||
ImageOutputSize: l.ImageOutputSize,
|
||||
ImageOutputTokens: l.ImageOutputTokens,
|
||||
ImageOutputCost: l.ImageOutputCost,
|
||||
ImageSizeSource: l.ImageSizeSource,
|
||||
ImageSizeBreakdown: l.ImageSizeBreakdown,
|
||||
MediaType: l.MediaType,
|
||||
UserAgent: l.UserAgent,
|
||||
IPAddress: l.IPAddress,
|
||||
CacheTTLOverridden: l.CacheTTLOverridden,
|
||||
BillingMode: l.BillingMode,
|
||||
CreatedAt: l.CreatedAt,
|
||||
User: UserFromServiceShallow(l.User),
|
||||
APIKey: APIKeyFromService(l.APIKey),
|
||||
Group: GroupFromServiceShallow(l.Group),
|
||||
Subscription: UserSubscriptionFromService(l.Subscription),
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -482,13 +482,14 @@ type UsageLog struct {
|
||||
CacheCreation5mTokens int `json:"cache_creation_5m_tokens"`
|
||||
CacheCreation1hTokens int `json:"cache_creation_1h_tokens"`
|
||||
|
||||
InputCost float64 `json:"input_cost"`
|
||||
OutputCost float64 `json:"output_cost"`
|
||||
CacheCreationCost float64 `json:"cache_creation_cost"`
|
||||
CacheReadCost float64 `json:"cache_read_cost"`
|
||||
TotalCost float64 `json:"total_cost"`
|
||||
ActualCost float64 `json:"actual_cost"`
|
||||
RateMultiplier float64 `json:"rate_multiplier"`
|
||||
InputCost float64 `json:"input_cost"`
|
||||
OutputCost float64 `json:"output_cost"`
|
||||
CacheCreationCost float64 `json:"cache_creation_cost"`
|
||||
CacheReadCost float64 `json:"cache_read_cost"`
|
||||
TotalCost float64 `json:"total_cost"`
|
||||
ActualCost float64 `json:"actual_cost"`
|
||||
RateMultiplier float64 `json:"rate_multiplier"`
|
||||
LongContextBillingApplied bool `json:"long_context_billing_applied"`
|
||||
|
||||
BillingType int8 `json:"billing_type"`
|
||||
RequestType string `json:"request_type"`
|
||||
|
||||
@@ -15,11 +15,13 @@ import (
|
||||
// Codex CLI and the Codex desktop app refresh their model picker from
|
||||
// GET {base_url}/models?client_version=... (custom provider mode) or
|
||||
// GET /backend-api/codex/models (chatgpt_base_url mode). Both routes land
|
||||
// here. The manifest is proxied verbatim from the ChatGPT backend with a
|
||||
// schedulable OAuth account's credentials, so clients pointed at the gateway
|
||||
// see the account's real, always-current model entitlements instead of a
|
||||
// frozen local cache.
|
||||
// here. The manifest is proxied verbatim from the selected account's ChatGPT
|
||||
// backend or custom API key upstream. API key manifests use a short-lived,
|
||||
// asynchronously revalidated cache to tolerate canceled client requests.
|
||||
func (h *OpenAIGatewayHandler) CodexModels(c *gin.Context) {
|
||||
if c.Request.Context().Err() != nil {
|
||||
return
|
||||
}
|
||||
apiKey, ok := middleware2.GetAPIKeyFromContext(c)
|
||||
if !ok || apiKey.Group == nil {
|
||||
h.errorResponse(c, http.StatusUnauthorized, "invalid_request_error", "API key group is required")
|
||||
@@ -30,24 +32,54 @@ func (h *OpenAIGatewayHandler) CodexModels(c *gin.Context) {
|
||||
return
|
||||
}
|
||||
|
||||
account, err := h.gatewayService.SelectAccountForModel(c.Request.Context(), apiKey.GroupID, "", "")
|
||||
if err != nil {
|
||||
h.errorResponse(c, http.StatusServiceUnavailable, "upstream_error", "No available OpenAI accounts")
|
||||
return
|
||||
maxAccountSwitches := h.maxAccountSwitches
|
||||
if maxAccountSwitches <= 0 {
|
||||
maxAccountSwitches = 3
|
||||
}
|
||||
failedAccountIDs := make(map[int64]struct{})
|
||||
switchCount := 0
|
||||
var lastUpstreamErr error
|
||||
|
||||
manifest, err := h.gatewayService.FetchCodexModelsManifest(c.Request.Context(), account, c.Query("client_version"), c.GetHeader("If-None-Match"))
|
||||
if err != nil {
|
||||
h.errorResponse(c, infraerrors.Code(err), "upstream_error", infraerrors.Message(err))
|
||||
return
|
||||
}
|
||||
for {
|
||||
account, err := h.gatewayService.SelectAccountForModelWithExclusions(c.Request.Context(), apiKey.GroupID, "", "", failedAccountIDs)
|
||||
if err != nil {
|
||||
if c.Request.Context().Err() != nil {
|
||||
return
|
||||
}
|
||||
if lastUpstreamErr != nil {
|
||||
h.errorResponse(c, infraerrors.Code(lastUpstreamErr), "upstream_error", infraerrors.Message(lastUpstreamErr))
|
||||
return
|
||||
}
|
||||
h.errorResponse(c, http.StatusServiceUnavailable, "upstream_error", "No available OpenAI accounts")
|
||||
return
|
||||
}
|
||||
|
||||
if manifest.ETag != "" {
|
||||
c.Header("ETag", manifest.ETag)
|
||||
}
|
||||
if manifest.NotModified {
|
||||
c.Status(http.StatusNotModified)
|
||||
manifest, err := h.gatewayService.FetchCodexModelsManifest(c.Request.Context(), account, c.Query("client_version"), c.GetHeader("If-None-Match"))
|
||||
if err != nil {
|
||||
if c.Request.Context().Err() != nil {
|
||||
return
|
||||
}
|
||||
if service.IsRetryableCodexModelsManifestError(err) && switchCount < maxAccountSwitches {
|
||||
failedAccountIDs[account.ID] = struct{}{}
|
||||
switchCount++
|
||||
lastUpstreamErr = err
|
||||
continue
|
||||
}
|
||||
h.errorResponse(c, infraerrors.Code(err), "upstream_error", infraerrors.Message(err))
|
||||
return
|
||||
}
|
||||
if c.Request.Context().Err() != nil {
|
||||
return
|
||||
}
|
||||
|
||||
if manifest.ETag != "" {
|
||||
c.Header("ETag", manifest.ETag)
|
||||
}
|
||||
if manifest.NotModified {
|
||||
c.Status(http.StatusNotModified)
|
||||
return
|
||||
}
|
||||
c.Data(http.StatusOK, "application/json", manifest.Body)
|
||||
return
|
||||
}
|
||||
c.Data(http.StatusOK, "application/json", manifest.Body)
|
||||
}
|
||||
|
||||
@@ -0,0 +1,288 @@
|
||||
package handler
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"net"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"strings"
|
||||
"sync"
|
||||
"testing"
|
||||
|
||||
"github.com/Wei-Shaw/sub2api/internal/config"
|
||||
middleware2 "github.com/Wei-Shaw/sub2api/internal/server/middleware"
|
||||
"github.com/Wei-Shaw/sub2api/internal/service"
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
|
||||
type codexModelsFailoverAccountRepo struct {
|
||||
service.AccountRepository
|
||||
accounts []service.Account
|
||||
}
|
||||
|
||||
func (r codexModelsFailoverAccountRepo) GetByID(_ context.Context, id int64) (*service.Account, error) {
|
||||
for i := range r.accounts {
|
||||
if r.accounts[i].ID == id {
|
||||
account := r.accounts[i]
|
||||
return &account, nil
|
||||
}
|
||||
}
|
||||
return nil, service.ErrNoAvailableAccounts
|
||||
}
|
||||
|
||||
func (r codexModelsFailoverAccountRepo) ListSchedulableByPlatform(_ context.Context, platform string) ([]service.Account, error) {
|
||||
accounts := make([]service.Account, 0, len(r.accounts))
|
||||
for _, account := range r.accounts {
|
||||
if account.Platform == platform {
|
||||
accounts = append(accounts, account)
|
||||
}
|
||||
}
|
||||
return accounts, nil
|
||||
}
|
||||
|
||||
type codexModelsFailoverHTTPUpstream struct {
|
||||
service.HTTPUpstream
|
||||
mu sync.Mutex
|
||||
accountIDs []int64
|
||||
firstErr error
|
||||
firstStatus int
|
||||
statuses map[int64]int
|
||||
}
|
||||
|
||||
func (u *codexModelsFailoverHTTPUpstream) Do(_ *http.Request, _ string, accountID int64, _ int) (*http.Response, error) {
|
||||
u.mu.Lock()
|
||||
u.accountIDs = append(u.accountIDs, accountID)
|
||||
u.mu.Unlock()
|
||||
|
||||
status, hasStatus := u.statuses[accountID]
|
||||
if accountID == 1 || hasStatus {
|
||||
if u.firstErr != nil {
|
||||
return nil, u.firstErr
|
||||
}
|
||||
if !hasStatus {
|
||||
status = u.firstStatus
|
||||
}
|
||||
if status == 0 {
|
||||
status = http.StatusServiceUnavailable
|
||||
}
|
||||
return &http.Response{
|
||||
StatusCode: status,
|
||||
Status: http.StatusText(status),
|
||||
Header: make(http.Header),
|
||||
Body: io.NopCloser(strings.NewReader(
|
||||
`{"error":{"message":"No available OpenAI accounts","type":"upstream_error"}}`,
|
||||
)),
|
||||
}, nil
|
||||
}
|
||||
return &http.Response{
|
||||
StatusCode: http.StatusOK,
|
||||
Status: "200 OK",
|
||||
Header: make(http.Header),
|
||||
Body: io.NopCloser(strings.NewReader(`{"models":[{"slug":"gpt-5.6-sol"}]}`)),
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (u *codexModelsFailoverHTTPUpstream) calls() []int64 {
|
||||
u.mu.Lock()
|
||||
defer u.mu.Unlock()
|
||||
return append([]int64(nil), u.accountIDs...)
|
||||
}
|
||||
|
||||
func TestCodexModelsCanceledRequestDoesNotWriteResponse(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
recorder := httptest.NewRecorder()
|
||||
c, _ := gin.CreateTestContext(recorder)
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
cancel()
|
||||
c.Request = httptest.NewRequest(http.MethodGet, "/v1/models", nil).WithContext(ctx)
|
||||
|
||||
h := &OpenAIGatewayHandler{}
|
||||
h.CodexModels(c)
|
||||
|
||||
if c.Writer.Written() {
|
||||
t.Fatalf("canceled request wrote an HTTP response: status=%d body=%q", recorder.Code, recorder.Body.String())
|
||||
}
|
||||
}
|
||||
|
||||
func TestCodexModelsFailsOverFromRetryableUpstreamStatus(t *testing.T) {
|
||||
retryableStatuses := []int{
|
||||
http.StatusTooManyRequests,
|
||||
http.StatusInternalServerError,
|
||||
http.StatusBadGateway,
|
||||
http.StatusServiceUnavailable,
|
||||
http.StatusGatewayTimeout,
|
||||
}
|
||||
for _, status := range retryableStatuses {
|
||||
t.Run(http.StatusText(status), func(t *testing.T) {
|
||||
handler, upstream, groupID := newCodexModelsFailoverTestHandler(status)
|
||||
recorder := performCodexModelsRequest(t, handler, groupID)
|
||||
|
||||
if got, want := upstream.calls(), []int64{1, 2}; !equalInt64Slices(got, want) {
|
||||
t.Fatalf("upstream account calls: got %v, want %v", got, want)
|
||||
}
|
||||
if recorder.Code != http.StatusOK {
|
||||
t.Fatalf("status: got %d, want %d; body=%s", recorder.Code, http.StatusOK, recorder.Body.String())
|
||||
}
|
||||
if got, want := recorder.Body.String(), `{"models":[{"slug":"gpt-5.6-sol"}]}`; got != want {
|
||||
t.Fatalf("body: got %q, want %q", got, want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestCodexModelsFailsOverFromUpstreamTransportError(t *testing.T) {
|
||||
handler, upstream, groupID := newCodexModelsFailoverTestHandler(http.StatusServiceUnavailable)
|
||||
upstream.firstErr = &net.OpError{
|
||||
Op: "read",
|
||||
Net: "tcp",
|
||||
Err: errors.New("connection reset"),
|
||||
}
|
||||
recorder := performCodexModelsRequest(t, handler, groupID)
|
||||
|
||||
if got, want := upstream.calls(), []int64{1, 2}; !equalInt64Slices(got, want) {
|
||||
t.Fatalf("upstream account calls: got %v, want %v", got, want)
|
||||
}
|
||||
if recorder.Code != http.StatusOK {
|
||||
t.Fatalf("status: got %d, want %d; body=%s", recorder.Code, http.StatusOK, recorder.Body.String())
|
||||
}
|
||||
}
|
||||
|
||||
func TestCodexModelsDoesNotFailOverFromPermanentUpstreamStatus(t *testing.T) {
|
||||
statuses := []int{
|
||||
http.StatusBadRequest,
|
||||
http.StatusUnauthorized,
|
||||
http.StatusForbidden,
|
||||
http.StatusNotFound,
|
||||
600,
|
||||
}
|
||||
for _, status := range statuses {
|
||||
t.Run(fmt.Sprintf("status_%d", status), func(t *testing.T) {
|
||||
handler, upstream, groupID := newCodexModelsFailoverTestHandler(status)
|
||||
recorder := performCodexModelsRequest(t, handler, groupID)
|
||||
|
||||
if got, want := upstream.calls(), []int64{1}; !equalInt64Slices(got, want) {
|
||||
t.Fatalf("upstream account calls: got %v, want %v", got, want)
|
||||
}
|
||||
if recorder.Code != http.StatusBadGateway {
|
||||
t.Fatalf("status: got %d, want %d; body=%s", recorder.Code, http.StatusBadGateway, recorder.Body.String())
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestCodexModelsDoesNotFailOverFromUpstreamConfigurationError(t *testing.T) {
|
||||
handler, upstream, groupID := newCodexModelsFailoverTestHandler(http.StatusServiceUnavailable)
|
||||
upstream.firstErr = errors.New("invalid proxy URL")
|
||||
recorder := performCodexModelsRequest(t, handler, groupID)
|
||||
|
||||
if got, want := upstream.calls(), []int64{1}; !equalInt64Slices(got, want) {
|
||||
t.Fatalf("upstream account calls: got %v, want %v", got, want)
|
||||
}
|
||||
if recorder.Code != http.StatusBadGateway {
|
||||
t.Fatalf("status: got %d, want %d; body=%s", recorder.Code, http.StatusBadGateway, recorder.Body.String())
|
||||
}
|
||||
}
|
||||
|
||||
func TestCodexModelsReturnsLastUpstreamErrorWhenAccountsAreExhausted(t *testing.T) {
|
||||
handler, upstream, groupID := newCodexModelsFailoverTestHandler(http.StatusServiceUnavailable)
|
||||
upstream.statuses = map[int64]int{
|
||||
1: http.StatusServiceUnavailable,
|
||||
2: http.StatusGatewayTimeout,
|
||||
}
|
||||
recorder := performCodexModelsRequest(t, handler, groupID)
|
||||
|
||||
if got, want := upstream.calls(), []int64{1, 2}; !equalInt64Slices(got, want) {
|
||||
t.Fatalf("upstream account calls: got %v, want %v", got, want)
|
||||
}
|
||||
if recorder.Code != http.StatusBadGateway {
|
||||
t.Fatalf("status: got %d, want %d; body=%s", recorder.Code, http.StatusBadGateway, recorder.Body.String())
|
||||
}
|
||||
if body := recorder.Body.String(); !strings.Contains(body, "upstream error 504") {
|
||||
t.Fatalf("body does not preserve the last upstream error: %s", body)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCodexModelsHonorsAccountSwitchLimit(t *testing.T) {
|
||||
handler, upstream, groupID := newCodexModelsFailoverTestHandlerWithAccountCount(http.StatusServiceUnavailable, 4, 2)
|
||||
upstream.statuses = map[int64]int{
|
||||
1: http.StatusServiceUnavailable,
|
||||
2: http.StatusBadGateway,
|
||||
3: http.StatusGatewayTimeout,
|
||||
4: http.StatusInternalServerError,
|
||||
}
|
||||
recorder := performCodexModelsRequest(t, handler, groupID)
|
||||
|
||||
if got, want := upstream.calls(), []int64{1, 2, 3}; !equalInt64Slices(got, want) {
|
||||
t.Fatalf("upstream account calls: got %v, want %v", got, want)
|
||||
}
|
||||
if recorder.Code != http.StatusBadGateway {
|
||||
t.Fatalf("status: got %d, want %d; body=%s", recorder.Code, http.StatusBadGateway, recorder.Body.String())
|
||||
}
|
||||
if body := recorder.Body.String(); !strings.Contains(body, "upstream error 504") {
|
||||
t.Fatalf("body does not preserve the limit-ending upstream error: %s", body)
|
||||
}
|
||||
}
|
||||
|
||||
func newCodexModelsFailoverTestHandler(firstStatus int) (*OpenAIGatewayHandler, *codexModelsFailoverHTTPUpstream, int64) {
|
||||
return newCodexModelsFailoverTestHandlerWithAccountCount(firstStatus, 2, 3)
|
||||
}
|
||||
|
||||
func newCodexModelsFailoverTestHandlerWithAccountCount(firstStatus, accountCount, maxSwitches int) (*OpenAIGatewayHandler, *codexModelsFailoverHTTPUpstream, int64) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
groupID := int64(42)
|
||||
accounts := make([]service.Account, 0, accountCount)
|
||||
for i := 1; i <= accountCount; i++ {
|
||||
accounts = append(accounts, service.Account{
|
||||
ID: int64(i),
|
||||
Name: fmt.Sprintf("upstream-%d", i),
|
||||
Platform: service.PlatformOpenAI,
|
||||
Type: service.AccountTypeAPIKey,
|
||||
Status: service.StatusActive,
|
||||
Schedulable: true,
|
||||
Priority: i - 1,
|
||||
Concurrency: 1,
|
||||
Credentials: map[string]any{
|
||||
"api_key": fmt.Sprintf("sk-%d", i),
|
||||
"base_url": fmt.Sprintf("https://upstream-%d.example/v1", i),
|
||||
},
|
||||
})
|
||||
}
|
||||
upstream := &codexModelsFailoverHTTPUpstream{firstStatus: firstStatus}
|
||||
cfg := &config.Config{RunMode: config.RunModeSimple}
|
||||
gatewayService := service.NewOpenAIGatewayService(
|
||||
codexModelsFailoverAccountRepo{accounts: accounts},
|
||||
nil, nil, nil, nil, nil, nil, cfg, nil, nil, nil, nil, nil,
|
||||
upstream,
|
||||
nil, nil, nil, nil, nil, nil, nil, nil,
|
||||
)
|
||||
return &OpenAIGatewayHandler{gatewayService: gatewayService, maxAccountSwitches: maxSwitches}, upstream, groupID
|
||||
}
|
||||
|
||||
func performCodexModelsRequest(t *testing.T, handler *OpenAIGatewayHandler, groupID int64) *httptest.ResponseRecorder {
|
||||
t.Helper()
|
||||
recorder := httptest.NewRecorder()
|
||||
c, _ := gin.CreateTestContext(recorder)
|
||||
c.Request = httptest.NewRequest(http.MethodGet, "/v1/models?client_version=0.144.0", nil)
|
||||
c.Set(string(middleware2.ContextKeyAPIKey), &service.APIKey{
|
||||
GroupID: &groupID,
|
||||
Group: &service.Group{ID: groupID, Platform: service.PlatformOpenAI},
|
||||
})
|
||||
|
||||
handler.CodexModels(c)
|
||||
return recorder
|
||||
}
|
||||
|
||||
func equalInt64Slices(got, want []int64) bool {
|
||||
if len(got) != len(want) {
|
||||
return false
|
||||
}
|
||||
for i := range got {
|
||||
if got[i] != want[i] {
|
||||
return false
|
||||
}
|
||||
}
|
||||
return true
|
||||
}
|
||||
@@ -415,11 +415,13 @@ func TestResolveOpenAIMessagesDispatchMappedModel(t *testing.T) {
|
||||
SonnetMappedModel: "gpt-5.2",
|
||||
ExactModelMappings: map[string]string{
|
||||
"claude-sonnet-4-5-20250929": "gpt-5.4-mini-high",
|
||||
"claude-fable-5": "gpt-5.6-sol",
|
||||
},
|
||||
},
|
||||
},
|
||||
}
|
||||
require.Equal(t, "gpt-5.4-mini", resolveOpenAIMessagesDispatchMappedModel(apiKey, "claude-sonnet-4-5-20250929"))
|
||||
require.Equal(t, "gpt-5.6-sol", resolveOpenAIMessagesDispatchMappedModel(apiKey, "claude-fable-5"))
|
||||
})
|
||||
|
||||
t.Run("uses_family_default_when_no_override", func(t *testing.T) {
|
||||
|
||||
Reference in New Issue
Block a user