mirror of
https://github.com/Wei-Shaw/sub2api.git
synced 2026-10-07 15:08:02 +08:00
修复 PR 5888 剩余审计问题
This commit is contained in:
@@ -96,18 +96,12 @@ func (h *OpenAIGatewayHandler) ResponsesInputTokens(c *gin.Context) {
|
||||
requestPlatform := openAICompatibleRequestPlatform(c.Request.Context(), apiKey)
|
||||
sessionHash := h.gatewayService.GenerateSessionHash(c, body)
|
||||
requestStart := time.Now()
|
||||
selection, _, err := h.gatewayService.SelectAccountWithSchedulerForCapability(
|
||||
account, err := h.gatewayService.SelectAccountForTokenCount(
|
||||
c.Request.Context(),
|
||||
apiKey.GroupID,
|
||||
"",
|
||||
sessionHash,
|
||||
routingModel,
|
||||
nil,
|
||||
service.OpenAIUpstreamTransportAny,
|
||||
service.OpenAIEndpointCapabilityChatCompletions,
|
||||
false,
|
||||
false,
|
||||
false,
|
||||
requestPlatform,
|
||||
)
|
||||
service.SetOpsLatencyMs(c, service.OpsAuthLatencyMsKey, time.Since(requestStart).Milliseconds())
|
||||
@@ -120,7 +114,7 @@ func (h *OpenAIGatewayHandler) ResponsesInputTokens(c *gin.Context) {
|
||||
h.errorResponse(c, cls.Status, cls.ErrType, cls.Message)
|
||||
return
|
||||
}
|
||||
if selection == nil || selection.Account == nil {
|
||||
if account == nil {
|
||||
cls := classifyOpenAICompatibleNoAccountErrorFromGin(c, h.gatewayService, apiKey, routingModel, reqModel)
|
||||
if !cls.ModelNotFound {
|
||||
markOpsRoutingCapacityLimited(c)
|
||||
@@ -129,16 +123,7 @@ func (h *OpenAIGatewayHandler) ResponsesInputTokens(c *gin.Context) {
|
||||
return
|
||||
}
|
||||
|
||||
account := selection.Account
|
||||
setOpsSelectedAccount(c, account.ID, account.Platform)
|
||||
accountRelease, acquired := h.acquireCountTokensAccountSlot(c, apiKey.GroupID, sessionHash, selection, false, reqLog)
|
||||
if !acquired {
|
||||
return
|
||||
}
|
||||
if accountRelease != nil {
|
||||
defer accountRelease()
|
||||
}
|
||||
account = selection.Account
|
||||
if err := h.gatewayService.ForwardResponsesInputTokens(c.Request.Context(), c, account, forwardBody); err != nil {
|
||||
reqLog.Error("openai_input_tokens.forward_failed", zap.Int64("account_id", account.ID), zap.Error(err))
|
||||
}
|
||||
@@ -283,18 +268,12 @@ func (h *OpenAIGatewayHandler) CountTokens(c *gin.Context) {
|
||||
if preferredMappedModel != "" {
|
||||
currentRoutingModel = preferredMappedModel
|
||||
}
|
||||
selection, _, err := h.gatewayService.SelectAccountWithSchedulerForCapability(
|
||||
account, err := h.gatewayService.SelectAccountForTokenCount(
|
||||
c.Request.Context(),
|
||||
apiKey.GroupID,
|
||||
"",
|
||||
sessionHash,
|
||||
currentRoutingModel,
|
||||
nil,
|
||||
service.OpenAIUpstreamTransportAny,
|
||||
service.OpenAIEndpointCapabilityChatCompletions,
|
||||
false,
|
||||
false,
|
||||
false,
|
||||
openAICompatibleRequestPlatform(c.Request.Context(), apiKey),
|
||||
)
|
||||
service.SetOpsLatencyMs(c, service.OpsAuthLatencyMsKey, time.Since(requestStart).Milliseconds())
|
||||
@@ -308,7 +287,7 @@ func (h *OpenAIGatewayHandler) CountTokens(c *gin.Context) {
|
||||
h.anthropicErrorResponse(c, cls.Status, cls.ErrType, cls.Message)
|
||||
return
|
||||
}
|
||||
if selection == nil || selection.Account == nil {
|
||||
if account == nil {
|
||||
cls := classifyOpenAICompatibleNoAccountErrorFromGin(c, h.gatewayService, apiKey, currentRoutingModel, reqModel)
|
||||
if !cls.ModelNotFound {
|
||||
markOpsRoutingCapacityLimited(c)
|
||||
@@ -317,16 +296,7 @@ func (h *OpenAIGatewayHandler) CountTokens(c *gin.Context) {
|
||||
return
|
||||
}
|
||||
|
||||
account := selection.Account
|
||||
setOpsSelectedAccount(c, account.ID, account.Platform)
|
||||
accountRelease, acquired := h.acquireCountTokensAccountSlot(c, apiKey.GroupID, sessionHash, selection, true, reqLog)
|
||||
if !acquired {
|
||||
return
|
||||
}
|
||||
if accountRelease != nil {
|
||||
defer accountRelease()
|
||||
}
|
||||
account = selection.Account
|
||||
forwardBody := mappedBodyForMessages(channelMapping.Mapped, channelMapping.MappedModel)
|
||||
defaultMappedModel := preferredMappedModel
|
||||
|
||||
@@ -334,41 +304,3 @@ func (h *OpenAIGatewayHandler) CountTokens(c *gin.Context) {
|
||||
reqLog.Error("openai_count_tokens.forward_failed", zap.Int64("account_id", account.ID), zap.Error(err))
|
||||
}
|
||||
}
|
||||
|
||||
func (h *OpenAIGatewayHandler) acquireCountTokensAccountSlot(
|
||||
c *gin.Context,
|
||||
groupID *int64,
|
||||
sessionHash string,
|
||||
selection *service.AccountSelectionResult,
|
||||
anthropicResponse bool,
|
||||
reqLog *zap.Logger,
|
||||
) (func(), bool) {
|
||||
writeError := func(status int, errType, message string) {
|
||||
if anthropicResponse {
|
||||
h.anthropicErrorResponse(c, status, errType, message)
|
||||
return
|
||||
}
|
||||
h.errorResponse(c, status, errType, message)
|
||||
}
|
||||
streamStarted := false
|
||||
release, result := h.acquireOpenAIAccountSlot(
|
||||
c,
|
||||
groupID,
|
||||
sessionHash,
|
||||
selection,
|
||||
false,
|
||||
&streamStarted,
|
||||
reqLog,
|
||||
writeError,
|
||||
)
|
||||
if result == openAISlotAcquireOK {
|
||||
return release, true
|
||||
}
|
||||
// Token-count requests suppress the profit gate before selection, so this
|
||||
// is defensive only. Never forward without a slot if a stale gate appears.
|
||||
if result == openAISlotAcquireProfitVetoed {
|
||||
markOpsRoutingCapacityLimited(c)
|
||||
writeError(http.StatusServiceUnavailable, "api_error", "No available accounts")
|
||||
}
|
||||
return nil, false
|
||||
}
|
||||
|
||||
@@ -2466,6 +2466,10 @@ func (h *OpenAIGatewayHandler) ResponsesWebSocket(c *gin.Context) {
|
||||
|
||||
// WebSocket 首包可能很大,hash 必须在 hooks 外算成字符串,避免 AfterTurn 闭包保活请求体。
|
||||
requestPayloadHash = service.HashUsageRequestPayload(wsFirstMessage)
|
||||
if preemptCtx, cleanupPreempt, armed := h.gatewayService.BeginOpenAIWSIngressSessionPreemption(ctx, c, account, wsFirstMessage); armed {
|
||||
ctx = preemptCtx
|
||||
defer cleanupPreempt()
|
||||
}
|
||||
|
||||
for {
|
||||
err := h.gatewayService.ProxyResponsesWebSocketFromClient(ctx, c, wsConn, account, token, wsFirstMessage, hooks)
|
||||
|
||||
@@ -1,17 +1,12 @@
|
||||
package handler
|
||||
|
||||
import (
|
||||
"context"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"sync/atomic"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/Wei-Shaw/sub2api/internal/service"
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/stretchr/testify/require"
|
||||
"go.uber.org/zap"
|
||||
)
|
||||
|
||||
func TestResponsesInputTokensIsClassifiedAsTokenCountRequest(t *testing.T) {
|
||||
@@ -24,81 +19,3 @@ func TestResponsesInputTokensIsClassifiedAsTokenCountRequest(t *testing.T) {
|
||||
require.True(t, isTokenCountRequestPath("/responses/input_tokens"))
|
||||
require.False(t, isTokenCountRequestPath("/v1/responses"))
|
||||
}
|
||||
|
||||
func TestCountTokensAccountSlot_CancellationStopsBeforeForward(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
|
||||
for _, tt := range []struct {
|
||||
name string
|
||||
anthropic bool
|
||||
}{
|
||||
{name: "responses input tokens", anthropic: false},
|
||||
{name: "anthropic count tokens", anthropic: true},
|
||||
} {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
cache := &concurrencyCacheMock{
|
||||
acquireAccountSlotFn: func(context.Context, int64, int, string) (bool, error) {
|
||||
return false, nil
|
||||
},
|
||||
}
|
||||
h := &OpenAIGatewayHandler{
|
||||
gatewayService: &service.OpenAIGatewayService{},
|
||||
concurrencyHelper: NewConcurrencyHelper(service.NewConcurrencyService(cache), SSEPingFormatNone, time.Second),
|
||||
}
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
cancel()
|
||||
recorder := httptest.NewRecorder()
|
||||
c, _ := gin.CreateTestContext(recorder)
|
||||
c.Request = httptest.NewRequest(http.MethodPost, "/v1/count_tokens", nil).WithContext(ctx)
|
||||
groupID := int64(41)
|
||||
selection := &service.AccountSelectionResult{
|
||||
Account: &service.Account{ID: 42, Platform: service.PlatformOpenAI},
|
||||
WaitPlan: &service.AccountWaitPlan{
|
||||
AccountID: 42,
|
||||
MaxConcurrency: 1,
|
||||
MaxWaiting: 1,
|
||||
Timeout: time.Second,
|
||||
},
|
||||
}
|
||||
|
||||
release, acquired := h.acquireCountTokensAccountSlot(c, &groupID, "", selection, tt.anthropic, zap.NewNop())
|
||||
forwarded := false
|
||||
if acquired {
|
||||
forwarded = true
|
||||
if release != nil {
|
||||
release()
|
||||
}
|
||||
}
|
||||
|
||||
require.False(t, acquired)
|
||||
require.Nil(t, release)
|
||||
require.False(t, forwarded, "a canceled WaitPlan must stop before count-token forwarding")
|
||||
require.Zero(t, atomic.LoadInt32(&cache.releaseAccountCalled))
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestCountTokensAccountSlot_SelectionReleaseRunsExactlyOnce(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
|
||||
var released atomic.Int32
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
recorder := httptest.NewRecorder()
|
||||
c, _ := gin.CreateTestContext(recorder)
|
||||
c.Request = httptest.NewRequest(http.MethodPost, "/v1/responses/input_tokens", nil).WithContext(ctx)
|
||||
h := &OpenAIGatewayHandler{gatewayService: &service.OpenAIGatewayService{}}
|
||||
groupID := int64(51)
|
||||
selection := &service.AccountSelectionResult{
|
||||
Account: &service.Account{ID: 52, Platform: service.PlatformOpenAI},
|
||||
Acquired: true,
|
||||
ReleaseFunc: func() { released.Add(1) },
|
||||
}
|
||||
|
||||
release, acquired := h.acquireCountTokensAccountSlot(c, &groupID, "", selection, false, zap.NewNop())
|
||||
require.True(t, acquired)
|
||||
require.NotNil(t, release)
|
||||
release()
|
||||
cancel()
|
||||
require.Eventually(t, func() bool { return released.Load() == 1 }, time.Second, 10*time.Millisecond)
|
||||
require.Equal(t, int32(1), released.Load())
|
||||
}
|
||||
|
||||
@@ -176,17 +176,10 @@ func TestOpsCaptureWriter_ReleaseWaitsForDelegatedWriteWithoutHoldingStateMutex(
|
||||
}()
|
||||
<-inner.writeStarted
|
||||
|
||||
mutexAvailable := make(chan struct{})
|
||||
go func() {
|
||||
w.state.mu.Lock()
|
||||
w.state.mu.Unlock()
|
||||
close(mutexAvailable)
|
||||
}()
|
||||
select {
|
||||
case <-mutexAvailable:
|
||||
case <-time.After(time.Second):
|
||||
if !w.state.mu.TryLock() {
|
||||
t.Fatal("state mutex remained held across the delegated network write")
|
||||
}
|
||||
w.state.mu.Unlock()
|
||||
|
||||
releaseDone := make(chan struct{})
|
||||
go func() {
|
||||
|
||||
@@ -1280,6 +1280,18 @@ func logOpsRecoveredUpstream(c *gin.Context, ops *service.OpsService, finalStatu
|
||||
|
||||
entry := &service.OpsInsertErrorLogInput{StatusCode: finalStatus}
|
||||
applyOpsUpstreamFieldsFromContext(c, entry)
|
||||
if len(entry.UpstreamErrors) > 0 {
|
||||
visibleEvents := make([]*service.OpsUpstreamErrorEvent, 0, len(entry.UpstreamErrors))
|
||||
for _, event := range entry.UpstreamErrors {
|
||||
if event != nil && !event.SkipMonitoring {
|
||||
visibleEvents = append(visibleEvents, event)
|
||||
}
|
||||
}
|
||||
if len(visibleEvents) == 0 {
|
||||
return
|
||||
}
|
||||
applyOpsUpstreamErrorEvents(entry, visibleEvents)
|
||||
}
|
||||
if entry.UpstreamStatusCode == nil && entry.UpstreamErrorMessage == nil &&
|
||||
entry.UpstreamErrorDetail == nil && len(entry.UpstreamErrors) == 0 {
|
||||
return
|
||||
@@ -1350,7 +1362,7 @@ func logOpsRecoveredUpstream(c *gin.Context, ops *service.OpsService, finalStatu
|
||||
|
||||
apiKey := getOpsAPIKey(c)
|
||||
fallbackPlatform := guessPlatformFromPath(entry.RequestPath)
|
||||
var requestContext context.Context = context.Background()
|
||||
requestContext := context.Background()
|
||||
if c.Request != nil {
|
||||
requestContext = c.Request.Context()
|
||||
}
|
||||
@@ -1668,49 +1680,42 @@ func applyOpsUpstreamFieldsFromContext(c *gin.Context, entry *service.OpsInsertE
|
||||
}
|
||||
if v, ok := c.Get(service.OpsUpstreamErrorsKey); ok {
|
||||
if events, ok := v.([]*service.OpsUpstreamErrorEvent); ok && len(events) > 0 {
|
||||
entry.UpstreamErrors = events
|
||||
var last *service.OpsUpstreamErrorEvent
|
||||
for i := len(events) - 1; i >= 0; i-- {
|
||||
if events[i] != nil {
|
||||
last = events[i]
|
||||
break
|
||||
}
|
||||
}
|
||||
if last == nil {
|
||||
return
|
||||
}
|
||||
if last.Stage == string(service.GatewayFailureStageAccountAuth) {
|
||||
code := 0
|
||||
entry.UpstreamStatusCode = &code
|
||||
entry.UpstreamErrorMessage = nil
|
||||
if message := strings.TrimSpace(last.Message); message != "" {
|
||||
entry.UpstreamErrorMessage = &message
|
||||
}
|
||||
entry.UpstreamErrorDetail = nil
|
||||
if detail := strings.TrimSpace(last.Detail); detail != "" {
|
||||
entry.UpstreamErrorDetail = &detail
|
||||
}
|
||||
} else {
|
||||
entry.UpstreamStatusCode = nil
|
||||
if last.UpstreamStatusCode > 0 {
|
||||
code := last.UpstreamStatusCode
|
||||
entry.UpstreamStatusCode = &code
|
||||
}
|
||||
entry.UpstreamErrorMessage = nil
|
||||
if strings.TrimSpace(last.Message) != "" {
|
||||
message := strings.TrimSpace(last.Message)
|
||||
entry.UpstreamErrorMessage = &message
|
||||
}
|
||||
entry.UpstreamErrorDetail = nil
|
||||
if strings.TrimSpace(last.Detail) != "" {
|
||||
detail := strings.TrimSpace(last.Detail)
|
||||
entry.UpstreamErrorDetail = &detail
|
||||
}
|
||||
}
|
||||
applyOpsUpstreamErrorEvents(entry, events)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func applyOpsUpstreamErrorEvents(entry *service.OpsInsertErrorLogInput, events []*service.OpsUpstreamErrorEvent) {
|
||||
entry.UpstreamErrors = events
|
||||
var last *service.OpsUpstreamErrorEvent
|
||||
for i := len(events) - 1; i >= 0; i-- {
|
||||
if events[i] != nil {
|
||||
last = events[i]
|
||||
break
|
||||
}
|
||||
}
|
||||
if last == nil {
|
||||
return
|
||||
}
|
||||
|
||||
entry.UpstreamStatusCode = nil
|
||||
entry.UpstreamErrorMessage = nil
|
||||
entry.UpstreamErrorDetail = nil
|
||||
if last.Stage == string(service.GatewayFailureStageAccountAuth) {
|
||||
code := 0
|
||||
entry.UpstreamStatusCode = &code
|
||||
} else if last.UpstreamStatusCode > 0 {
|
||||
code := last.UpstreamStatusCode
|
||||
entry.UpstreamStatusCode = &code
|
||||
}
|
||||
if message := strings.TrimSpace(last.Message); message != "" {
|
||||
entry.UpstreamErrorMessage = &message
|
||||
}
|
||||
if detail := strings.TrimSpace(last.Detail); detail != "" {
|
||||
entry.UpstreamErrorDetail = &detail
|
||||
}
|
||||
}
|
||||
|
||||
func suppressOpsUpstreamAttributionForLocalModelConfiguration(c *gin.Context, entry *service.OpsInsertErrorLogInput) {
|
||||
if entry == nil || !service.HasOpsClientBusinessLimited(c) || service.OpsClientBusinessLimitedReason(c) != service.OpsClientBusinessLimitedReasonLocalModelConfiguration {
|
||||
return
|
||||
|
||||
@@ -375,6 +375,56 @@ func TestOpsErrorLoggerMiddleware_RecordsRecoveredUpstreamTelemetryOutsideFailur
|
||||
require.Equal(t, http.StatusTooManyRequests, persistedEvents[0].UpstreamStatusCode)
|
||||
}
|
||||
|
||||
func TestOpsErrorLoggerMiddleware_RecoveredTelemetryFiltersSkipMonitoringAttempts(t *testing.T) {
|
||||
setupOpsErrorLogTestQueue(t, 2)
|
||||
gin.SetMode(gin.TestMode)
|
||||
|
||||
ops := service.NewOpsService(nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil)
|
||||
router := gin.New()
|
||||
router.Use(OpsErrorLoggerMiddleware(ops))
|
||||
router.POST("/v1/responses", func(c *gin.Context) {
|
||||
c.Set(service.OpsUpstreamErrorsKey, []*service.OpsUpstreamErrorEvent{
|
||||
{UpstreamStatusCode: http.StatusTooManyRequests, Message: "visible retry"},
|
||||
{UpstreamStatusCode: http.StatusBadGateway, Message: "hidden retry", SkipMonitoring: true},
|
||||
})
|
||||
c.JSON(http.StatusOK, gin.H{"status": "completed"})
|
||||
})
|
||||
|
||||
recorder := httptest.NewRecorder()
|
||||
router.ServeHTTP(recorder, httptest.NewRequest(http.MethodPost, "/v1/responses", nil))
|
||||
|
||||
require.Equal(t, int64(1), OpsErrorLogQueueLength())
|
||||
job := <-opsErrorLogQueue
|
||||
require.Equal(t, "Recovered upstream error 429: visible retry", job.entry.ErrorMessage)
|
||||
require.NotNil(t, job.entry.UpstreamErrorsJSON)
|
||||
events, err := service.ParseOpsUpstreamErrors(*job.entry.UpstreamErrorsJSON)
|
||||
require.NoError(t, err)
|
||||
require.Len(t, events, 1)
|
||||
require.Equal(t, "visible retry", events[0].Message)
|
||||
}
|
||||
|
||||
func TestOpsErrorLoggerMiddleware_RecoveredTelemetrySkipsAllHiddenAttempts(t *testing.T) {
|
||||
setupOpsErrorLogTestQueue(t, 2)
|
||||
gin.SetMode(gin.TestMode)
|
||||
|
||||
ops := service.NewOpsService(nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil)
|
||||
router := gin.New()
|
||||
router.Use(OpsErrorLoggerMiddleware(ops))
|
||||
router.POST("/v1/responses", func(c *gin.Context) {
|
||||
c.Set(service.OpsUpstreamErrorsKey, []*service.OpsUpstreamErrorEvent{{
|
||||
UpstreamStatusCode: http.StatusTooManyRequests,
|
||||
Message: "hidden retry",
|
||||
SkipMonitoring: true,
|
||||
}})
|
||||
c.JSON(http.StatusOK, gin.H{"status": "completed"})
|
||||
})
|
||||
|
||||
recorder := httptest.NewRecorder()
|
||||
router.ServeHTTP(recorder, httptest.NewRequest(http.MethodPost, "/v1/responses", nil))
|
||||
|
||||
require.Equal(t, int64(0), OpsErrorLogQueueLength())
|
||||
}
|
||||
|
||||
func TestOpsErrorLoggerMiddleware_IntermediateSkipMonitoringDoesNotHideFinalVisibleFailure(t *testing.T) {
|
||||
setupOpsErrorLogTestQueue(t, 2)
|
||||
gin.SetMode(gin.TestMode)
|
||||
|
||||
@@ -67,7 +67,7 @@ func TestCheckErrorPolicy(t *testing.T) {
|
||||
expected: ErrorPolicySkipped,
|
||||
},
|
||||
{
|
||||
name: "global_529_bypasses_custom_error_code_filter",
|
||||
name: "custom_error_codes_excluding_529_skip_global_cooldown",
|
||||
account: &Account{
|
||||
ID: 33,
|
||||
Type: AccountTypeAPIKey,
|
||||
@@ -79,10 +79,10 @@ func TestCheckErrorPolicy(t *testing.T) {
|
||||
},
|
||||
statusCode: 529,
|
||||
body: []byte(`{"error":{"message":"overloaded"}}`),
|
||||
expected: ErrorPolicyMatched,
|
||||
expected: ErrorPolicySkipped,
|
||||
},
|
||||
{
|
||||
name: "global_529_bypasses_pool_mode",
|
||||
name: "pool_mode_skips_global_529_cooldown",
|
||||
account: &Account{
|
||||
ID: 34,
|
||||
Type: AccountTypeAPIKey,
|
||||
@@ -93,6 +93,32 @@ func TestCheckErrorPolicy(t *testing.T) {
|
||||
},
|
||||
statusCode: 529,
|
||||
body: []byte(`{"error":{"message":"overloaded"}}`),
|
||||
expected: ErrorPolicySkipped,
|
||||
},
|
||||
{
|
||||
name: "ordinary_account_uses_global_529_cooldown",
|
||||
account: &Account{
|
||||
ID: 35,
|
||||
Type: AccountTypeAPIKey,
|
||||
Platform: PlatformOpenAI,
|
||||
},
|
||||
statusCode: 529,
|
||||
body: []byte(`{"error":{"message":"overloaded"}}`),
|
||||
expected: ErrorPolicyMatched,
|
||||
},
|
||||
{
|
||||
name: "custom_error_codes_including_529_take_precedence",
|
||||
account: &Account{
|
||||
ID: 36,
|
||||
Type: AccountTypeAPIKey,
|
||||
Platform: PlatformOpenAI,
|
||||
Credentials: map[string]any{
|
||||
"custom_error_codes_enabled": true,
|
||||
"custom_error_codes": []any{float64(529)},
|
||||
},
|
||||
},
|
||||
statusCode: 529,
|
||||
body: []byte(`{"error":{"message":"overloaded"}}`),
|
||||
expected: ErrorPolicyMatched,
|
||||
},
|
||||
{
|
||||
|
||||
@@ -620,6 +620,59 @@ func TestOpenAIGatewayService_SelectAccountWithScheduler_DefaultDisabled_Embeddi
|
||||
require.Equal(t, openAIAccountScheduleLayerLoadBalance, decision.Layer)
|
||||
}
|
||||
|
||||
func TestOpenAIGatewayService_SelectAccountForTokenCount_DoesNotAcquireGenerationSlot(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
groupID := int64(10115)
|
||||
acquiredIDs := make([]int64, 0)
|
||||
accounts := []Account{
|
||||
{
|
||||
ID: 36501, Platform: PlatformOpenAI, Type: AccountTypeAPIKey,
|
||||
Status: StatusActive, Schedulable: true, Concurrency: 1, Priority: 0,
|
||||
Credentials: map[string]any{"openai_capabilities": []any{"chat_completions"}},
|
||||
},
|
||||
{
|
||||
ID: 36502, Platform: PlatformOpenAI, Type: AccountTypeAPIKey,
|
||||
Status: StatusActive, Schedulable: true, Concurrency: 1, Priority: 5,
|
||||
Credentials: map[string]any{"openai_capabilities": []any{"embeddings"}},
|
||||
},
|
||||
{
|
||||
ID: 36503, Platform: PlatformGrok, Type: AccountTypeAPIKey,
|
||||
Status: StatusActive, Schedulable: true, Concurrency: 1, Priority: 10,
|
||||
Credentials: map[string]any{"openai_capabilities": []any{"chat_completions"}},
|
||||
},
|
||||
{
|
||||
ID: 36504, Platform: PlatformOpenAI, Type: AccountTypeAPIKey,
|
||||
Status: StatusActive, Schedulable: true, Concurrency: 1, Priority: 15,
|
||||
Credentials: map[string]any{
|
||||
"openai_capabilities": []any{"chat_completions"},
|
||||
"model_mapping": map[string]any{"gpt-4o": "gpt-4o"},
|
||||
},
|
||||
},
|
||||
}
|
||||
svc := &OpenAIGatewayService{
|
||||
accountRepo: schedulerTestOpenAIAccountRepo{accounts: accounts},
|
||||
cache: &schedulerTestGatewayCache{},
|
||||
cfg: &config.Config{},
|
||||
concurrencyService: NewConcurrencyService(schedulerTestConcurrencyCache{
|
||||
acquireResults: map[int64]bool{36501: false},
|
||||
acquiredIDs: &acquiredIDs,
|
||||
}),
|
||||
}
|
||||
|
||||
account, err := svc.SelectAccountForTokenCount(
|
||||
ctx,
|
||||
&groupID,
|
||||
"",
|
||||
"gpt-5.1",
|
||||
OpenAIEndpointCapabilityChatCompletions,
|
||||
PlatformOpenAI,
|
||||
)
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, account)
|
||||
require.Equal(t, int64(36501), account.ID)
|
||||
require.Empty(t, acquiredIDs, "token counting must not acquire a generation slot")
|
||||
}
|
||||
|
||||
// 生图意图的 /v1/responses 请求要求 OpenAIEndpointCapabilityResponses:探测确认
|
||||
// 不支持 Responses API 的 APIKey 账号必须被排除,避免 forward 阶段降级为无法生图
|
||||
// 的 Chat Completions 直转(#4417)。
|
||||
|
||||
@@ -142,9 +142,11 @@ func TestApplyCodexOAuthTransform_NormalizesIsolatedLegacyReferenceAcrossTurns(t
|
||||
input, ok := reqBody["input"].([]any)
|
||||
require.True(t, ok)
|
||||
require.Len(t, input, 3)
|
||||
require.Equal(t, "fc_previous_turn", input[0].(map[string]any)["id"])
|
||||
require.Equal(t, "fc_remote_item", input[1].(map[string]any)["id"])
|
||||
require.Equal(t, "vendor_remote_item", input[2].(map[string]any)["id"])
|
||||
for i, expectedID := range []string{"fc_previous_turn", "fc_remote_item", "vendor_remote_item"} {
|
||||
item, itemOK := input[i].(map[string]any)
|
||||
require.True(t, itemOK)
|
||||
require.Equal(t, expectedID, item["id"])
|
||||
}
|
||||
}
|
||||
|
||||
func TestApplyCodexOAuthTransform_BoundsLongCallIDsAndPreservesPairing(t *testing.T) {
|
||||
|
||||
@@ -257,6 +257,33 @@ func (s *OpenAIGatewayService) SelectAccountForModelWithExclusions(ctx context.C
|
||||
return s.selectAccountForModelWithExclusions(s.withOpenAIQuotaAutoPauseContext(ctx), groupID, PlatformOpenAI, sessionHash, requestedModel, excludedIDs, false, 0, "", false)
|
||||
}
|
||||
|
||||
// SelectAccountForTokenCount selects an account for a non-billable token-count
|
||||
// request. It applies the normal platform, model, capability, and runtime
|
||||
// eligibility checks without acquiring or waiting for a generation slot.
|
||||
func (s *OpenAIGatewayService) SelectAccountForTokenCount(
|
||||
ctx context.Context,
|
||||
groupID *int64,
|
||||
sessionHash string,
|
||||
requestedModel string,
|
||||
requiredCapability OpenAIEndpointCapability,
|
||||
platform string,
|
||||
) (*Account, error) {
|
||||
ctx = WithOpenAIProfitControlSuppressed(ctx)
|
||||
ctx = s.withOpenAIQuotaAutoPauseContext(ctx)
|
||||
return s.selectAccountForModelWithExclusions(
|
||||
ctx,
|
||||
groupID,
|
||||
platform,
|
||||
sessionHash,
|
||||
requestedModel,
|
||||
nil,
|
||||
false,
|
||||
0,
|
||||
requiredCapability,
|
||||
false,
|
||||
)
|
||||
}
|
||||
|
||||
// NormalizeOpenAICompatiblePlatform 保留 grok 与国产 OpenAI 兼容供应商(kimi/zhipu/
|
||||
// deepseek)的原值,其他值一律归一为 openai。调度器据此对账号与请求做精确平台匹配:
|
||||
// kimi 分组请求只命中 kimi 账号,语义与 openai/grok 一致。
|
||||
|
||||
@@ -99,22 +99,9 @@ func (s *OpenAIGatewayService) ProxyResponsesWebSocketFromClient(
|
||||
}
|
||||
}
|
||||
|
||||
// Only persistent inbound WebSocket sessions participate in preemption.
|
||||
// HTTP ingress that opportunistically uses an upstream WS is handled by
|
||||
// forwardOpenAIWSV2 and deliberately never reaches this registration.
|
||||
preemptSessionHash := ""
|
||||
preemptGroupID := getOpenAIGroupIDFromContext(c)
|
||||
if account.Platform == PlatformOpenAI && account.Type == AccountTypeOAuth {
|
||||
preemptSessionHash = s.GenerateSessionHash(c, firstClientMessage)
|
||||
}
|
||||
if preemptCtx, cleanupPreempt, armed, preemptedPrevious := s.beginOpenAIWSSessionPreemptContext(
|
||||
ctx,
|
||||
account,
|
||||
preemptGroupID,
|
||||
getAPIKeyIDFromContext(c),
|
||||
preemptSessionHash,
|
||||
false,
|
||||
); armed {
|
||||
// The handler normally owns this registration across retry attempts. Direct
|
||||
// callers still get the same session-scoped preemption behavior here.
|
||||
if preemptCtx, cleanupPreempt, armed := s.BeginOpenAIWSIngressSessionPreemption(ctx, c, account, firstClientMessage); armed {
|
||||
ctx = preemptCtx
|
||||
defer cleanupPreempt()
|
||||
defer func() {
|
||||
@@ -122,12 +109,6 @@ func (s *OpenAIGatewayService) ProxyResponsesWebSocketFromClient(
|
||||
returnErr = errOpenAIWSSessionPreempted
|
||||
}
|
||||
}()
|
||||
if preemptedPrevious {
|
||||
if stateStore := s.getOpenAIWSStateStore(); stateStore != nil {
|
||||
stateStore.DeleteSessionTurnState(preemptGroupID, preemptSessionHash)
|
||||
stateStore.DeleteSessionConn(preemptGroupID, preemptSessionHash)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
wsDecision := s.getOpenAIWSProtocolResolver().Resolve(account)
|
||||
|
||||
@@ -8,6 +8,7 @@ import (
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/google/uuid"
|
||||
)
|
||||
|
||||
@@ -38,6 +39,49 @@ type openAIWSSessionPreemptKey struct {
|
||||
sessionHash string
|
||||
}
|
||||
|
||||
type openAIWSSessionPreemptContextKey struct{}
|
||||
|
||||
// BeginOpenAIWSIngressSessionPreemption keeps a persistent inbound WS session
|
||||
// registered across upstream retry attempts. Nested forwarding calls reuse the
|
||||
// registration so returning from one attempt cannot create a preemption gap.
|
||||
func (s *OpenAIGatewayService) BeginOpenAIWSIngressSessionPreemption(
|
||||
ctx context.Context,
|
||||
c *gin.Context,
|
||||
account *Account,
|
||||
firstClientMessage []byte,
|
||||
) (context.Context, func(), bool) {
|
||||
if ctx == nil {
|
||||
ctx = context.Background()
|
||||
}
|
||||
if armed, _ := ctx.Value(openAIWSSessionPreemptContextKey{}).(bool); armed {
|
||||
return ctx, func() {}, true
|
||||
}
|
||||
|
||||
preemptSessionHash := ""
|
||||
preemptGroupID := getOpenAIGroupIDFromContext(c)
|
||||
if account != nil && account.Platform == PlatformOpenAI && account.Type == AccountTypeOAuth {
|
||||
preemptSessionHash = s.GenerateSessionHash(c, firstClientMessage)
|
||||
}
|
||||
preemptCtx, cleanup, armed, preemptedPrevious := s.beginOpenAIWSSessionPreemptContext(
|
||||
ctx,
|
||||
account,
|
||||
preemptGroupID,
|
||||
getAPIKeyIDFromContext(c),
|
||||
preemptSessionHash,
|
||||
false,
|
||||
)
|
||||
if !armed {
|
||||
return ctx, func() {}, false
|
||||
}
|
||||
if preemptedPrevious {
|
||||
if stateStore := s.getOpenAIWSStateStore(); stateStore != nil {
|
||||
stateStore.DeleteSessionTurnState(preemptGroupID, preemptSessionHash)
|
||||
stateStore.DeleteSessionConn(preemptGroupID, preemptSessionHash)
|
||||
}
|
||||
}
|
||||
return context.WithValue(preemptCtx, openAIWSSessionPreemptContextKey{}, true), cleanup, true
|
||||
}
|
||||
|
||||
func newOpenAIWSSessionPreemptKey(groupID, apiKeyID int64, sessionHash string) (openAIWSSessionPreemptKey, bool) {
|
||||
sessionHash = strings.TrimSpace(sessionHash)
|
||||
if groupID <= 0 || apiKeyID <= 0 || sessionHash == "" {
|
||||
|
||||
@@ -5,10 +5,13 @@ package service
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"sync"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
@@ -110,6 +113,44 @@ func TestOpenAIWSSessionPreemptContextEligibilityAndLocalCancellation(t *testing
|
||||
secondCleanup()
|
||||
}
|
||||
|
||||
func TestOpenAIWSIngressSessionPreemptionSurvivesNestedForwardCleanup(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
groupID := int64(7)
|
||||
newContext := func() *gin.Context {
|
||||
c, _ := gin.CreateTestContext(httptest.NewRecorder())
|
||||
c.Request = httptest.NewRequest(http.MethodGet, "/v1/responses", nil)
|
||||
c.Set("api_key", &APIKey{ID: 11, GroupID: &groupID})
|
||||
return c
|
||||
}
|
||||
|
||||
svc := &OpenAIGatewayService{}
|
||||
account := &Account{ID: 1, Platform: PlatformOpenAI, Type: AccountTypeOAuth}
|
||||
firstMessage := []byte(`{"type":"response.create","prompt_cache_key":"session-1","input":"hello"}`)
|
||||
|
||||
firstCtx, firstCleanup, armed := svc.BeginOpenAIWSIngressSessionPreemption(
|
||||
context.Background(), newContext(), account, firstMessage,
|
||||
)
|
||||
require.True(t, armed)
|
||||
defer firstCleanup()
|
||||
|
||||
// ProxyResponsesWebSocketFromClient enters the same helper for each upstream
|
||||
// attempt. Its cleanup must not release the handler-owned registration.
|
||||
nestedCtx, nestedCleanup, armed := svc.BeginOpenAIWSIngressSessionPreemption(
|
||||
firstCtx, newContext(), account, firstMessage,
|
||||
)
|
||||
require.True(t, armed)
|
||||
require.Equal(t, firstCtx, nestedCtx)
|
||||
nestedCleanup()
|
||||
require.NoError(t, firstCtx.Err())
|
||||
|
||||
_, secondCleanup, armed := svc.BeginOpenAIWSIngressSessionPreemption(
|
||||
context.Background(), newContext(), account, firstMessage,
|
||||
)
|
||||
require.True(t, armed)
|
||||
defer secondCleanup()
|
||||
require.True(t, IsOpenAIWSSessionPreemptedError(context.Cause(firstCtx)))
|
||||
}
|
||||
|
||||
func TestOpenAIWSSessionPreemptRemoteClaimAndStaleReleaseAreAtomic(t *testing.T) {
|
||||
cache := &openAIWSSessionPreemptCacheStub{}
|
||||
svc := &OpenAIGatewayService{cache: cache}
|
||||
|
||||
@@ -36,10 +36,16 @@ func (r *errSettingRepo) Get(_ context.Context, _ string) (*Setting, error) {
|
||||
type overloadAccountRepoStub struct {
|
||||
mockAccountRepoForGemini
|
||||
overloadCalls int
|
||||
errorCalls int
|
||||
lastOverloadID int64
|
||||
lastOverloadEnd time.Time
|
||||
}
|
||||
|
||||
func (r *overloadAccountRepoStub) SetError(_ context.Context, _ int64, _ string) error {
|
||||
r.errorCalls++
|
||||
return nil
|
||||
}
|
||||
|
||||
func (r *overloadAccountRepoStub) SetOverloaded(_ context.Context, id int64, until time.Time) error {
|
||||
r.overloadCalls++
|
||||
r.lastOverloadID = id
|
||||
@@ -269,7 +275,7 @@ func TestHandle529_DBReadError_FallsBackToConfig(t *testing.T) {
|
||||
require.WithinDuration(t, before.Add(7*time.Minute), accountRepo.lastOverloadEnd, 2*time.Second)
|
||||
}
|
||||
|
||||
func TestHandleUpstreamError_529BypassesPoolAndCustomCodeGates(t *testing.T) {
|
||||
func TestHandleUpstreamError_529RespectsAccountPolicies(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
credentials map[string]any
|
||||
@@ -301,12 +307,32 @@ func TestHandleUpstreamError_529BypassesPoolAndCustomCodeGates(t *testing.T) {
|
||||
shouldDisable := svc.HandleUpstreamError(context.Background(), account, 529, nil, []byte(`{"error":{"message":"overloaded"}}`))
|
||||
|
||||
require.False(t, shouldDisable)
|
||||
require.Equal(t, 1, repo.overloadCalls)
|
||||
require.Equal(t, account.ID, repo.lastOverloadID)
|
||||
require.Zero(t, repo.overloadCalls)
|
||||
require.Zero(t, repo.errorCalls)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestHandleUpstreamError_529CustomCodeDisablesInsteadOfOverloadCooldown(t *testing.T) {
|
||||
repo := &overloadAccountRepoStub{}
|
||||
svc := NewRateLimitService(repo, nil, &config.Config{}, nil, nil)
|
||||
account := &Account{
|
||||
ID: 102,
|
||||
Platform: PlatformOpenAI,
|
||||
Type: AccountTypeAPIKey,
|
||||
Credentials: map[string]any{
|
||||
"custom_error_codes_enabled": true,
|
||||
"custom_error_codes": []any{float64(529)},
|
||||
},
|
||||
}
|
||||
|
||||
shouldDisable := svc.HandleUpstreamError(context.Background(), account, 529, nil, []byte(`{"error":{"message":"overloaded"}}`))
|
||||
|
||||
require.True(t, shouldDisable)
|
||||
require.Equal(t, 1, repo.errorCalls)
|
||||
require.Zero(t, repo.overloadCalls)
|
||||
}
|
||||
|
||||
// ===========================================================================
|
||||
// Model: defaults & JSON round-trip
|
||||
// ===========================================================================
|
||||
|
||||
@@ -252,11 +252,6 @@ const (
|
||||
// 自定义错误码开启时覆盖后续所有逻辑(包括临时不可调度)。
|
||||
func (s *RateLimitService) CheckErrorPolicy(ctx context.Context, account *Account, statusCode int, responseBody []byte, requestedModel ...string) ErrorPolicyResult {
|
||||
ctx = withTempUnschedulableModel(ctx, requestedModel)
|
||||
// 529 is governed by the global overload cooldown. Return Matched before
|
||||
// local pool/custom-code filters so every caller reaches HandleUpstreamError.
|
||||
if statusCode == 529 {
|
||||
return ErrorPolicyMatched
|
||||
}
|
||||
if account.IsCustomErrorCodesEnabled() {
|
||||
if account.ShouldHandleErrorCode(statusCode) {
|
||||
return ErrorPolicyMatched
|
||||
@@ -272,6 +267,11 @@ func (s *RateLimitService) CheckErrorPolicy(ctx context.Context, account *Accoun
|
||||
}
|
||||
return ErrorPolicySkipped
|
||||
}
|
||||
// The global overload cooldown is the default for ordinary accounts. Explicit
|
||||
// account policies above retain precedence over this fallback.
|
||||
if statusCode == 529 {
|
||||
return ErrorPolicyMatched
|
||||
}
|
||||
if s.tryTempUnschedulable(ctx, account, statusCode, responseBody, firstRequestedModel(requestedModel)) {
|
||||
return ErrorPolicyTempUnscheduled
|
||||
}
|
||||
@@ -287,14 +287,6 @@ func (s *RateLimitService) HandleUpstreamError(ctx context.Context, account *Acc
|
||||
s.maybeHandleOpenAITeamLinkedError(ctx, account, statusCode, responseBody)
|
||||
customErrorCodesEnabled := account.IsCustomErrorCodesEnabled()
|
||||
|
||||
// The configured 529 cooldown is a global overload policy. Apply it before
|
||||
// pool-mode and custom-code gates so those local retry policies cannot leave
|
||||
// an overloaded account eligible for new requests.
|
||||
if statusCode == 529 {
|
||||
s.handle529(ctx, account)
|
||||
return false
|
||||
}
|
||||
|
||||
// 池模式默认不标记本地账号状态;但管理员显式配置的临时不可调度规则优先。
|
||||
// 401 保留现有认证错误语义,不在这里改变池模式的认证处理。
|
||||
if account.IsPoolMode() && !customErrorCodesEnabled {
|
||||
@@ -312,6 +304,15 @@ func (s *RateLimitService) HandleUpstreamError(ctx context.Context, account *Acc
|
||||
return false
|
||||
}
|
||||
|
||||
if statusCode == 529 {
|
||||
if customErrorCodesEnabled {
|
||||
s.handleCustomErrorCode(ctx, account, statusCode, extractUpstreamErrorMessage(responseBody))
|
||||
return true
|
||||
}
|
||||
s.handle529(ctx, account)
|
||||
return false
|
||||
}
|
||||
|
||||
if len(requestedModel) > 0 && s.HandleUpstreamModelNotFound(ctx, account, requestedModel[0], statusCode, responseBody) {
|
||||
return true
|
||||
}
|
||||
@@ -499,7 +500,7 @@ func (s *RateLimitService) HandleUpstreamError(ctx context.Context, account *Acc
|
||||
s.handle429(ctx, account, headers, responseBody)
|
||||
shouldDisable = false
|
||||
case 529:
|
||||
// Handled before pool/custom-code policy gates above.
|
||||
// Handled after pool/custom-code policy gates above.
|
||||
shouldDisable = false
|
||||
default:
|
||||
// 自定义错误码启用时:在列表中的错误码都应该停止调度
|
||||
|
||||
Reference in New Issue
Block a user