mirror of
https://github.com/Wei-Shaw/sub2api.git
synced 2026-10-07 15:57:53 +08:00
完善 OpenAI 网关切换与运维错误语义
This commit is contained in:
@@ -56,6 +56,9 @@ const (
|
||||
const profitVetoExhaustedMessage = "No available accounts: all candidates rejected by group profit control"
|
||||
|
||||
func sameAccountRetryDelayFor(failoverErr *service.UpstreamFailoverError, retryCount int) time.Duration {
|
||||
if failoverErr != nil && failoverErr.SameAccountRetryDelay > 0 {
|
||||
return failoverErr.SameAccountRetryDelay
|
||||
}
|
||||
if failoverErr == nil || !failoverErr.RequestScopedTransient || retryCount <= 1 {
|
||||
return sameAccountRetryDelay
|
||||
}
|
||||
@@ -70,6 +73,16 @@ func sameAccountRetryDelayFor(failoverErr *service.UpstreamFailoverError, retryC
|
||||
return delay
|
||||
}
|
||||
|
||||
func sameAccountRetryAllowed(failoverErr *service.UpstreamFailoverError, retryCount, retryLimit int) bool {
|
||||
if failoverErr == nil || !failoverErr.RetryableOnSameAccount {
|
||||
return false
|
||||
}
|
||||
if !failoverErr.SameAccountRetryDeadline.IsZero() {
|
||||
return time.Now().Before(failoverErr.SameAccountRetryDeadline)
|
||||
}
|
||||
return retryCount < retryLimit
|
||||
}
|
||||
|
||||
// FailoverState 跨循环迭代共享的 failover 状态
|
||||
type FailoverState struct {
|
||||
SwitchCount int
|
||||
@@ -158,14 +171,14 @@ func (s *FailoverState) HandleFailoverError(
|
||||
}
|
||||
|
||||
// 同账号重试不算切换账号,粘性会话仅在实际切换时强制缓存计费。
|
||||
sameAccountRetry := failoverErr.RetryableOnSameAccount && s.SameAccountRetryCount[accountID] < retryLimit
|
||||
sameAccountRetry := sameAccountRetryAllowed(failoverErr, s.SameAccountRetryCount[accountID], retryLimit)
|
||||
if needForceCacheBilling(s.hasBoundSession, failoverErr, sameAccountRetry) {
|
||||
s.ForceCacheBilling = true
|
||||
}
|
||||
|
||||
// 同账号重试:对 RetryableOnSameAccount 的临时性错误,先在同一账号上重试。
|
||||
// 重试次数上限 retryLimit 由调用方传入(账号级 pool_mode_retry_count 配置)。
|
||||
if failoverErr.RetryableOnSameAccount && s.SameAccountRetryCount[accountID] < retryLimit {
|
||||
if sameAccountRetryAllowed(failoverErr, s.SameAccountRetryCount[accountID], retryLimit) {
|
||||
s.SameAccountRetryCount[accountID]++
|
||||
retryDelay := sameAccountRetryDelayFor(failoverErr, s.SameAccountRetryCount[accountID])
|
||||
logger.FromContext(ctx).Warn("gateway.failover_same_account_retry",
|
||||
|
||||
@@ -58,6 +58,21 @@ func TestSameAccountRetryDelayFor(t *testing.T) {
|
||||
t.Run("nil error keeps fixed delay", func(t *testing.T) {
|
||||
require.Equal(t, 500*time.Millisecond, sameAccountRetryDelayFor(nil, 10))
|
||||
})
|
||||
|
||||
t.Run("explicit oauth delay wins", func(t *testing.T) {
|
||||
err := &service.UpstreamFailoverError{SameAccountRetryDelay: 3 * time.Second}
|
||||
require.Equal(t, 3*time.Second, sameAccountRetryDelayFor(err, 1))
|
||||
})
|
||||
}
|
||||
|
||||
func TestSameAccountRetryAllowedUsesDeadlineInsteadOfPoolCount(t *testing.T) {
|
||||
err := &service.UpstreamFailoverError{
|
||||
RetryableOnSameAccount: true,
|
||||
SameAccountRetryDeadline: time.Now().Add(time.Minute),
|
||||
}
|
||||
require.True(t, sameAccountRetryAllowed(err, 100, 0))
|
||||
err.SameAccountRetryDeadline = time.Now().Add(-time.Second)
|
||||
require.False(t, sameAccountRetryAllowed(err, 0, 100))
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
@@ -320,6 +335,20 @@ func TestHandleFailoverError_CacheBilling(t *testing.T) {
|
||||
require.Zero(t, fs.SwitchCount)
|
||||
})
|
||||
|
||||
t.Run("deadline允许超过计数上限时仍不强制缓存计费", func(t *testing.T) {
|
||||
mock := &mockTempUnscheduler{}
|
||||
fs := NewFailoverState(3, true)
|
||||
fs.SameAccountRetryCount[100] = maxSameAccountRetries
|
||||
err := newTestFailoverErr(http.StatusTooManyRequests, true, false)
|
||||
err.SameAccountRetryDeadline = time.Now().Add(time.Minute)
|
||||
err.SameAccountRetryDelay = time.Nanosecond
|
||||
|
||||
fs.HandleFailoverError(context.Background(), mock, 100, "openai", maxSameAccountRetries, err)
|
||||
|
||||
require.False(t, fs.ForceCacheBilling)
|
||||
require.Zero(t, fs.SwitchCount)
|
||||
})
|
||||
|
||||
t.Run("同账号重试耗尽并实际切换时设置ForceCacheBilling", func(t *testing.T) {
|
||||
mock := &mockTempUnscheduler{}
|
||||
fs := NewFailoverState(3, true)
|
||||
|
||||
@@ -5,6 +5,7 @@ import (
|
||||
"errors"
|
||||
"net/http"
|
||||
"strconv"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/Wei-Shaw/sub2api/internal/pkg/ip"
|
||||
@@ -173,6 +174,7 @@ func (h *GatewayHandler) ChatCompletions(c *gin.Context) {
|
||||
if err != nil {
|
||||
if len(fs.FailedAccountIDs) == 0 {
|
||||
cls := classifyNoAccountErrorFromGin(c, h.gatewayService, apiKey, reqModel, reqModel, groupPlatform)
|
||||
cls = classifySelectionFailureError(err, cls)
|
||||
if !cls.ModelNotFound {
|
||||
markOpsRoutingCapacityLimitedIfNoAvailable(c, err)
|
||||
}
|
||||
@@ -384,6 +386,14 @@ func (h *GatewayHandler) handleCCFailoverExhausted(c *gin.Context, lastErr *serv
|
||||
h.chatCompletionsErrorResponse(c, status, "server_error", message)
|
||||
return
|
||||
}
|
||||
if lastErr != nil && lastErr.IsOpenAICapacityShed() && strings.TrimSpace(lastErr.ClientMessage) != "" {
|
||||
status := lastErr.ClientStatusCode
|
||||
if status <= 0 {
|
||||
status = http.StatusServiceUnavailable
|
||||
}
|
||||
h.chatCompletionsErrorResponse(c, status, "server_error", lastErr.ClientMessage)
|
||||
return
|
||||
}
|
||||
statusCode := http.StatusBadGateway
|
||||
if lastErr != nil && lastErr.StatusCode > 0 {
|
||||
statusCode = lastErr.StatusCode
|
||||
|
||||
@@ -5,6 +5,7 @@ import (
|
||||
"errors"
|
||||
"net/http"
|
||||
"strconv"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/Wei-Shaw/sub2api/internal/pkg/ip"
|
||||
@@ -175,6 +176,7 @@ func (h *GatewayHandler) Responses(c *gin.Context) {
|
||||
if err != nil {
|
||||
if len(fs.FailedAccountIDs) == 0 {
|
||||
cls := classifyNoAccountErrorFromGin(c, h.gatewayService, apiKey, reqModel, reqModel, effectiveAPIKeyPlatform(c, apiKey))
|
||||
cls = classifySelectionFailureError(err, cls)
|
||||
if !cls.ModelNotFound {
|
||||
markOpsRoutingCapacityLimitedIfNoAvailable(c, err)
|
||||
}
|
||||
@@ -361,25 +363,38 @@ func (h *GatewayHandler) responsesErrorResponse(c *gin.Context, status int, code
|
||||
|
||||
// handleResponsesFailoverExhausted writes a failover-exhausted error in Responses format.
|
||||
func (h *GatewayHandler) handleResponsesFailoverExhausted(c *gin.Context, lastErr *service.UpstreamFailoverError, streamStarted bool) {
|
||||
if streamStarted {
|
||||
return // Can't write error after stream started
|
||||
}
|
||||
if lastErr != nil {
|
||||
copyFailoverRetryAfter(c, lastErr.ResponseHeaders)
|
||||
}
|
||||
if lastErr != nil && lastErr.IsCredentialFailure() {
|
||||
status, message := credentialFailoverClientResponse(lastErr)
|
||||
h.responsesErrorResponse(c, status, "server_error", message)
|
||||
return
|
||||
}
|
||||
statusCode := http.StatusBadGateway
|
||||
if lastErr != nil && lastErr.StatusCode > 0 {
|
||||
statusCode = lastErr.StatusCode
|
||||
}
|
||||
if lastErr != nil && service.IsOpenAISilentRefusalErrorBody(lastErr.ResponseBody) {
|
||||
status, code, message := statusCode, "server_error", "All available accounts exhausted"
|
||||
if lastErr != nil && lastErr.IsCredentialFailure() {
|
||||
status, message = credentialFailoverClientResponse(lastErr)
|
||||
} else if lastErr != nil && lastErr.IsOpenAICapacityShed() && strings.TrimSpace(lastErr.ClientMessage) != "" {
|
||||
status = lastErr.ClientStatusCode
|
||||
if status <= 0 {
|
||||
status = http.StatusServiceUnavailable
|
||||
}
|
||||
message = lastErr.ClientMessage
|
||||
} else if lastErr != nil && service.IsOpenAISilentRefusalErrorBody(lastErr.ResponseBody) {
|
||||
service.SetOpsUpstreamError(c, statusCode, service.OpenAISilentRefusalClientMessage(), "")
|
||||
h.responsesErrorResponse(c, http.StatusBadGateway, "upstream_error", service.OpenAISilentRefusalClientMessage())
|
||||
status, code, message = http.StatusBadGateway, "upstream_error", service.OpenAISilentRefusalClientMessage()
|
||||
} else if lastErr != nil && statusCode == http.StatusTooManyRequests {
|
||||
status, code, message = http.StatusTooManyRequests, "rate_limit_error", "All available accounts are currently rate-limited. Please retry later."
|
||||
}
|
||||
if streamStarted {
|
||||
// A slot-wait heartbeat commits HTTP 200 before any upstream response.
|
||||
// In that case a terminal frame is still required; once any semantic or
|
||||
// official terminal bytes exist, preserve them without appending a second
|
||||
// generic response.failed.
|
||||
service.MarkOpsStreamError(c, code, message, status)
|
||||
if c != nil && c.Writer != nil && (c.Writer.Size() <= 0 || gatewayStreamHasOnlyHeartbeats(c)) {
|
||||
writeResponsesFailedSSE(c, code, message)
|
||||
}
|
||||
return
|
||||
}
|
||||
h.responsesErrorResponse(c, statusCode, "server_error", "All available accounts exhausted")
|
||||
h.responsesErrorResponse(c, status, code, message)
|
||||
}
|
||||
|
||||
@@ -16,6 +16,29 @@ import (
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
|
||||
const gatewayStreamHeartbeatBytesKey = "gateway_stream_heartbeat_bytes"
|
||||
|
||||
func recordGatewayStreamHeartbeat(c *gin.Context, written int) {
|
||||
if c == nil || written <= 0 {
|
||||
return
|
||||
}
|
||||
total, _ := c.Get(gatewayStreamHeartbeatBytesKey)
|
||||
bytes, _ := total.(int)
|
||||
c.Set(gatewayStreamHeartbeatBytesKey, bytes+written)
|
||||
}
|
||||
|
||||
func gatewayStreamHasOnlyHeartbeats(c *gin.Context) bool {
|
||||
if c == nil || c.Writer == nil {
|
||||
return false
|
||||
}
|
||||
value, ok := c.Get(gatewayStreamHeartbeatBytesKey)
|
||||
if !ok {
|
||||
return false
|
||||
}
|
||||
heartbeatBytes, _ := value.(int)
|
||||
return heartbeatBytes > 0 && c.Writer.Size() == heartbeatBytes
|
||||
}
|
||||
|
||||
// claudeCodeValidator is a singleton validator for Claude Code client detection
|
||||
var claudeCodeValidator = service.NewClaudeCodeValidator()
|
||||
|
||||
@@ -396,9 +419,11 @@ func (h *ConcurrencyHelper) waitForSlotWithPingTimeout(c *gin.Context, slotType
|
||||
c.Header("X-Accel-Buffering", "no")
|
||||
*streamStarted = true
|
||||
}
|
||||
if _, err := fmt.Fprint(c.Writer, string(h.pingFormat)); err != nil {
|
||||
written, err := fmt.Fprint(c.Writer, string(h.pingFormat))
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
recordGatewayStreamHeartbeat(c, written)
|
||||
flusher.Flush()
|
||||
|
||||
case <-timer.C:
|
||||
|
||||
@@ -346,7 +346,7 @@ func (h *OpenAIGatewayHandler) handleGrokMedia(c *gin.Context, endpoint service.
|
||||
return
|
||||
}
|
||||
if failoverErr.ShouldReportAccountScheduleFailure() {
|
||||
h.gatewayService.ReportOpenAIAccountScheduleResult(account.ID, grokMediaScheduleModel(account, routingModel, nil), false, nil)
|
||||
h.gatewayService.ReportOpenAIAccountScheduleResult(account, grokMediaScheduleModel(account, routingModel, nil), false, nil)
|
||||
}
|
||||
if c.Writer.Size() != writerSizeBeforeForward {
|
||||
h.handleFailoverExhausted(c, failoverErr, true)
|
||||
@@ -398,7 +398,7 @@ func (h *OpenAIGatewayHandler) handleGrokMedia(c *gin.Context, endpoint service.
|
||||
)
|
||||
continue
|
||||
}
|
||||
h.gatewayService.ReportOpenAIAccountScheduleResult(account.ID, grokMediaScheduleModel(account, routingModel, nil), false, nil)
|
||||
h.gatewayService.ReportOpenAIAccountScheduleResult(account, grokMediaScheduleModel(account, routingModel, nil), false, nil)
|
||||
if !service.IsResponseCommitted(c) && c.Writer.Size() == writerSizeBeforeForward {
|
||||
h.errorResponse(c, http.StatusBadGateway, "upstream_error", "Upstream request failed")
|
||||
}
|
||||
@@ -409,7 +409,7 @@ func (h *OpenAIGatewayHandler) handleGrokMedia(c *gin.Context, endpoint service.
|
||||
return
|
||||
}
|
||||
|
||||
h.gatewayService.ReportOpenAIAccountScheduleResult(account.ID, grokMediaScheduleModel(account, routingModel, result), true, nil)
|
||||
h.gatewayService.ReportOpenAIAccountScheduleResult(account, grokMediaScheduleModel(account, routingModel, result), true, nil)
|
||||
if isGrokVideoCreateEndpoint(endpoint) && strings.TrimSpace(result.ResponseID) != "" {
|
||||
if err := h.gatewayService.BindGrokMediaVideoRequestAccount(
|
||||
requestCtx, apiKey.GroupID, result.ResponseID, subject.UserID, apiKey.ID, account.ID,
|
||||
|
||||
@@ -4,6 +4,8 @@ import (
|
||||
"context"
|
||||
"fmt"
|
||||
"net/http"
|
||||
"regexp"
|
||||
"strconv"
|
||||
"strings"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
@@ -32,6 +34,29 @@ type noAccountErrorClassification struct {
|
||||
ModelNotFound bool // true when this is a 404 model_not_found classification
|
||||
}
|
||||
|
||||
var selectionModelRateLimitedPattern = regexp.MustCompile(`(?:model_rate_limited|rate_limited)=(\d+)`)
|
||||
|
||||
// classifySelectionFailureError preserves the scheduler's compact reason when
|
||||
// every model-capable account is temporarily rate limited.
|
||||
func classifySelectionFailureError(err error, fallback noAccountErrorClassification) noAccountErrorClassification {
|
||||
if err == nil {
|
||||
return fallback
|
||||
}
|
||||
match := selectionModelRateLimitedPattern.FindStringSubmatch(strings.ToLower(err.Error()))
|
||||
if len(match) != 2 {
|
||||
return fallback
|
||||
}
|
||||
count, parseErr := strconv.Atoi(match[1])
|
||||
if parseErr != nil || count <= 0 {
|
||||
return fallback
|
||||
}
|
||||
return noAccountErrorClassification{
|
||||
Status: http.StatusTooManyRequests,
|
||||
ErrType: "rate_limit_error",
|
||||
Message: "All available accounts are currently rate-limited. Please retry later.",
|
||||
}
|
||||
}
|
||||
|
||||
// classifyNoAccountError decides between 404 model_not_found and 503
|
||||
// api_error for "no available accounts" failures.
|
||||
//
|
||||
|
||||
@@ -61,6 +61,21 @@ func TestClassifyNoAccountError_NilDiagnoser_Falls503(t *testing.T) {
|
||||
require.False(t, cls.ModelNotFound)
|
||||
}
|
||||
|
||||
func TestClassifySelectionFailureError_RateLimitedPool(t *testing.T) {
|
||||
fallback := noAccountErrorClassification{Status: http.StatusServiceUnavailable, ErrType: "api_error", Message: "Service temporarily unavailable"}
|
||||
|
||||
got := classifySelectionFailureError(
|
||||
fmt.Errorf("no available accounts supporting model: gpt-5.6-sol (total=3 eligible=0 model_rate_limited=3)"),
|
||||
fallback,
|
||||
)
|
||||
|
||||
require.Equal(t, http.StatusTooManyRequests, got.Status)
|
||||
require.Equal(t, "rate_limit_error", got.ErrType)
|
||||
require.Contains(t, got.Message, "rate-limited")
|
||||
require.Equal(t, fallback, classifySelectionFailureError(fmt.Errorf("model_rate_limited=0"), fallback))
|
||||
require.Equal(t, fallback, classifySelectionFailureError(fmt.Errorf("no available accounts"), fallback))
|
||||
}
|
||||
|
||||
func TestClassifyNoAccountError_NilAPIKey_Falls503(t *testing.T) {
|
||||
c := newTestGinContextWithRequest()
|
||||
fd := &fakeDiagnoser{resp: service.ModelAvailabilityDiagnosis{HasAccountsInPool: true, HasModelSupport: false}}
|
||||
|
||||
@@ -113,6 +113,7 @@ func (h *OpenAIGatewayHandler) AlphaSearch(c *gin.Context) {
|
||||
sessionHash := h.gatewayService.GenerateSessionHashWithFallback(c, nil, searchID)
|
||||
profitVetoCount := 0
|
||||
failedAccountIDs := make(map[int64]struct{})
|
||||
sameAccountRetryCount := make(map[int64]int)
|
||||
var lastFailoverErr *service.UpstreamFailoverError
|
||||
switchCount := 0
|
||||
var oauth429FailoverState service.OpenAIOAuth429FailoverState
|
||||
@@ -186,7 +187,7 @@ func (h *OpenAIGatewayHandler) AlphaSearch(c *gin.Context) {
|
||||
service.SetOpsLatencyMs(c, service.OpsResponseLatencyMsKey, time.Since(forwardStart).Milliseconds())
|
||||
|
||||
if err == nil {
|
||||
h.gatewayService.ReportOpenAIAccountScheduleResult(account.ID, account.GetMappedModel(requestedModel), true, nil)
|
||||
h.gatewayService.ReportOpenAIAccountScheduleResult(account, openAIAccountScheduleModel(c, account, requestedModel, false, result), true, nil)
|
||||
if result != nil {
|
||||
h.recordAlphaSearchUsage(c, apiKey, account, subscription, channelMapping, requestedModel, body, result, subject.UserID)
|
||||
}
|
||||
@@ -195,7 +196,7 @@ func (h *OpenAIGatewayHandler) AlphaSearch(c *gin.Context) {
|
||||
|
||||
var failoverErr *service.UpstreamFailoverError
|
||||
if !errors.As(err, &failoverErr) {
|
||||
h.gatewayService.ReportOpenAIAccountScheduleResult(account.ID, account.GetMappedModel(requestedModel), false, nil)
|
||||
h.gatewayService.ReportOpenAIAccountScheduleResult(account, openAIAccountScheduleModel(c, account, requestedModel, false, result), false, nil, err)
|
||||
if c.Writer.Size() == writerSizeBeforeForward {
|
||||
h.errorResponse(c, http.StatusBadGateway, "upstream_error", "Upstream request failed")
|
||||
}
|
||||
@@ -203,7 +204,7 @@ func (h *OpenAIGatewayHandler) AlphaSearch(c *gin.Context) {
|
||||
return
|
||||
}
|
||||
|
||||
h.gatewayService.ReportOpenAIAccountScheduleResult(account.ID, account.GetMappedModel(requestedModel), false, nil)
|
||||
h.gatewayService.ReportOpenAIAccountScheduleResult(account, openAIAccountScheduleModel(c, account, requestedModel, false, result), false, nil, err)
|
||||
if c.Writer.Size() != writerSizeBeforeForward {
|
||||
h.handleFailoverExhausted(c, failoverErr, true)
|
||||
return
|
||||
@@ -215,6 +216,26 @@ func (h *OpenAIGatewayHandler) AlphaSearch(c *gin.Context) {
|
||||
)
|
||||
return
|
||||
}
|
||||
if failoverErr.RetryableOnSameAccount {
|
||||
retryLimit := account.GetPoolModeRetryCount()
|
||||
if sameAccountRetryAllowed(failoverErr, sameAccountRetryCount[account.ID], retryLimit) {
|
||||
sameAccountRetryCount[account.ID]++
|
||||
retryDelay := sameAccountRetryDelayFor(failoverErr, sameAccountRetryCount[account.ID])
|
||||
reqLog.Warn("openai_alpha_search.same_account_retry",
|
||||
zap.Int64("account_id", account.ID),
|
||||
zap.Int("upstream_status", failoverErr.StatusCode),
|
||||
zap.Int("retry_limit", retryLimit),
|
||||
zap.Int("retry_count", sameAccountRetryCount[account.ID]),
|
||||
zap.Duration("retry_delay", retryDelay),
|
||||
)
|
||||
select {
|
||||
case <-c.Request.Context().Done():
|
||||
return
|
||||
case <-time.After(retryDelay):
|
||||
}
|
||||
continue
|
||||
}
|
||||
}
|
||||
h.gatewayService.RecordOpenAIAccountSwitch()
|
||||
failedAccountIDs[account.ID] = struct{}{}
|
||||
lastFailoverErr = failoverErr
|
||||
|
||||
@@ -183,6 +183,7 @@ func (h *OpenAIGatewayHandler) ChatCompletions(c *gin.Context) {
|
||||
)
|
||||
if len(failedAccountIDs) == 0 {
|
||||
cls := classifyOpenAICompatibleNoAccountErrorFromGin(c, h.gatewayService, apiKey, reqModel, reqModel)
|
||||
cls = classifySelectionFailureError(err, cls)
|
||||
if !cls.ModelNotFound {
|
||||
markOpsRoutingCapacityLimitedIfNoAvailable(c, err)
|
||||
}
|
||||
@@ -240,11 +241,11 @@ func (h *OpenAIGatewayHandler) ChatCompletions(c *gin.Context) {
|
||||
}()
|
||||
return h.gatewayService.ForwardAsChatCompletions(c.Request.Context(), c, account, forwardBody, promptCacheKey, "")
|
||||
}()
|
||||
cyberBlockKeyChat := ""
|
||||
var cyberBlockBodyChat []byte
|
||||
if service.GetOpsCyberPolicy(c) != nil {
|
||||
cyberBlockKeyChat = service.CyberSessionBlockKey(apiKey.ID, c, body)
|
||||
cyberBlockBodyChat = body
|
||||
}
|
||||
h.recordCyberPolicyIfMarked(c, apiKey, account, subscription, reqModel, err != nil, cyberBlockKeyChat, clientRequestedUsageFields(c, channelMapping, reqModel, ""), service.HashUsageRequestPayload(body))
|
||||
h.recordCyberPolicyIfMarked(c, apiKey, account, subscription, reqModel, err != nil, cyberBlockBodyChat, clientRequestedUsageFields(c, channelMapping, reqModel, ""), service.HashUsageRequestPayload(body))
|
||||
|
||||
forwardDurationMs := time.Since(forwardStart).Milliseconds()
|
||||
upstreamLatencyMs, _ := getContextInt64(c, service.OpsUpstreamLatencyMsKey)
|
||||
@@ -316,11 +317,12 @@ func (h *OpenAIGatewayHandler) ChatCompletions(c *gin.Context) {
|
||||
return
|
||||
}
|
||||
if c.Writer.Size() != writerSizeBeforeForward {
|
||||
h.gatewayService.ObserveOpenAIAccountHealthFailure(c.Request.Context(), account, err)
|
||||
h.handleFailoverExhausted(c, failoverErr, true)
|
||||
return
|
||||
}
|
||||
if failoverErr.ShouldReportAccountScheduleFailure() {
|
||||
h.gatewayService.ReportOpenAIAccountScheduleResult(account.ID, account.GetMappedModel(reqModel), false, nil)
|
||||
h.gatewayService.ReportOpenAIAccountScheduleResult(account, openAIAccountScheduleModel(c, account, reqModel, false, nil), false, nil, err)
|
||||
}
|
||||
if !failoverErr.ShouldRetryNextAccount() {
|
||||
h.handleFailoverExhausted(c, failoverErr, streamStarted)
|
||||
@@ -329,7 +331,7 @@ func (h *OpenAIGatewayHandler) ChatCompletions(c *gin.Context) {
|
||||
// Pool mode: retry on the same account
|
||||
if failoverErr.RetryableOnSameAccount {
|
||||
retryLimit := account.GetPoolModeRetryCount()
|
||||
if sameAccountRetryCount[account.ID] < retryLimit {
|
||||
if sameAccountRetryAllowed(failoverErr, sameAccountRetryCount[account.ID], retryLimit) {
|
||||
sameAccountRetryCount[account.ID]++
|
||||
retryDelay := sameAccountRetryDelayFor(failoverErr, sameAccountRetryCount[account.ID])
|
||||
reqLog.Warn("openai_chat_completions.pool_mode_same_account_retry",
|
||||
@@ -367,7 +369,7 @@ func (h *OpenAIGatewayHandler) ChatCompletions(c *gin.Context) {
|
||||
)
|
||||
continue
|
||||
}
|
||||
h.gatewayService.ReportOpenAIAccountScheduleResult(account.ID, account.GetMappedModel(reqModel), false, nil)
|
||||
h.gatewayService.ReportOpenAIAccountScheduleResult(account, openAIAccountScheduleModel(c, account, reqModel, false, nil), false, nil, err)
|
||||
upstreamErrorAlreadyCommunicated := openAIForwardErrorAlreadyCommunicated(c, writerSizeBeforeForward, err)
|
||||
wroteFallback := false
|
||||
if !upstreamErrorAlreadyCommunicated {
|
||||
@@ -387,9 +389,9 @@ func (h *OpenAIGatewayHandler) ChatCompletions(c *gin.Context) {
|
||||
}
|
||||
}
|
||||
if result != nil {
|
||||
h.gatewayService.ReportOpenAIAccountScheduleResult(account.ID, account.GetMappedModel(reqModel), true, result.FirstTokenMs)
|
||||
h.gatewayService.ReportOpenAIAccountScheduleResult(account, openAIAccountScheduleModel(c, account, reqModel, false, result), true, result.FirstTokenMs)
|
||||
} else {
|
||||
h.gatewayService.ReportOpenAIAccountScheduleResult(account.ID, account.GetMappedModel(reqModel), true, nil)
|
||||
h.gatewayService.ReportOpenAIAccountScheduleResult(account, openAIAccountScheduleModel(c, account, reqModel, false, result), true, nil)
|
||||
}
|
||||
|
||||
submitChatUsage(result)
|
||||
|
||||
@@ -215,7 +215,7 @@ func (h *OpenAIGatewayHandler) Embeddings(c *gin.Context) {
|
||||
h.handleFailoverExhausted(c, failoverErr, true)
|
||||
return
|
||||
}
|
||||
h.gatewayService.ReportOpenAIAccountScheduleResult(account.ID, account.GetMappedModel(reqModel), false, nil)
|
||||
h.gatewayService.ReportOpenAIAccountScheduleResult(account, openAIAccountScheduleModel(c, account, reqModel, false, result), false, nil, err)
|
||||
if failoverClientGone(c) {
|
||||
reqLog.Info("openai_embeddings.failover_aborted_client_disconnected",
|
||||
zap.Int64("account_id", account.ID),
|
||||
@@ -239,7 +239,7 @@ func (h *OpenAIGatewayHandler) Embeddings(c *gin.Context) {
|
||||
)
|
||||
continue
|
||||
}
|
||||
h.gatewayService.ReportOpenAIAccountScheduleResult(account.ID, account.GetMappedModel(reqModel), false, nil)
|
||||
h.gatewayService.ReportOpenAIAccountScheduleResult(account, openAIAccountScheduleModel(c, account, reqModel, false, result), false, nil, err)
|
||||
if c.Writer.Size() == writerSizeBeforeForward {
|
||||
h.errorResponse(c, http.StatusBadGateway, "upstream_error", "Upstream request failed")
|
||||
}
|
||||
@@ -250,7 +250,7 @@ func (h *OpenAIGatewayHandler) Embeddings(c *gin.Context) {
|
||||
return
|
||||
}
|
||||
|
||||
h.gatewayService.ReportOpenAIAccountScheduleResult(account.ID, account.GetMappedModel(reqModel), true, nil)
|
||||
h.gatewayService.ReportOpenAIAccountScheduleResult(account, openAIAccountScheduleModel(c, account, reqModel, false, result), true, nil)
|
||||
userAgent := c.GetHeader("User-Agent")
|
||||
clientIP := ip.GetClientIP(c)
|
||||
inboundEndpoint := GetInboundEndpoint(c)
|
||||
|
||||
@@ -131,9 +131,14 @@ func (h *OpenAIGatewayHandler) ResponsesInputTokens(c *gin.Context) {
|
||||
|
||||
account := selection.Account
|
||||
setOpsSelectedAccount(c, account.ID, account.Platform)
|
||||
if selection.Acquired && selection.ReleaseFunc != nil {
|
||||
defer selection.ReleaseFunc()
|
||||
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))
|
||||
}
|
||||
@@ -182,8 +187,7 @@ func (h *OpenAIGatewayHandler) GrokCountTokens(c *gin.Context) {
|
||||
}
|
||||
|
||||
// CountTokens handles Anthropic-compatible POST /v1/messages/count_tokens for OpenAI groups.
|
||||
// It validates billing and routes to an OpenAI token-count bridge without taking concurrency slots
|
||||
// or recording usage.
|
||||
// It validates billing and routes to an OpenAI token-count bridge without recording usage.
|
||||
func (h *OpenAIGatewayHandler) CountTokens(c *gin.Context) {
|
||||
apiKey, ok := middleware2.GetAPIKeyFromContext(c)
|
||||
if !ok {
|
||||
@@ -315,9 +319,14 @@ func (h *OpenAIGatewayHandler) CountTokens(c *gin.Context) {
|
||||
|
||||
account := selection.Account
|
||||
setOpsSelectedAccount(c, account.ID, account.Platform)
|
||||
if selection.Acquired && selection.ReleaseFunc != nil {
|
||||
defer selection.ReleaseFunc()
|
||||
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
|
||||
|
||||
@@ -325,3 +334,41 @@ 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
|
||||
}
|
||||
|
||||
@@ -13,6 +13,7 @@ import (
|
||||
"github.com/Wei-Shaw/sub2api/internal/service"
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/stretchr/testify/require"
|
||||
"github.com/tidwall/gjson"
|
||||
)
|
||||
|
||||
func TestGatewayChatCredentialStopDoesNotSelectAnotherAccountAndReturnsSafe503(t *testing.T) {
|
||||
@@ -64,6 +65,112 @@ func TestGatewayChatAntigravityCredentialFailureReturnsActionableMessage(t *test
|
||||
require.NotContains(t, strings.ToLower(recorder.Body.String()), "refresh_token")
|
||||
}
|
||||
|
||||
func TestOpenAIAccessStateCredentialFailureUsesTypedSafeResponse(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
recorder := httptest.NewRecorder()
|
||||
c, _ := gin.CreateTestContext(recorder)
|
||||
|
||||
(&OpenAIGatewayHandler{}).handleFailoverExhausted(c, &service.UpstreamFailoverError{
|
||||
StatusCode: http.StatusForbidden,
|
||||
Stage: service.GatewayFailureStageAccountAuth,
|
||||
Scope: service.GatewayFailureScopeAccount,
|
||||
Reason: service.OpenAIUpstreamAccessStateReason,
|
||||
NextAccountAction: service.NextAccountRetry,
|
||||
ClientStatusCode: http.StatusBadGateway,
|
||||
ClientMessage: "Upstream access is temporarily unavailable, please retry later",
|
||||
ResponseBody: []byte(`{"error":{"message":"Your workspace is deactivated","token":"must-not-leak"}}`),
|
||||
}, false)
|
||||
|
||||
require.Equal(t, http.StatusBadGateway, recorder.Code)
|
||||
require.Contains(t, recorder.Body.String(), "Upstream access is temporarily unavailable")
|
||||
require.NotContains(t, strings.ToLower(recorder.Body.String()), "deactivated")
|
||||
require.NotContains(t, recorder.Body.String(), "must-not-leak")
|
||||
}
|
||||
|
||||
func TestOpenAICapacityFailoverExhaustionPreservesMessageAsServerError(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
message := "Our servers are currently overloaded. Please try again later."
|
||||
failoverErr := &service.UpstreamFailoverError{
|
||||
StatusCode: http.StatusBadRequest,
|
||||
ResponseBody: []byte(`{"error":{"code":"server_is_overloaded","message":"` + message + `"}}`),
|
||||
RetryableOnSameAccount: true,
|
||||
RequestScopedTransient: true,
|
||||
ClientStatusCode: http.StatusServiceUnavailable,
|
||||
ClientMessage: message,
|
||||
}
|
||||
|
||||
t.Run("native_openai", func(t *testing.T) {
|
||||
recorder := httptest.NewRecorder()
|
||||
c, _ := gin.CreateTestContext(recorder)
|
||||
(&OpenAIGatewayHandler{}).handleFailoverExhausted(c, failoverErr, false)
|
||||
require.Equal(t, http.StatusServiceUnavailable, recorder.Code)
|
||||
require.Equal(t, "server_error", gjson.Get(recorder.Body.String(), "error.type").String())
|
||||
require.Equal(t, message, gjson.Get(recorder.Body.String(), "error.message").String())
|
||||
require.NotContains(t, recorder.Body.String(), "server_is_overloaded")
|
||||
})
|
||||
|
||||
t.Run("responses_compat", func(t *testing.T) {
|
||||
recorder := httptest.NewRecorder()
|
||||
c, _ := gin.CreateTestContext(recorder)
|
||||
(&GatewayHandler{}).handleResponsesFailoverExhausted(c, failoverErr, false)
|
||||
require.Equal(t, http.StatusServiceUnavailable, recorder.Code)
|
||||
require.Equal(t, "server_error", gjson.Get(recorder.Body.String(), "error.code").String())
|
||||
require.Equal(t, message, gjson.Get(recorder.Body.String(), "error.message").String())
|
||||
})
|
||||
|
||||
t.Run("anthropic_compat", func(t *testing.T) {
|
||||
recorder := httptest.NewRecorder()
|
||||
c, _ := gin.CreateTestContext(recorder)
|
||||
(&OpenAIGatewayHandler{}).handleAnthropicFailoverExhausted(c, failoverErr, false)
|
||||
require.Equal(t, http.StatusServiceUnavailable, recorder.Code)
|
||||
require.Equal(t, "api_error", gjson.Get(recorder.Body.String(), "error.type").String())
|
||||
require.Equal(t, message, gjson.Get(recorder.Body.String(), "error.message").String())
|
||||
})
|
||||
}
|
||||
|
||||
func TestResponsesFailoverExhaustedAfterForwardedTerminalMarksOpsWithoutDuplicateFrame(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
recorder := httptest.NewRecorder()
|
||||
c, _ := gin.CreateTestContext(recorder)
|
||||
official := "event: response.failed\ndata: {\"type\":\"response.failed\",\"response\":{\"status\":\"failed\",\"error\":{\"code\":\"server_error\",\"message\":\"official failure\"}}}\n\n"
|
||||
_, _ = recorder.Write([]byte(official))
|
||||
service.MarkOpsStreamError(c, "server_error", "official failure", http.StatusBadGateway)
|
||||
|
||||
(&GatewayHandler{}).handleResponsesFailoverExhausted(c, &service.UpstreamFailoverError{
|
||||
StatusCode: http.StatusBadGateway,
|
||||
ResponseBody: []byte(`{"error":{"message":"fallback failure"}}`),
|
||||
}, true)
|
||||
|
||||
require.Equal(t, official, recorder.Body.String())
|
||||
streamErr, ok := service.GetOpsStreamError(c)
|
||||
require.True(t, ok)
|
||||
require.Equal(t, "official failure", streamErr.Message)
|
||||
|
||||
markerRecorder := httptest.NewRecorder()
|
||||
markerContext, _ := gin.CreateTestContext(markerRecorder)
|
||||
(&GatewayHandler{}).handleResponsesFailoverExhausted(markerContext, &service.UpstreamFailoverError{
|
||||
StatusCode: http.StatusTooManyRequests,
|
||||
}, true)
|
||||
require.Contains(t, markerRecorder.Body.String(), "event: response.failed")
|
||||
require.Equal(t, 1, strings.Count(markerRecorder.Body.String(), "event: response.failed"))
|
||||
streamErr, ok = service.GetOpsStreamError(markerContext)
|
||||
require.True(t, ok)
|
||||
require.Equal(t, http.StatusTooManyRequests, streamErr.IntendedStatus)
|
||||
require.Equal(t, "rate_limit_error", streamErr.ErrType)
|
||||
|
||||
heartbeatRecorder := httptest.NewRecorder()
|
||||
heartbeatContext, _ := gin.CreateTestContext(heartbeatRecorder)
|
||||
heartbeat := ": keepalive\n\n"
|
||||
written, err := heartbeatRecorder.Write([]byte(heartbeat))
|
||||
require.NoError(t, err)
|
||||
recordGatewayStreamHeartbeat(heartbeatContext, written)
|
||||
(&GatewayHandler{}).handleResponsesFailoverExhausted(heartbeatContext, &service.UpstreamFailoverError{
|
||||
StatusCode: http.StatusBadGateway,
|
||||
}, true)
|
||||
require.True(t, strings.HasPrefix(heartbeatRecorder.Body.String(), heartbeat))
|
||||
require.Equal(t, 1, strings.Count(heartbeatRecorder.Body.String(), "event: response.failed"))
|
||||
}
|
||||
|
||||
func TestGatewayChatInferenceExhaustionRestoresRetryAfter(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
recorder := httptest.NewRecorder()
|
||||
|
||||
@@ -24,7 +24,7 @@ func TestRecordCyberPolicyIfMarked_NoMark(t *testing.T) {
|
||||
c := newTestGinContext()
|
||||
h := &OpenAIGatewayHandler{}
|
||||
|
||||
h.recordCyberPolicyIfMarked(c, nil, nil, nil, "gpt-5", true, "", service.ChannelUsageFields{}, "")
|
||||
h.recordCyberPolicyIfMarked(c, nil, nil, nil, "gpt-5", true, nil, service.ChannelUsageFields{}, "")
|
||||
|
||||
// Flag must NOT be set when there was no mark.
|
||||
require.False(t, c.GetBool(cyberPolicyRecordedKey),
|
||||
@@ -47,14 +47,14 @@ func TestRecordCyberPolicyIfMarked_WithMark(t *testing.T) {
|
||||
|
||||
// First call: should set the flag.
|
||||
require.NotPanics(t, func() {
|
||||
h.recordCyberPolicyIfMarked(c, nil, nil, nil, "gpt-5", true, "", service.ChannelUsageFields{}, "")
|
||||
h.recordCyberPolicyIfMarked(c, nil, nil, nil, "gpt-5", true, nil, service.ChannelUsageFields{}, "")
|
||||
})
|
||||
require.True(t, c.GetBool(cyberPolicyRecordedKey),
|
||||
"cyberPolicyRecordedKey must be true after first call with a mark")
|
||||
|
||||
// Second call: flag already set — must be a no-op (idempotent).
|
||||
require.NotPanics(t, func() {
|
||||
h.recordCyberPolicyIfMarked(c, nil, nil, nil, "gpt-5", false, "", service.ChannelUsageFields{}, "")
|
||||
h.recordCyberPolicyIfMarked(c, nil, nil, nil, "gpt-5", false, nil, service.ChannelUsageFields{}, "")
|
||||
})
|
||||
// Flag should still be true (not toggled or cleared).
|
||||
require.True(t, c.GetBool(cyberPolicyRecordedKey),
|
||||
@@ -75,7 +75,7 @@ func TestRecordCyberPolicyIfMarked_ForwardSuccessSkipsUsageLog(t *testing.T) {
|
||||
h := &OpenAIGatewayHandler{}
|
||||
|
||||
require.NotPanics(t, func() {
|
||||
h.recordCyberPolicyIfMarked(c, nil, nil, nil, "gpt-5", false /* forwardErrored=false */, "", service.ChannelUsageFields{}, "")
|
||||
h.recordCyberPolicyIfMarked(c, nil, nil, nil, "gpt-5", false /* forwardErrored=false */, nil, service.ChannelUsageFields{}, "")
|
||||
})
|
||||
require.True(t, c.GetBool(cyberPolicyRecordedKey))
|
||||
}
|
||||
@@ -88,7 +88,7 @@ func TestClearCyberPolicyTurnState(t *testing.T) {
|
||||
h := &OpenAIGatewayHandler{}
|
||||
|
||||
service.MarkOpsCyberPolicy(c, service.CyberPolicyMark{Message: "turn1", UpstreamStatus: 200})
|
||||
h.recordCyberPolicyIfMarked(c, nil, nil, nil, "gpt-5", false, "", service.ChannelUsageFields{}, "")
|
||||
h.recordCyberPolicyIfMarked(c, nil, nil, nil, "gpt-5", false, nil, service.ChannelUsageFields{}, "")
|
||||
require.True(t, c.GetBool(cyberPolicyRecordedKey))
|
||||
|
||||
clearCyberPolicyTurnState(c)
|
||||
@@ -97,7 +97,7 @@ func TestClearCyberPolicyTurnState(t *testing.T) {
|
||||
|
||||
// turn2: a fresh cyber hit must be recordable again.
|
||||
service.MarkOpsCyberPolicy(c, service.CyberPolicyMark{Message: "turn2", UpstreamStatus: 200})
|
||||
h.recordCyberPolicyIfMarked(c, nil, nil, nil, "gpt-5", false, "", service.ChannelUsageFields{}, "")
|
||||
h.recordCyberPolicyIfMarked(c, nil, nil, nil, "gpt-5", false, nil, service.ChannelUsageFields{}, "")
|
||||
require.True(t, c.GetBool(cyberPolicyRecordedKey))
|
||||
require.Equal(t, "turn2", service.GetOpsCyberPolicy(c).Message)
|
||||
}
|
||||
@@ -139,6 +139,23 @@ func TestRejectIfCyberSessionBlocked_FailOpen(t *testing.T) {
|
||||
require.False(t, h2.rejectIfCyberSessionBlocked(c, key, []byte(`{}`), "gpt-5", cyberBlockFormatResponses), "nil gateway service → pass")
|
||||
}
|
||||
|
||||
func TestBuildCyberSessionBlockWritePlanCombinesExplicitAndTranscriptKeys(t *testing.T) {
|
||||
body := []byte(`{"messages":[{"role":"user","content":"setup"},{"role":"assistant","content":"ready"},{"role":"user","content":"trigger"}]}`)
|
||||
c := newTestGinContext()
|
||||
c.Request = httptest.NewRequest("POST", "/openai/v1/responses", strings.NewReader(string(body)))
|
||||
c.Request.RemoteAddr = "203.0.113.44:12345"
|
||||
c.Request.Header.Set("User-Agent", "client/1.2.3")
|
||||
|
||||
plan := buildCyberSessionBlockWritePlan(7, c, body)
|
||||
require.Len(t, plan.keys, 2)
|
||||
require.NotEmpty(t, plan.scopeKey)
|
||||
|
||||
c.Request.Header.Set("session_id", "sess-explicit")
|
||||
plan = buildCyberSessionBlockWritePlan(7, c, body)
|
||||
require.Len(t, plan.keys, 3)
|
||||
require.NotEmpty(t, plan.scopeKey)
|
||||
}
|
||||
|
||||
// TestRecordCyberPolicyIfMarked_BlockKeyPlumbed verifies the 6th param is
|
||||
// accepted and a non-empty key with nil gateway service does not panic
|
||||
// (write-side guards live in the service layer).
|
||||
@@ -147,7 +164,7 @@ func TestRecordCyberPolicyIfMarked_BlockKeyPlumbed(t *testing.T) {
|
||||
service.MarkOpsCyberPolicy(c, service.CyberPolicyMark{Message: "x", UpstreamStatus: 400})
|
||||
h := &OpenAIGatewayHandler{}
|
||||
require.NotPanics(t, func() {
|
||||
h.recordCyberPolicyIfMarked(c, nil, nil, nil, "gpt-5", true, "deadbeef", service.ChannelUsageFields{}, "")
|
||||
h.recordCyberPolicyIfMarked(c, nil, nil, nil, "gpt-5", true, []byte(`{"input":"deadbeef"}`), service.ChannelUsageFields{}, "")
|
||||
})
|
||||
}
|
||||
|
||||
|
||||
@@ -58,7 +58,7 @@ func newOpenAIWSUnsupportedModelSwitchError(model string) error {
|
||||
}
|
||||
|
||||
func shouldReportOpenAIWSProxyAccountFailure(err error) bool {
|
||||
return err != nil && !errors.Is(err, errOpenAIWSUnsupportedModelSwitch)
|
||||
return err != nil && !errors.Is(err, errOpenAIWSUnsupportedModelSwitch) && !service.IsOpenAIWSSessionPreemptedError(err)
|
||||
}
|
||||
|
||||
func openAIWSTurnBillingModel(result *service.OpenAIForwardResult, mapping service.ChannelMappingResult, requestedModel, upstreamModel string) string {
|
||||
@@ -98,6 +98,22 @@ func openAIForwardSucceededForScheduling(result *service.OpenAIForwardResult) bo
|
||||
return result.SucceededForScheduling()
|
||||
}
|
||||
|
||||
func openAIAccountScheduleModel(c *gin.Context, account *service.Account, forwardModel string, requireCompact bool, result *service.OpenAIForwardResult) string {
|
||||
if result != nil {
|
||||
if actual := strings.TrimSpace(result.UpstreamModel); actual != "" {
|
||||
return actual
|
||||
}
|
||||
}
|
||||
if c != nil {
|
||||
if value, ok := c.Get(service.OpsUpstreamModelKey); ok {
|
||||
if actual, ok := value.(string); ok && strings.TrimSpace(actual) != "" {
|
||||
return strings.TrimSpace(actual)
|
||||
}
|
||||
}
|
||||
}
|
||||
return service.ResolveOpenAIAccountUpstreamModelForRequest(account, forwardModel, requireCompact)
|
||||
}
|
||||
|
||||
func resolveOpenAIMessagesDispatchMappedModel(c *gin.Context, apiKey *service.APIKey, requestedModel string) string {
|
||||
if apiKey == nil || apiKey.Group == nil {
|
||||
return ""
|
||||
@@ -393,11 +409,6 @@ func (h *OpenAIGatewayHandler) Responses(c *gin.Context) {
|
||||
h.errorResponse(c, http.StatusBadRequest, "invalid_request_error", "previous_response_id must be a response.id (resp_*), not a message id")
|
||||
return
|
||||
}
|
||||
reqLog.Warn("openai.request_validation_failed",
|
||||
zap.String("reason", "previous_response_id_requires_wsv2"),
|
||||
)
|
||||
h.errorResponse(c, http.StatusBadRequest, "invalid_request_error", "previous_response_id is only supported on Responses WebSocket v2")
|
||||
return
|
||||
}
|
||||
|
||||
setOpsRequestContext(c, reqModel, reqStream)
|
||||
@@ -550,6 +561,7 @@ func (h *OpenAIGatewayHandler) Responses(c *gin.Context) {
|
||||
return
|
||||
}
|
||||
cls := classifyNoAccountErrorFromGin(c, h.gatewayService, apiKey, reqModel, reqModel, requestPlatform)
|
||||
cls = classifySelectionFailureError(err, cls)
|
||||
if !cls.ModelNotFound {
|
||||
markOpsRoutingCapacityLimitedIfNoAvailable(c, err)
|
||||
}
|
||||
@@ -584,6 +596,29 @@ func (h *OpenAIGatewayHandler) Responses(c *gin.Context) {
|
||||
zap.Float64("load_skew", scheduleDecision.LoadSkew),
|
||||
)
|
||||
account := selection.Account
|
||||
if previousResponseID != "" && requestPlatform == service.PlatformOpenAI && !account.IsOpenAIApiKey() {
|
||||
// The public Responses HTTP API supports previous_response_id on API-key
|
||||
// accounts. OAuth/SetupToken upstreams do not, so keep searching instead
|
||||
// of silently deleting continuation state from a mixed account pool.
|
||||
failedAccountIDs[account.ID] = struct{}{}
|
||||
if selection.ReleaseFunc != nil {
|
||||
selection.ReleaseFunc()
|
||||
selection.ReleaseFunc = nil
|
||||
}
|
||||
lastFailoverErr = &service.UpstreamFailoverError{
|
||||
StatusCode: http.StatusBadRequest,
|
||||
Stage: service.GatewayFailureStageInference,
|
||||
Scope: service.GatewayFailureScopeRequest,
|
||||
Reason: service.OpenAIHTTPContinuationUnsupportedReason,
|
||||
ClientStatusCode: http.StatusBadRequest,
|
||||
ClientMessage: "previous_response_id requires an OpenAI API-key account for HTTP requests",
|
||||
}
|
||||
reqLog.Debug("openai.account_skipped_http_continuation_unsupported",
|
||||
zap.Int64("account_id", account.ID),
|
||||
zap.String("account_type", account.Type),
|
||||
)
|
||||
continue
|
||||
}
|
||||
sessionHash = ensureOpenAIPoolModeSessionHash(sessionHash, account)
|
||||
reqLog.Debug("openai.account_selected", zap.Int64("account_id", account.ID), zap.String("account_name", account.Name))
|
||||
setOpsSelectedAccount(c, account.ID, account.Platform)
|
||||
@@ -620,11 +655,11 @@ func (h *OpenAIGatewayHandler) Responses(c *gin.Context) {
|
||||
}()
|
||||
return h.gatewayService.Forward(c.Request.Context(), c, account, attemptBody)
|
||||
}()
|
||||
cyberBlockKeyHTTP := ""
|
||||
var cyberBlockBodyHTTP []byte
|
||||
if service.GetOpsCyberPolicy(c) != nil {
|
||||
cyberBlockKeyHTTP = service.CyberSessionBlockKey(apiKey.ID, c, sessionHashBody)
|
||||
cyberBlockBodyHTTP = sessionHashBody
|
||||
}
|
||||
h.recordCyberPolicyIfMarked(c, apiKey, account, subscription, reqModel, err != nil, cyberBlockKeyHTTP, clientRequestedUsageFields(c, channelMapping, reqModel, ""), service.HashUsageRequestPayload(body))
|
||||
h.recordCyberPolicyIfMarked(c, apiKey, account, subscription, reqModel, err != nil, cyberBlockBodyHTTP, clientRequestedUsageFields(c, channelMapping, reqModel, ""), service.HashUsageRequestPayload(body))
|
||||
forwardDurationMs := time.Since(forwardStart).Milliseconds()
|
||||
upstreamLatencyMs, _ := getContextInt64(c, service.OpsUpstreamLatencyMsKey)
|
||||
responseLatencyMs := forwardDurationMs
|
||||
@@ -697,6 +732,7 @@ func (h *OpenAIGatewayHandler) Responses(c *gin.Context) {
|
||||
return
|
||||
}
|
||||
if !openAIForwardMayFailover(c, writerSizeBeforeForward, failoverErr) {
|
||||
h.gatewayService.ObserveOpenAIAccountHealthFailure(c.Request.Context(), account, err)
|
||||
h.handleFailoverExhausted(c, failoverErr, true)
|
||||
return
|
||||
}
|
||||
@@ -706,7 +742,7 @@ func (h *OpenAIGatewayHandler) Responses(c *gin.Context) {
|
||||
streamStarted = true
|
||||
}
|
||||
if failoverErr.ShouldReportAccountScheduleFailure() {
|
||||
h.gatewayService.ReportOpenAIAccountScheduleResult(account.ID, account.GetMappedModel(reqModel), false, nil)
|
||||
h.gatewayService.ReportOpenAIAccountScheduleResult(account, openAIAccountScheduleModel(c, account, forwardModel, requireCompact, nil), false, nil, err)
|
||||
}
|
||||
if !failoverErr.ShouldRetryNextAccount() {
|
||||
h.handleFailoverExhausted(c, failoverErr, streamStarted)
|
||||
@@ -719,7 +755,7 @@ func (h *OpenAIGatewayHandler) Responses(c *gin.Context) {
|
||||
// 池模式:同账号重试
|
||||
if failoverErr.RetryableOnSameAccount {
|
||||
retryLimit := account.GetPoolModeRetryCount()
|
||||
if sameAccountRetryCount[account.ID] < retryLimit {
|
||||
if sameAccountRetryAllowed(failoverErr, sameAccountRetryCount[account.ID], retryLimit) {
|
||||
sameAccountRetryCount[account.ID]++
|
||||
retryDelay := sameAccountRetryDelayFor(failoverErr, sameAccountRetryCount[account.ID])
|
||||
reqLog.Warn("openai.pool_mode_same_account_retry",
|
||||
@@ -768,7 +804,7 @@ func (h *OpenAIGatewayHandler) Responses(c *gin.Context) {
|
||||
reqLog.Warn("openai.upstream_failover_switching", failoverSwitchFields...)
|
||||
continue
|
||||
}
|
||||
h.gatewayService.ReportOpenAIAccountScheduleResult(account.ID, account.GetMappedModel(reqModel), false, nil)
|
||||
h.gatewayService.ReportOpenAIAccountScheduleResult(account, openAIAccountScheduleModel(c, account, forwardModel, requireCompact, result), false, nil, err)
|
||||
upstreamErrorAlreadyCommunicated := openAIForwardErrorAlreadyCommunicated(c, writerSizeBeforeForward, err)
|
||||
wroteFallback := false
|
||||
if !upstreamErrorAlreadyCommunicated {
|
||||
@@ -794,9 +830,9 @@ func (h *OpenAIGatewayHandler) Responses(c *gin.Context) {
|
||||
if account.Type == service.AccountTypeOAuth && !account.IsShadow() {
|
||||
h.gatewayService.UpdateCodexUsageSnapshotFromHeaders(c.Request.Context(), account.ID, result.ResponseHeaders)
|
||||
}
|
||||
h.gatewayService.ReportOpenAIAccountScheduleResult(account.ID, account.GetMappedModel(reqModel), openAIForwardSucceededForScheduling(result), result.FirstTokenMs)
|
||||
h.gatewayService.ReportOpenAIAccountScheduleResult(account, openAIAccountScheduleModel(c, account, forwardModel, requireCompact, result), openAIForwardSucceededForScheduling(result), result.FirstTokenMs)
|
||||
} else {
|
||||
h.gatewayService.ReportOpenAIAccountScheduleResult(account.ID, account.GetMappedModel(reqModel), openAIForwardSucceededForScheduling(result), nil)
|
||||
h.gatewayService.ReportOpenAIAccountScheduleResult(account, openAIAccountScheduleModel(c, account, forwardModel, requireCompact, result), openAIForwardSucceededForScheduling(result), nil)
|
||||
}
|
||||
|
||||
// 使用量记录通过有界 worker 池提交,避免请求热路径创建无界 goroutine。
|
||||
@@ -1182,11 +1218,11 @@ func (h *OpenAIGatewayHandler) Messages(c *gin.Context) {
|
||||
}()
|
||||
return h.gatewayService.ForwardAsAnthropic(c.Request.Context(), c, account, forwardBody, promptCacheKey, defaultMappedModel)
|
||||
}()
|
||||
cyberBlockKeyMsg := ""
|
||||
var cyberBlockBodyMsg []byte
|
||||
if service.GetOpsCyberPolicy(c) != nil {
|
||||
cyberBlockKeyMsg = service.CyberSessionBlockKey(apiKey.ID, c, body)
|
||||
cyberBlockBodyMsg = body
|
||||
}
|
||||
h.recordCyberPolicyIfMarked(c, apiKey, account, subscription, reqModel, err != nil, cyberBlockKeyMsg, clientRequestedUsageFields(c, channelMappingMsg, reqModel, ""), service.HashUsageRequestPayload(body))
|
||||
h.recordCyberPolicyIfMarked(c, apiKey, account, subscription, reqModel, err != nil, cyberBlockBodyMsg, clientRequestedUsageFields(c, channelMappingMsg, reqModel, ""), service.HashUsageRequestPayload(body))
|
||||
forwardDurationMs := time.Since(forwardStart).Milliseconds()
|
||||
upstreamLatencyMs, _ := getContextInt64(c, service.OpsUpstreamLatencyMsKey)
|
||||
responseLatencyMs := forwardDurationMs
|
||||
@@ -1260,11 +1296,12 @@ func (h *OpenAIGatewayHandler) Messages(c *gin.Context) {
|
||||
return
|
||||
}
|
||||
if c.Writer.Size() != writerSizeBeforeForward {
|
||||
h.gatewayService.ObserveOpenAIAccountHealthFailure(c.Request.Context(), account, err)
|
||||
h.handleAnthropicFailoverExhausted(c, failoverErr, true)
|
||||
return
|
||||
}
|
||||
if failoverErr.ShouldReportAccountScheduleFailure() {
|
||||
h.gatewayService.ReportOpenAIAccountScheduleResult(account.ID, account.GetMappedModel(currentRoutingModel), false, nil)
|
||||
h.gatewayService.ReportOpenAIAccountScheduleResult(account, openAIAccountScheduleModel(c, account, currentRoutingModel, false, nil), false, nil, err)
|
||||
}
|
||||
if !failoverErr.ShouldRetryNextAccount() {
|
||||
h.handleAnthropicFailoverExhausted(c, failoverErr, streamStarted)
|
||||
@@ -1273,7 +1310,7 @@ func (h *OpenAIGatewayHandler) Messages(c *gin.Context) {
|
||||
// 池模式:同账号重试
|
||||
if failoverErr.RetryableOnSameAccount {
|
||||
retryLimit := account.GetPoolModeRetryCount()
|
||||
if sameAccountRetryCount[account.ID] < retryLimit {
|
||||
if sameAccountRetryAllowed(failoverErr, sameAccountRetryCount[account.ID], retryLimit) {
|
||||
sameAccountRetryCount[account.ID]++
|
||||
retryDelay := sameAccountRetryDelayFor(failoverErr, sameAccountRetryCount[account.ID])
|
||||
reqLog.Warn("openai_messages.pool_mode_same_account_retry",
|
||||
@@ -1321,7 +1358,7 @@ func (h *OpenAIGatewayHandler) Messages(c *gin.Context) {
|
||||
submitMessagesUsage(result)
|
||||
return
|
||||
}
|
||||
h.gatewayService.ReportOpenAIAccountScheduleResult(account.ID, account.GetMappedModel(currentRoutingModel), false, nil)
|
||||
h.gatewayService.ReportOpenAIAccountScheduleResult(account, openAIAccountScheduleModel(c, account, currentRoutingModel, false, result), false, nil, err)
|
||||
wroteFallback := h.ensureAnthropicErrorResponse(c, streamStarted)
|
||||
reqLog.Warn("openai_messages.forward_failed",
|
||||
zap.Int64("account_id", account.ID),
|
||||
@@ -1333,9 +1370,9 @@ func (h *OpenAIGatewayHandler) Messages(c *gin.Context) {
|
||||
}
|
||||
}
|
||||
if result != nil {
|
||||
h.gatewayService.ReportOpenAIAccountScheduleResult(account.ID, account.GetMappedModel(currentRoutingModel), true, result.FirstTokenMs)
|
||||
h.gatewayService.ReportOpenAIAccountScheduleResult(account, openAIAccountScheduleModel(c, account, currentRoutingModel, false, result), true, result.FirstTokenMs)
|
||||
} else {
|
||||
h.gatewayService.ReportOpenAIAccountScheduleResult(account.ID, account.GetMappedModel(currentRoutingModel), true, nil)
|
||||
h.gatewayService.ReportOpenAIAccountScheduleResult(account, openAIAccountScheduleModel(c, account, currentRoutingModel, false, result), true, nil)
|
||||
}
|
||||
|
||||
submitMessagesUsage(result)
|
||||
@@ -1403,6 +1440,14 @@ func (h *OpenAIGatewayHandler) handleAnthropicFailoverExhausted(c *gin.Context,
|
||||
h.anthropicStreamingAwareError(c, status, "api_error", message, streamStarted)
|
||||
return
|
||||
}
|
||||
if failoverErr != nil && failoverErr.IsOpenAICapacityShed() && strings.TrimSpace(failoverErr.ClientMessage) != "" {
|
||||
status := failoverErr.ClientStatusCode
|
||||
if status <= 0 {
|
||||
status = http.StatusServiceUnavailable
|
||||
}
|
||||
h.anthropicStreamingAwareError(c, status, "api_error", failoverErr.ClientMessage, streamStarted)
|
||||
return
|
||||
}
|
||||
status, errType, errMsg := h.mapUpstreamError(failoverErr.StatusCode)
|
||||
h.anthropicStreamingAwareError(c, status, errType, errMsg, streamStarted)
|
||||
}
|
||||
@@ -1545,9 +1590,32 @@ func (h *OpenAIGatewayHandler) acquireResponsesAccountSlot(
|
||||
streamStarted *bool,
|
||||
reqLog *zap.Logger,
|
||||
) (func(), openAISlotAcquireResult) {
|
||||
return h.acquireOpenAIAccountSlot(c, groupID, sessionHash, selection, reqStream, streamStarted, reqLog, nil)
|
||||
}
|
||||
|
||||
type openAISlotErrorWriter func(status int, errType, message string)
|
||||
|
||||
// acquireOpenAIAccountSlot centralizes scheduler selection admission. The
|
||||
// optional error writer lets non-Responses endpoints retain their wire format
|
||||
// while sharing the same WaitPlan, cancellation, and release semantics.
|
||||
func (h *OpenAIGatewayHandler) acquireOpenAIAccountSlot(
|
||||
c *gin.Context,
|
||||
groupID *int64,
|
||||
sessionHash string,
|
||||
selection *service.AccountSelectionResult,
|
||||
reqStream bool,
|
||||
streamStarted *bool,
|
||||
reqLog *zap.Logger,
|
||||
writeError openAISlotErrorWriter,
|
||||
) (func(), openAISlotAcquireResult) {
|
||||
if writeError == nil {
|
||||
writeError = func(status int, errType, message string) {
|
||||
h.handleStreamingAwareError(c, status, errType, message, *streamStarted)
|
||||
}
|
||||
}
|
||||
if selection == nil || selection.Account == nil {
|
||||
markOpsRoutingCapacityLimited(c)
|
||||
h.handleStreamingAwareError(c, http.StatusServiceUnavailable, "api_error", "No available accounts", *streamStarted)
|
||||
writeError(http.StatusServiceUnavailable, "api_error", "No available accounts")
|
||||
return nil, openAISlotAcquireFailed
|
||||
}
|
||||
|
||||
@@ -1577,7 +1645,7 @@ func (h *OpenAIGatewayHandler) acquireResponsesAccountSlot(
|
||||
}
|
||||
if selection.WaitPlan == nil {
|
||||
markOpsRoutingCapacityLimited(c)
|
||||
h.handleStreamingAwareError(c, http.StatusServiceUnavailable, "api_error", "No available accounts", *streamStarted)
|
||||
writeError(http.StatusServiceUnavailable, "api_error", "No available accounts")
|
||||
return nil, openAISlotAcquireFailed
|
||||
}
|
||||
|
||||
@@ -1588,7 +1656,8 @@ func (h *OpenAIGatewayHandler) acquireResponsesAccountSlot(
|
||||
)
|
||||
if err != nil {
|
||||
reqLog.Warn("openai.account_slot_quick_acquire_failed", zap.Int64("account_id", account.ID), zap.Error(err))
|
||||
h.handleConcurrencyError(c, err, "account", *streamStarted)
|
||||
status, errType, message := concurrencyErrorResponse(err, "account")
|
||||
writeError(status, errType, message)
|
||||
return nil, openAISlotAcquireFailed
|
||||
}
|
||||
if fastAcquired {
|
||||
@@ -1618,7 +1687,7 @@ func (h *OpenAIGatewayHandler) acquireResponsesAccountSlot(
|
||||
zap.Int64("account_id", account.ID),
|
||||
zap.Int("max_waiting", selection.WaitPlan.MaxWaiting),
|
||||
)
|
||||
h.handleStreamingAwareError(c, http.StatusTooManyRequests, "rate_limit_error", "Too many pending requests, please retry later", *streamStarted)
|
||||
writeError(http.StatusTooManyRequests, "rate_limit_error", "Too many pending requests, please retry later")
|
||||
return nil, openAISlotAcquireFailed
|
||||
}
|
||||
|
||||
@@ -1641,7 +1710,8 @@ func (h *OpenAIGatewayHandler) acquireResponsesAccountSlot(
|
||||
)
|
||||
if err != nil {
|
||||
reqLog.Warn("openai.account_slot_acquire_failed", zap.Int64("account_id", account.ID), zap.Error(err))
|
||||
h.handleConcurrencyError(c, err, "account", *streamStarted)
|
||||
status, errType, message := concurrencyErrorResponse(err, "account")
|
||||
writeError(status, errType, message)
|
||||
return nil, openAISlotAcquireFailed
|
||||
}
|
||||
|
||||
@@ -1820,19 +1890,36 @@ func (h *OpenAIGatewayHandler) ResponsesWebSocket(c *gin.Context) {
|
||||
return
|
||||
}
|
||||
|
||||
// F5a: 握手层会话屏蔽检查。WS 握手无 body,显式标识仅来自握手 header
|
||||
// (session_id / conversation_id);无标识则放行,连接内仍有本地 flag 兜底。
|
||||
cyberBlockKey := service.CyberSessionBlockKey(apiKey.ID, c, nil)
|
||||
if cyberBlockKey != "" && h.gatewayService.IsCyberSessionBlocked(c.Request.Context(), cyberBlockKey) {
|
||||
// The first response.create frame is available here, so explicit IDs are
|
||||
// checked directly and body-derived sessions use the coarse scope gate.
|
||||
if cyberBlockKey := findBlockedCyberSessionKey(c.Request.Context(), h.gatewayService, apiKey.ID, c, firstMessage); cyberBlockKey != "" {
|
||||
writeCyberSessionBlockedWSError(c.Request.Context(), wsConn)
|
||||
closeOpenAIClientWS(wsConn, coderws.StatusPolicyViolation, "session blocked by cyber-security policy")
|
||||
h.enqueueCyberSessionBlockedOpsEntry(c, apiKey, reqModel, cyberBlockKey)
|
||||
return
|
||||
}
|
||||
cyberBlockedThisConn := false
|
||||
var cyberTurnBodiesMu sync.Mutex
|
||||
cyberTurnBodies := map[int][]byte{1: append([]byte(nil), firstMessage...)}
|
||||
setCyberTurnBody := func(turn int, payload []byte) {
|
||||
cyberTurnBodiesMu.Lock()
|
||||
cyberTurnBodies[turn] = append([]byte(nil), payload...)
|
||||
cyberTurnBodiesMu.Unlock()
|
||||
}
|
||||
takeCyberTurnBody := func(turn int) []byte {
|
||||
cyberTurnBodiesMu.Lock()
|
||||
body := cyberTurnBodies[turn]
|
||||
delete(cyberTurnBodies, turn)
|
||||
cyberTurnBodiesMu.Unlock()
|
||||
return body
|
||||
}
|
||||
|
||||
// 解析渠道级模型映射
|
||||
channelMappingWS, _ := h.gatewayService.ResolveChannelMappingAndRestrict(ctx, apiKey.GroupID, reqModel)
|
||||
wsForwardModel := reqModel
|
||||
if channelMappingWS.Mapped && strings.TrimSpace(channelMappingWS.MappedModel) != "" {
|
||||
wsForwardModel = strings.TrimSpace(channelMappingWS.MappedModel)
|
||||
}
|
||||
|
||||
var currentUserRelease func()
|
||||
var currentAccountRelease func()
|
||||
@@ -1902,15 +1989,38 @@ func (h *OpenAIGatewayHandler) ResponsesWebSocket(c *gin.Context) {
|
||||
switchCount := 0
|
||||
profitVetoCount := 0
|
||||
failedAccountIDs := make(map[int64]struct{})
|
||||
sameAccountRetryCount := make(map[int64]int)
|
||||
var lastFailoverErr *service.UpstreamFailoverError
|
||||
var oauth429FailoverState service.OpenAIOAuth429FailoverState
|
||||
wsAttemptMessage := append([]byte(nil), firstMessage...)
|
||||
waitForWSSameAccountRetry := func(account *service.Account, failoverErr *service.UpstreamFailoverError) bool {
|
||||
if account == nil || failoverErr == nil || failoverErr.StatusCode != http.StatusTooManyRequests || failoverErr.SameAccountRetryDeadline.IsZero() {
|
||||
return false
|
||||
}
|
||||
if !sameAccountRetryAllowed(failoverErr, sameAccountRetryCount[account.ID], 0) {
|
||||
return false
|
||||
}
|
||||
sameAccountRetryCount[account.ID]++
|
||||
retryDelay := sameAccountRetryDelayFor(failoverErr, sameAccountRetryCount[account.ID])
|
||||
reqLog.Warn("openai.websocket.same_account_retry",
|
||||
zap.Int64("account_id", account.ID),
|
||||
zap.Int("upstream_status", failoverErr.StatusCode),
|
||||
zap.Int("retry_count", sameAccountRetryCount[account.ID]),
|
||||
zap.Duration("retry_delay", retryDelay),
|
||||
)
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return false
|
||||
case <-time.After(retryDelay):
|
||||
return true
|
||||
}
|
||||
}
|
||||
handleWSFailover := func(account *service.Account, failoverErr *service.UpstreamFailoverError) bool {
|
||||
if ctx.Err() != nil {
|
||||
return false
|
||||
}
|
||||
if failoverErr.ShouldReportAccountScheduleFailure() {
|
||||
h.gatewayService.ReportOpenAIAccountScheduleResult(account.ID, account.GetMappedModel(reqModel), false, nil)
|
||||
h.gatewayService.ReportOpenAIAccountScheduleResult(account, openAIAccountScheduleModel(c, account, wsForwardModel, false, nil), false, nil, failoverErr)
|
||||
}
|
||||
releaseAccountSlot()
|
||||
if !failoverErr.ShouldRetryNextAccount() {
|
||||
@@ -2067,6 +2177,9 @@ func (h *OpenAIGatewayHandler) ResponsesWebSocket(c *gin.Context) {
|
||||
}
|
||||
// 准入完成:门并入连接 ctx,turn 级复核与 failover 重选共用。
|
||||
ctx = admissionCtx
|
||||
// Account selection starts a fresh upstream attempt. Clear any model
|
||||
// captured by the previous failover account before credential lookup.
|
||||
setOpsSelectedAccount(c, account.ID, account.Platform)
|
||||
currentAccountRelease = wrapReleaseOnDone(ctx, accountReleaseFunc)
|
||||
if err := h.gatewayService.BindStickySessionAfterProfitAdmission(ctx, apiKey.GroupID, sessionHash, account.ID); err != nil {
|
||||
reqLog.Warn("openai.websocket_bind_sticky_session_after_profit_admission_failed", zap.Int64("account_id", account.ID), zap.Error(err))
|
||||
@@ -2131,6 +2244,8 @@ func (h *OpenAIGatewayHandler) ResponsesWebSocket(c *gin.Context) {
|
||||
TurnStarted: recordTurnStart,
|
||||
BeforeRequest: func(turn int, payload []byte, originalModel string) error {
|
||||
c.Set(securityAuditWSTurnContextKey, turn)
|
||||
service.BeginOpsStreamTurn(c, turn)
|
||||
setCyberTurnBody(turn, payload)
|
||||
if turn == 1 {
|
||||
return nil
|
||||
}
|
||||
@@ -2155,6 +2270,7 @@ func (h *OpenAIGatewayHandler) ResponsesWebSocket(c *gin.Context) {
|
||||
if model == "" {
|
||||
model = reqModel
|
||||
}
|
||||
setOpsRequestContext(c, model, true)
|
||||
mapping, _ := h.gatewayService.ResolveChannelMappingAndRestrict(ctx, apiKey.GroupID, model)
|
||||
mappedModelUnchanged := false
|
||||
if previous := turnChannelMapping.Load(); previous != nil && previous.turn < turn {
|
||||
@@ -2215,6 +2331,7 @@ func (h *OpenAIGatewayHandler) ResponsesWebSocket(c *gin.Context) {
|
||||
},
|
||||
AfterTurn: func(turn int, result *service.OpenAIForwardResult, turnErr error) {
|
||||
turnStart := getTurnStart(turn)
|
||||
cyberBlockBody := takeCyberTurnBody(turn)
|
||||
// F1: cyber 标记按 turn 生命周期清理——defer 保证任意早返回路径都执行;
|
||||
// CyberBlocked 必须在 submit 前同步预捕获(task 闭包由 worker 池异步执行,
|
||||
// 届时 defer 已清除标记)。
|
||||
@@ -2240,7 +2357,7 @@ func (h *OpenAIGatewayHandler) ResponsesWebSocket(c *gin.Context) {
|
||||
turnUpstreamModel = turnRequestedModel
|
||||
}
|
||||
turnUsageFields := turnMapping.ToUsageFields(turnRequestedModel, turnUpstreamModel)
|
||||
h.recordCyberPolicyIfMarked(c, apiKey, account, subscription, turnRequestedModel, turnErr != nil, cyberBlockKey, turnUsageFields, requestPayloadHash)
|
||||
h.recordCyberPolicyIfMarked(c, apiKey, account, subscription, turnRequestedModel, turnErr != nil, cyberBlockBody, turnUsageFields, requestPayloadHash)
|
||||
if service.GetOpsCyberPolicy(c) != nil {
|
||||
cyberBlockedThisConn = true
|
||||
}
|
||||
@@ -2277,7 +2394,7 @@ func (h *OpenAIGatewayHandler) ResponsesWebSocket(c *gin.Context) {
|
||||
if scheduleModel == "" {
|
||||
scheduleModel = turnRequestedModel
|
||||
}
|
||||
h.gatewayService.ReportOpenAIAccountScheduleResult(account.ID, scheduleModel, openAIForwardSucceededForScheduling(result), result.FirstTokenMs)
|
||||
h.gatewayService.ReportOpenAIAccountScheduleResult(account, scheduleModel, openAIForwardSucceededForScheduling(result), result.FirstTokenMs)
|
||||
inboundEndpoint := GetInboundEndpoint(c)
|
||||
upstreamEndpoint := resolveOpenAIUpstreamEndpoint(c, account, result)
|
||||
quotaPlatform := service.QuotaPlatform(c.Request.Context(), apiKey)
|
||||
@@ -2329,7 +2446,15 @@ func (h *OpenAIGatewayHandler) ResponsesWebSocket(c *gin.Context) {
|
||||
// WebSocket 首包可能很大,hash 必须在 hooks 外算成字符串,避免 AfterTurn 闭包保活请求体。
|
||||
requestPayloadHash = service.HashUsageRequestPayload(wsFirstMessage)
|
||||
|
||||
if err := h.gatewayService.ProxyResponsesWebSocketFromClient(ctx, c, wsConn, account, token, wsFirstMessage, hooks); err != nil {
|
||||
for {
|
||||
err := h.gatewayService.ProxyResponsesWebSocketFromClient(ctx, c, wsConn, account, token, wsFirstMessage, hooks)
|
||||
if err == nil {
|
||||
reqLog.Info("openai.websocket_ingress_closed", zap.Int64("account_id", account.ID))
|
||||
return
|
||||
}
|
||||
if service.IsOpenAIWSSessionPreemptedError(err) {
|
||||
return
|
||||
}
|
||||
var failoverErr *service.UpstreamFailoverError
|
||||
if errors.As(err, &failoverErr) {
|
||||
retryPayload, retryCurrentTurn := service.OpenAIWSCurrentTurnRetryPayload(err)
|
||||
@@ -2347,9 +2472,31 @@ func (h *OpenAIGatewayHandler) ResponsesWebSocket(c *gin.Context) {
|
||||
zap.Int("retry_payload_bytes", len(retryPayload)),
|
||||
)
|
||||
}
|
||||
if handleWSFailover(account, failoverErr) {
|
||||
if waitForWSSameAccountRetry(account, failoverErr) {
|
||||
if failoverErr.ShouldReportAccountScheduleFailure() {
|
||||
h.gatewayService.ReportOpenAIAccountScheduleResult(account, openAIAccountScheduleModel(c, account, wsForwardModel, false, nil), false, nil, err)
|
||||
}
|
||||
if !ensureUserSlotHeld() {
|
||||
return
|
||||
}
|
||||
if currentAccountRelease == nil {
|
||||
accountRelease, acquired, acquireErr := h.concurrencyHelper.TryAcquireAccountSlot(ctx, account.ID, accountMaxConcurrency)
|
||||
if acquireErr != nil || !acquired {
|
||||
reqLog.Warn("openai.websocket_same_account_retry_slot_unavailable",
|
||||
zap.Int64("account_id", account.ID),
|
||||
zap.Error(acquireErr),
|
||||
)
|
||||
closeOpenAIClientWS(wsConn, coderws.StatusTryAgainLater, "account is busy, please retry later")
|
||||
return
|
||||
}
|
||||
currentAccountRelease = wrapReleaseOnDone(ctx, accountRelease)
|
||||
}
|
||||
wsFirstMessage = wsAttemptMessage
|
||||
continue
|
||||
}
|
||||
if handleWSFailover(account, failoverErr) {
|
||||
break
|
||||
}
|
||||
return
|
||||
}
|
||||
|
||||
@@ -2373,7 +2520,7 @@ func (h *OpenAIGatewayHandler) ResponsesWebSocket(c *gin.Context) {
|
||||
}
|
||||
|
||||
if shouldReportOpenAIWSProxyAccountFailure(err) {
|
||||
h.gatewayService.ReportOpenAIAccountScheduleResult(account.ID, account.GetMappedModel(reqModel), false, nil)
|
||||
h.gatewayService.ReportOpenAIAccountScheduleResult(account, openAIAccountScheduleModel(c, account, wsForwardModel, false, nil), false, nil, err)
|
||||
}
|
||||
closeStatus, closeReason := summarizeWSCloseErrorForLog(err)
|
||||
proxyFailedFields := []zap.Field{
|
||||
@@ -2400,8 +2547,6 @@ func (h *OpenAIGatewayHandler) ResponsesWebSocket(c *gin.Context) {
|
||||
closeOpenAIClientWS(wsConn, coderws.StatusInternalError, "upstream websocket proxy failed")
|
||||
return
|
||||
}
|
||||
reqLog.Info("openai.websocket_ingress_closed", zap.Int64("account_id", account.ID))
|
||||
return
|
||||
}
|
||||
|
||||
}
|
||||
@@ -2617,12 +2762,28 @@ func (h *OpenAIGatewayHandler) handleFailoverExhausted(c *gin.Context, failoverE
|
||||
)
|
||||
return
|
||||
}
|
||||
if failoverErr.Reason == service.OpenAIHTTPContinuationUnsupportedReason {
|
||||
message := strings.TrimSpace(failoverErr.ClientMessage)
|
||||
if message == "" {
|
||||
message = "previous_response_id requires an OpenAI API-key account for HTTP requests"
|
||||
}
|
||||
h.handleStreamingAwareError(c, http.StatusBadRequest, "invalid_request_error", message, streamStarted)
|
||||
return
|
||||
}
|
||||
copyFailoverRetryAfter(c, failoverErr.ResponseHeaders)
|
||||
if failoverErr.IsCredentialFailure() {
|
||||
status, message := credentialFailoverClientResponse(failoverErr)
|
||||
h.handleStreamingAwareError(c, status, "upstream_error", message, streamStarted)
|
||||
return
|
||||
}
|
||||
if failoverErr.IsOpenAICapacityShed() && strings.TrimSpace(failoverErr.ClientMessage) != "" {
|
||||
status := failoverErr.ClientStatusCode
|
||||
if status <= 0 {
|
||||
status = http.StatusServiceUnavailable
|
||||
}
|
||||
h.handleStreamingAwareError(c, status, "server_error", failoverErr.ClientMessage, streamStarted)
|
||||
return
|
||||
}
|
||||
statusCode := failoverErr.StatusCode
|
||||
responseBody := failoverErr.ResponseBody
|
||||
if service.IsOpenAISilentRefusalErrorBody(responseBody) {
|
||||
@@ -2665,6 +2826,13 @@ func (h *OpenAIGatewayHandler) handleFailoverExhausted(c *gin.Context, failoverE
|
||||
}
|
||||
|
||||
func credentialFailoverClientResponse(failoverErr *service.UpstreamFailoverError) (int, string) {
|
||||
if failoverErr != nil && failoverErr.Reason == service.OpenAIUpstreamAccessStateReason && strings.TrimSpace(failoverErr.ClientMessage) != "" {
|
||||
status := failoverErr.ClientStatusCode
|
||||
if status <= 0 {
|
||||
status = http.StatusServiceUnavailable
|
||||
}
|
||||
return status, failoverErr.ClientMessage
|
||||
}
|
||||
if failoverErr != nil && failoverErr.Reason == service.AntigravityCredentialRejectedReason {
|
||||
return http.StatusBadGateway, service.AntigravityCredentialRejectedClientMessage
|
||||
}
|
||||
@@ -3231,13 +3399,10 @@ func (h *OpenAIGatewayHandler) rejectIfCyberSessionBlocked(c *gin.Context, apiKe
|
||||
if enabled, _ := h.gatewayService.CyberSessionBlockRuntime(c.Request.Context()); !enabled {
|
||||
return false
|
||||
}
|
||||
key := service.CyberSessionBlockKey(apiKey.ID, c, body)
|
||||
key := findBlockedCyberSessionKey(c.Request.Context(), h.gatewayService, apiKey.ID, c, body)
|
||||
if key == "" {
|
||||
return false
|
||||
}
|
||||
if !h.gatewayService.IsCyberSessionBlocked(c.Request.Context(), key) {
|
||||
return false
|
||||
}
|
||||
// body-signal compact 心跳可能已把响应头提交为 200(cyber 检查在用户槽位
|
||||
// 长等待之后执行):以 response.failed 终止事件回传;未提交时停拍后照常
|
||||
// 写 JSON(#3887)。
|
||||
@@ -3265,12 +3430,56 @@ func (h *OpenAIGatewayHandler) rejectIfCyberSessionBlocked(c *gin.Context, apiKe
|
||||
return true
|
||||
}
|
||||
|
||||
type cyberSessionBlockWritePlan struct {
|
||||
scopeKey string
|
||||
keys []string
|
||||
}
|
||||
|
||||
func buildCyberSessionBlockWritePlan(apiKeyID int64, c *gin.Context, body []byte) cyberSessionBlockWritePlan {
|
||||
plan := cyberSessionBlockWritePlan{}
|
||||
if key := service.CyberSessionExplicitBlockKey(apiKeyID, c, body); key != "" {
|
||||
plan.keys = append(plan.keys, key)
|
||||
}
|
||||
transcriptKeys := service.CyberSessionTranscriptBlockKeys(apiKeyID, body)
|
||||
for _, key := range transcriptKeys {
|
||||
if len(plan.keys) == 0 || key != plan.keys[0] {
|
||||
plan.keys = append(plan.keys, key)
|
||||
}
|
||||
}
|
||||
if len(transcriptKeys) > 0 {
|
||||
plan.scopeKey = cyberSessionScopeKey(apiKeyID, c)
|
||||
}
|
||||
return plan
|
||||
}
|
||||
|
||||
func findBlockedCyberSessionKey(ctx context.Context, gatewayService *service.OpenAIGatewayService, apiKeyID int64, c *gin.Context, body []byte) string {
|
||||
if gatewayService == nil {
|
||||
return ""
|
||||
}
|
||||
clientIP, userAgent := "", ""
|
||||
if c != nil {
|
||||
clientIP = strings.TrimSpace(ip.GetClientIP(c))
|
||||
userAgent = c.GetHeader("User-Agent")
|
||||
}
|
||||
return gatewayService.FindCyberSessionBlockedForRequest(ctx, apiKeyID, c, body, clientIP, userAgent)
|
||||
}
|
||||
|
||||
func cyberSessionScopeKey(apiKeyID int64, c *gin.Context) string {
|
||||
if c == nil {
|
||||
return ""
|
||||
}
|
||||
return service.CyberSessionScopeKey(apiKeyID, strings.TrimSpace(ip.GetClientIP(c)), c.GetHeader("User-Agent"))
|
||||
}
|
||||
|
||||
// enqueueCyberSessionBlockedOpsEntry captures request meta and enqueues the
|
||||
// ops_error_logs entry for a locally blocked request.
|
||||
func (h *OpenAIGatewayHandler) enqueueCyberSessionBlockedOpsEntry(c *gin.Context, apiKey *service.APIKey, model string, sessionBlockKey string) {
|
||||
if h.opsService == nil {
|
||||
return
|
||||
}
|
||||
// The dedicated cyber_session_blocked entry owns Ops semantics for this
|
||||
// request; suppress the generic middleware record of the same 403 response.
|
||||
c.Set(opsDedicatedErrorRecordedKey, true)
|
||||
meta := cyberPolicyOpsErrorMeta{Model: model, InboundEndpoint: GetInboundEndpoint(c), CreatedAt: time.Now(), SessionBlockKey: sessionBlockKey}
|
||||
meta.RequestID = c.Writer.Header().Get("X-Request-Id")
|
||||
if c.Request != nil && c.Request.URL != nil {
|
||||
@@ -3304,7 +3513,7 @@ func (h *OpenAIGatewayHandler) enqueueCyberSessionBlockedOpsEntry(c *gin.Context
|
||||
// 并在 forward 返回错误时写一条 tokens=0 用量行。标记由 gateway 服务层在透传 cyber 后设置;
|
||||
// 当前请求已发给用户,本方法只做事后记录,不影响响应。forwardErrored 为 true 时才写用量行,
|
||||
// 避免与正常 RecordUsage(forward 成功路径)重复。每请求至多记录一次。
|
||||
func (h *OpenAIGatewayHandler) recordCyberPolicyIfMarked(c *gin.Context, apiKey *service.APIKey, account *service.Account, subscription *service.UserSubscription, model string, forwardErrored bool, cyberBlockKey string, channelFields service.ChannelUsageFields, requestPayloadHash string) {
|
||||
func (h *OpenAIGatewayHandler) recordCyberPolicyIfMarked(c *gin.Context, apiKey *service.APIKey, account *service.Account, subscription *service.UserSubscription, model string, forwardErrored bool, cyberBlockBody []byte, channelFields service.ChannelUsageFields, requestPayloadHash string) {
|
||||
mark := service.GetOpsCyberPolicy(c)
|
||||
if mark == nil {
|
||||
return
|
||||
@@ -3386,6 +3595,14 @@ func (h *OpenAIGatewayHandler) recordCyberPolicyIfMarked(c *gin.Context, apiKey
|
||||
ClientIP: clientIPStr,
|
||||
CreatedAt: time.Now(),
|
||||
}
|
||||
if gwSvc != nil && apiKey != nil {
|
||||
plan := buildCyberSessionBlockWritePlan(apiKey.ID, c, cyberBlockBody)
|
||||
if len(plan.keys) > 0 {
|
||||
blockCtx, cancel := context.WithTimeout(context.Background(), 500*time.Millisecond)
|
||||
gwSvc.MarkCyberSessionBlocked(blockCtx, plan.scopeKey, plan.keys)
|
||||
cancel()
|
||||
}
|
||||
}
|
||||
go func() {
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second)
|
||||
defer cancel()
|
||||
@@ -3427,9 +3644,6 @@ func (h *OpenAIGatewayHandler) recordCyberPolicyIfMarked(c *gin.Context, apiKey
|
||||
ChannelUsageFields: channelFields,
|
||||
})
|
||||
}
|
||||
if gwSvc != nil && cyberBlockKey != "" {
|
||||
gwSvc.MarkCyberSessionBlocked(ctx, cyberBlockKey)
|
||||
}
|
||||
if opsSvc != nil {
|
||||
enqueueOpsErrorLog(opsSvc, buildCyberPolicyOpsErrorEntry(opsMeta, mark))
|
||||
}
|
||||
|
||||
@@ -855,7 +855,7 @@ func TestOpenAIResponses_RejectsMessageIDAsPreviousResponseID(t *testing.T) {
|
||||
require.Contains(t, w.Body.String(), "previous_response_id must be a response.id")
|
||||
}
|
||||
|
||||
func TestOpenAIResponses_RejectsHTTPContinuationPreviousResponseID(t *testing.T) {
|
||||
func TestOpenAIResponses_AcceptsHTTPContinuationPreviousResponseIDBeforeRouting(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
|
||||
w := httptest.NewRecorder()
|
||||
@@ -879,9 +879,8 @@ func TestOpenAIResponses_RejectsHTTPContinuationPreviousResponseID(t *testing.T)
|
||||
h := newOpenAIHandlerForPreviousResponseIDValidation(t, nil)
|
||||
h.Responses(c)
|
||||
|
||||
require.Equal(t, http.StatusBadRequest, w.Code)
|
||||
require.Contains(t, w.Body.String(), "Responses WebSocket v2")
|
||||
require.Contains(t, w.Body.String(), "previous_response_id")
|
||||
require.NotEqual(t, http.StatusBadRequest, w.Code)
|
||||
require.NotContains(t, w.Body.String(), "Responses WebSocket v2")
|
||||
}
|
||||
|
||||
func TestOpenAIResponses_FunctionCallOutputHTTPGuidanceDoesNotSuggestPreviousResponseReuse(t *testing.T) {
|
||||
@@ -1412,8 +1411,8 @@ func TestOpenAIResponsesWebSocket_PassthroughTracksModelPerTurn(t *testing.T) {
|
||||
})
|
||||
|
||||
require.Len(t, got.upstreamPayloads, 2)
|
||||
require.Equal(t, "gpt-5.6-sol", gjson.GetBytes(got.upstreamPayloads[0], "model").String())
|
||||
require.Equal(t, "gpt-5.6-terra", gjson.GetBytes(got.upstreamPayloads[1], "model").String())
|
||||
require.Equal(t, "sol-channel", gjson.GetBytes(got.upstreamPayloads[0], "model").String())
|
||||
require.Equal(t, "terra-channel", gjson.GetBytes(got.upstreamPayloads[1], "model").String())
|
||||
require.Len(t, got.clientEvents, 2)
|
||||
require.Equal(t, "sol", gjson.GetBytes(got.clientEvents[0], "response.model").String())
|
||||
require.Equal(t, "terra", gjson.GetBytes(got.clientEvents[1], "response.model").String())
|
||||
@@ -1422,16 +1421,16 @@ func TestOpenAIResponsesWebSocket_PassthroughTracksModelPerTurn(t *testing.T) {
|
||||
require.Equal(t, "sol", got.logs[0].Model)
|
||||
require.Equal(t, "sol", got.logs[0].RequestedModel)
|
||||
require.NotNil(t, got.logs[0].UpstreamModel)
|
||||
require.Equal(t, "gpt-5.6-sol", *got.logs[0].UpstreamModel)
|
||||
require.Equal(t, "sol-channel", *got.logs[0].UpstreamModel)
|
||||
require.NotNil(t, got.logs[0].ModelMappingChain)
|
||||
require.Equal(t, "sol→sol-channel→gpt-5.6-sol", *got.logs[0].ModelMappingChain)
|
||||
require.Equal(t, "sol→sol-channel", *got.logs[0].ModelMappingChain)
|
||||
|
||||
require.Equal(t, "terra", got.logs[1].Model)
|
||||
require.Equal(t, "terra", got.logs[1].RequestedModel)
|
||||
require.NotNil(t, got.logs[1].UpstreamModel)
|
||||
require.Equal(t, "gpt-5.6-terra", *got.logs[1].UpstreamModel)
|
||||
require.Equal(t, "terra-channel", *got.logs[1].UpstreamModel)
|
||||
require.NotNil(t, got.logs[1].ModelMappingChain)
|
||||
require.Equal(t, "terra→terra-channel→gpt-5.6-terra", *got.logs[1].ModelMappingChain)
|
||||
require.Equal(t, "terra→terra-channel", *got.logs[1].ModelMappingChain)
|
||||
require.InDelta(t, got.logs[1].TotalCost*2.5, got.logs[0].TotalCost, 1e-12,
|
||||
"each turn must be billed with its own channel-mapped model")
|
||||
}
|
||||
@@ -1593,6 +1592,29 @@ func TestOpenAIWSTurnBillingModelPreservesImagePricingModel(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestOpenAIAccountScheduleModelUsesActualOrSharedResolver(t *testing.T) {
|
||||
account := &service.Account{
|
||||
Platform: service.PlatformOpenAI,
|
||||
Type: service.AccountTypeOAuth,
|
||||
Credentials: map[string]any{
|
||||
"model_mapping": map[string]any{"public": "billing"},
|
||||
"compact_model_mapping": map[string]any{"public": "compact-actual"},
|
||||
},
|
||||
}
|
||||
|
||||
reported := &service.OpenAIForwardResult{UpstreamModel: "observed-actual"}
|
||||
require.Equal(t, "observed-actual", openAIAccountScheduleModel(nil, account, "public", true, reported))
|
||||
require.Equal(t, "compact-actual", openAIAccountScheduleModel(nil, account, "public", true, nil))
|
||||
require.Equal(t, "billing", openAIAccountScheduleModel(nil, account, "public", false, nil))
|
||||
|
||||
c, _ := gin.CreateTestContext(nil)
|
||||
service.SetOpsUpstreamModel(c, "attempt-actual")
|
||||
require.Equal(t, "attempt-actual", openAIAccountScheduleModel(c, account, "public", true, nil))
|
||||
|
||||
setOpsSelectedAccount(c, account.ID, account.Platform)
|
||||
require.Equal(t, "attempt-actual", openAIAccountScheduleModel(c, account, "public", true, nil))
|
||||
}
|
||||
|
||||
func TestShouldReportOpenAIWSProxyAccountFailure(t *testing.T) {
|
||||
t.Run("unsupported client model switch does not penalize account", func(t *testing.T) {
|
||||
err := fmt.Errorf("wrapped ingress turn: %w", newOpenAIWSUnsupportedModelSwitchError("gpt-unsupported"))
|
||||
|
||||
@@ -271,7 +271,11 @@ func (h *OpenAIGatewayHandler) Images(c *gin.Context) {
|
||||
var imageUpstreamErr *service.OpenAIImagesUpstreamError
|
||||
if errors.As(err, &imageUpstreamErr) {
|
||||
retryableServerError := service.IsOpenAIImagesRetryableUpstreamError(imageUpstreamErr)
|
||||
h.gatewayService.ReportOpenAIAccountScheduleResult(account.ID, account.GetMappedModel(requestModel), !retryableServerError, nil)
|
||||
if retryableServerError {
|
||||
h.gatewayService.ReportOpenAIAccountScheduleResult(account, openAIAccountScheduleModel(c, account, requestModel, false, result), false, nil, err)
|
||||
} else {
|
||||
h.gatewayService.ReportOpenAIAccountScheduleResult(account, openAIAccountScheduleModel(c, account, requestModel, false, result), true, nil)
|
||||
}
|
||||
logEvent := "openai.images.upstream_user_error"
|
||||
if retryableServerError {
|
||||
logEvent = "openai.images.upstream_server_error_after_flush"
|
||||
@@ -287,7 +291,7 @@ func (h *OpenAIGatewayHandler) Images(c *gin.Context) {
|
||||
}
|
||||
var failoverErr *service.UpstreamFailoverError
|
||||
if errors.As(err, &failoverErr) {
|
||||
h.gatewayService.ReportOpenAIAccountScheduleResult(account.ID, account.GetMappedModel(requestModel), false, nil)
|
||||
h.gatewayService.ReportOpenAIAccountScheduleResult(account, openAIAccountScheduleModel(c, account, requestModel, false, result), false, nil, err)
|
||||
if service.OpenAIImagesJSONKeepaliveAdjustedWrittenSize(c) != writerSizeBeforeForward {
|
||||
reqLog.Warn("openai.images.upstream_failover_skipped_after_flush",
|
||||
zap.Int64("account_id", account.ID),
|
||||
@@ -305,18 +309,20 @@ func (h *OpenAIGatewayHandler) Images(c *gin.Context) {
|
||||
}
|
||||
if failoverErr.RetryableOnSameAccount {
|
||||
retryLimit := account.GetPoolModeRetryCount()
|
||||
if sameAccountRetryCount[account.ID] < retryLimit {
|
||||
if sameAccountRetryAllowed(failoverErr, sameAccountRetryCount[account.ID], retryLimit) {
|
||||
sameAccountRetryCount[account.ID]++
|
||||
retryDelay := sameAccountRetryDelayFor(failoverErr, sameAccountRetryCount[account.ID])
|
||||
reqLog.Warn("openai.images.pool_mode_same_account_retry",
|
||||
zap.Int64("account_id", account.ID),
|
||||
zap.Int("upstream_status", failoverErr.StatusCode),
|
||||
zap.Int("retry_limit", retryLimit),
|
||||
zap.Int("retry_count", sameAccountRetryCount[account.ID]),
|
||||
zap.Duration("retry_delay", retryDelay),
|
||||
)
|
||||
select {
|
||||
case <-requestCtx.Done():
|
||||
return
|
||||
case <-time.After(sameAccountRetryDelay):
|
||||
case <-time.After(retryDelay):
|
||||
}
|
||||
continue
|
||||
}
|
||||
@@ -341,7 +347,7 @@ func (h *OpenAIGatewayHandler) Images(c *gin.Context) {
|
||||
)
|
||||
continue
|
||||
}
|
||||
h.gatewayService.ReportOpenAIAccountScheduleResult(account.ID, account.GetMappedModel(requestModel), false, nil)
|
||||
h.gatewayService.ReportOpenAIAccountScheduleResult(account, openAIAccountScheduleModel(c, account, requestModel, false, result), false, nil, err)
|
||||
upstreamErrorAlreadyCommunicated := openAIForwardErrorAlreadyCommunicated(c, writerSizeBeforeForward, err)
|
||||
wroteFallback := false
|
||||
if !upstreamErrorAlreadyCommunicated {
|
||||
@@ -366,9 +372,9 @@ func (h *OpenAIGatewayHandler) Images(c *gin.Context) {
|
||||
if account.Type == service.AccountTypeOAuth && !account.IsShadow() {
|
||||
h.gatewayService.UpdateCodexUsageSnapshotFromHeaders(c.Request.Context(), account.ID, result.ResponseHeaders)
|
||||
}
|
||||
h.gatewayService.ReportOpenAIAccountScheduleResult(account.ID, account.GetMappedModel(requestModel), true, result.FirstTokenMs)
|
||||
h.gatewayService.ReportOpenAIAccountScheduleResult(account, openAIAccountScheduleModel(c, account, requestModel, false, result), true, result.FirstTokenMs)
|
||||
} else {
|
||||
h.gatewayService.ReportOpenAIAccountScheduleResult(account.ID, account.GetMappedModel(requestModel), true, nil)
|
||||
h.gatewayService.ReportOpenAIAccountScheduleResult(account, openAIAccountScheduleModel(c, account, requestModel, false, result), true, nil)
|
||||
}
|
||||
|
||||
userAgent := c.GetHeader("User-Agent")
|
||||
|
||||
@@ -1,12 +1,17 @@
|
||||
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) {
|
||||
@@ -19,3 +24,81 @@ 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())
|
||||
}
|
||||
|
||||
@@ -12,9 +12,28 @@ import (
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
type deterministicOpsCaptureWriterStatePool struct {
|
||||
states []*opsCaptureWriterState
|
||||
}
|
||||
|
||||
func (p *deterministicOpsCaptureWriterStatePool) Get() any {
|
||||
if len(p.states) == 0 {
|
||||
return &opsCaptureWriterState{limit: opsCaptureWriterLimit}
|
||||
}
|
||||
last := len(p.states) - 1
|
||||
state := p.states[last]
|
||||
p.states = p.states[:last]
|
||||
return state
|
||||
}
|
||||
|
||||
func (p *deterministicOpsCaptureWriterStatePool) Put(value any) {
|
||||
if state, ok := value.(*opsCaptureWriterState); ok && state != nil {
|
||||
p.states = append(p.states, state)
|
||||
}
|
||||
}
|
||||
|
||||
func TestOpsCaptureWriter_NilInnerWriter_NoPanic(t *testing.T) {
|
||||
w := &opsCaptureWriter{}
|
||||
w.ResponseWriter = nil
|
||||
|
||||
assert.NotPanics(t, func() {
|
||||
assert.Equal(t, 0, w.Status())
|
||||
@@ -88,3 +107,40 @@ func TestOpsCaptureWriter_CompactKeepaliveRestoresOriginalWriter(t *testing.T) {
|
||||
require.Equal(t, http.StatusOK, outerStatus)
|
||||
require.Equal(t, http.StatusOK, recorder.Code)
|
||||
}
|
||||
|
||||
func TestOpsCaptureWriter_StaleLeaseCannotReachReacquiredState(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
pool := &deterministicOpsCaptureWriterStatePool{}
|
||||
|
||||
firstRecorder := httptest.NewRecorder()
|
||||
firstContext, _ := gin.CreateTestContext(firstRecorder)
|
||||
stale := acquireOpsCaptureWriterFromPool(pool, firstContext.Writer)
|
||||
releaseOpsCaptureWriter(stale)
|
||||
|
||||
secondRecorder := httptest.NewRecorder()
|
||||
secondContext, _ := gin.CreateTestContext(secondRecorder)
|
||||
current := acquireOpsCaptureWriterFromPool(pool, secondContext.Writer)
|
||||
defer releaseOpsCaptureWriter(current)
|
||||
require.NotSame(t, stale, current)
|
||||
require.Same(t, stale.state, current.state)
|
||||
|
||||
current.WriteHeader(http.StatusInternalServerError)
|
||||
_, err := current.WriteString("current")
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, []byte("current"), current.capturedBytes())
|
||||
|
||||
n, err := stale.WriteString("stale")
|
||||
require.NoError(t, err)
|
||||
require.Zero(t, n)
|
||||
require.Nil(t, stale.capturedBytes())
|
||||
require.Equal(t, []byte("current"), current.capturedBytes())
|
||||
require.NotContains(t, secondRecorder.Body.String(), "stale")
|
||||
|
||||
// Releasing the stale handle must not return an active state to the pool.
|
||||
releaseOpsCaptureWriter(stale)
|
||||
thirdRecorder := httptest.NewRecorder()
|
||||
thirdContext, _ := gin.CreateTestContext(thirdRecorder)
|
||||
other := acquireOpsCaptureWriterFromPool(pool, thirdContext.Writer)
|
||||
defer releaseOpsCaptureWriter(other)
|
||||
require.NotSame(t, current.state, other.state)
|
||||
}
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -188,21 +188,23 @@ func TestOpsCaptureWriterPool_ResetOnRelease(t *testing.T) {
|
||||
|
||||
writer := acquireOpsCaptureWriter(c.Writer)
|
||||
require.NotNil(t, writer)
|
||||
_, err := writer.buf.WriteString("temp-error-body")
|
||||
c.Writer.WriteHeader(http.StatusInternalServerError)
|
||||
_, err := writer.WriteString("temp-error-body")
|
||||
require.NoError(t, err)
|
||||
require.NotEmpty(t, writer.capturedBytes())
|
||||
|
||||
releaseOpsCaptureWriter(writer)
|
||||
|
||||
reused := acquireOpsCaptureWriter(c.Writer)
|
||||
defer releaseOpsCaptureWriter(reused)
|
||||
|
||||
require.Zero(t, reused.buf.Len(), "writer should be reset before reuse")
|
||||
require.Empty(t, reused.capturedBytes(), "writer should be reset before reuse")
|
||||
}
|
||||
|
||||
func TestOpsCaptureWriterPool_DropsLargeBuffers(t *testing.T) {
|
||||
w := &opsCaptureWriter{}
|
||||
w.buf.Grow(opsCaptureWriterPoolMaxRetainedCapacity + 1)
|
||||
require.False(t, shouldPoolOpsCaptureWriter(w))
|
||||
state := &opsCaptureWriterState{}
|
||||
state.buf.Grow(opsCaptureWriterPoolMaxRetainedCapacity + 1)
|
||||
require.False(t, shouldPoolOpsCaptureWriterState(state))
|
||||
}
|
||||
|
||||
func TestEnqueueOpsErrorLog_SanitizesAndBoundsBodyBeforeQueue(t *testing.T) {
|
||||
@@ -280,6 +282,52 @@ func TestOpsErrorLoggerMiddleware_HardSkipsIngressRejection(t *testing.T) {
|
||||
require.Zero(t, OpsErrorLogEnqueuedTotal(), "ingress rejection must not enter the error queue")
|
||||
}
|
||||
|
||||
func TestOpsErrorLoggerMiddleware_DedicatedCyberSessionBlockRecordsExactlyOnce(t *testing.T) {
|
||||
setupOpsErrorLogTestQueue(t, 3)
|
||||
gin.SetMode(gin.TestMode)
|
||||
ops := service.NewOpsService(nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil)
|
||||
h := &OpenAIGatewayHandler{opsService: ops}
|
||||
apiKey := &service.APIKey{ID: 41, Key: "sk-dedicated-test"}
|
||||
router := gin.New()
|
||||
router.Use(OpsErrorLoggerMiddleware(ops))
|
||||
router.POST("/v1/responses", func(c *gin.Context) {
|
||||
h.enqueueCyberSessionBlockedOpsEntry(c, apiKey, "gpt-test", "session-block-hash")
|
||||
c.JSON(http.StatusForbidden, gin.H{"error": gin.H{
|
||||
"type": "permission_error", "code": "session_blocked_by_cyber_policy", "message": "blocked",
|
||||
}})
|
||||
})
|
||||
|
||||
recorder := httptest.NewRecorder()
|
||||
router.ServeHTTP(recorder, httptest.NewRequest(http.MethodPost, "/v1/responses", nil))
|
||||
|
||||
require.Equal(t, http.StatusForbidden, recorder.Code)
|
||||
require.Equal(t, int64(1), OpsErrorLogQueueLength())
|
||||
job := <-opsErrorLogQueue
|
||||
require.Equal(t, "cyber_policy_session_blocked", job.entry.ErrorType)
|
||||
require.Equal(t, http.StatusForbidden, job.entry.StatusCode)
|
||||
}
|
||||
|
||||
func TestOpsErrorLoggerMiddleware_OrdinaryPermissionStillRecords(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.JSON(http.StatusForbidden, gin.H{"error": gin.H{
|
||||
"type": "permission_error", "code": "permission_denied", "message": "denied",
|
||||
}})
|
||||
})
|
||||
|
||||
recorder := httptest.NewRecorder()
|
||||
router.ServeHTTP(recorder, httptest.NewRequest(http.MethodPost, "/v1/responses", nil))
|
||||
|
||||
require.Equal(t, int64(1), OpsErrorLogQueueLength())
|
||||
job := <-opsErrorLogQueue
|
||||
require.Equal(t, "permission_error", job.entry.ErrorType)
|
||||
require.Equal(t, http.StatusForbidden, job.entry.StatusCode)
|
||||
}
|
||||
|
||||
func TestOpsErrorLoggerMiddleware_SkipsRecoveredUpstreamErrorOnSuccessfulRequest(t *testing.T) {
|
||||
setupOpsErrorLogTestQueue(t, 2)
|
||||
gin.SetMode(gin.TestMode)
|
||||
@@ -374,18 +422,18 @@ func TestOpsCaptureWriter_CapturesSplitDataOnlyTerminalMarkers(t *testing.T) {
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
writer := &opsCaptureWriter{limit: opsCaptureWriterLimit}
|
||||
writer.captureResponseChunk([]byte(tt.prefix), http.StatusOK)
|
||||
require.Empty(t, writer.buf.Bytes(), "partial marker must remain in the bounded probe")
|
||||
writer.captureResponseChunk([]byte(tt.suffix+"\n\n"), http.StatusOK)
|
||||
state := &opsCaptureWriterState{limit: opsCaptureWriterLimit}
|
||||
state.captureResponseChunk([]byte(tt.prefix), http.StatusOK)
|
||||
require.Empty(t, state.buf.Bytes(), "partial frame must remain in the bounded probe")
|
||||
state.captureResponseChunk([]byte(tt.suffix+"\n\n"), http.StatusOK)
|
||||
|
||||
parsed := parseOpsErrorResponse(writer.buf.Bytes())
|
||||
require.True(t, writer.sseCapturing)
|
||||
parsed := parseOpsErrorResponse(state.buf.Bytes())
|
||||
require.True(t, state.sseCapturing)
|
||||
require.True(t, parsed.StreamFailure)
|
||||
require.Equal(t, tt.wantType, parsed.ErrorType)
|
||||
require.Equal(t, tt.wantCode, parsed.Code)
|
||||
require.Equal(t, tt.wantError, parsed.Message)
|
||||
require.LessOrEqual(t, len(writer.probe), opsTerminalSSEProbeSize)
|
||||
require.LessOrEqual(t, len(state.probe), opsTerminalSSEFrameProbeLimit)
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -545,6 +593,23 @@ func TestLogOpsStreamError_SkipWhenPassthroughSkipMonitoring(t *testing.T) {
|
||||
require.Equal(t, int64(0), OpsErrorLogEnqueuedTotal())
|
||||
}
|
||||
|
||||
func TestShouldSkipFinalOpsFailureUsesOnlyFinalAttemptRule(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
c, _ := gin.CreateTestContext(httptest.NewRecorder())
|
||||
c.Set(service.OpsUpstreamErrorsKey, []*service.OpsUpstreamErrorEvent{
|
||||
{UpstreamStatusCode: http.StatusBadGateway, Message: "hidden intermediate", SkipMonitoring: true},
|
||||
{UpstreamStatusCode: http.StatusServiceUnavailable, Message: "visible final"},
|
||||
})
|
||||
require.False(t, shouldSkipFinalOpsFailure(c))
|
||||
|
||||
c.Set(service.OpsUpstreamErrorsKey, []*service.OpsUpstreamErrorEvent{
|
||||
{UpstreamStatusCode: http.StatusBadGateway, Message: "visible intermediate"},
|
||||
nil,
|
||||
{UpstreamStatusCode: http.StatusServiceUnavailable, Message: "hidden final", SkipMonitoring: true},
|
||||
})
|
||||
require.True(t, shouldSkipFinalOpsFailure(c))
|
||||
}
|
||||
|
||||
// MarkOpsStreamError 采用「首个标记生效」:后续的通用兜底帧不得覆盖根因错误。
|
||||
func TestMarkOpsStreamError_FirstWins(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
@@ -561,6 +626,35 @@ func TestMarkOpsStreamError_FirstWins(t *testing.T) {
|
||||
require.Equal(t, http.StatusTooManyRequests, se.IntendedStatus)
|
||||
}
|
||||
|
||||
func TestLogOpsStreamError_RecordsOneFailurePerWebSocketTurn(t *testing.T) {
|
||||
setupOpsErrorLogTestQueue(t, 4)
|
||||
gin.SetMode(gin.TestMode)
|
||||
c, _ := gin.CreateTestContext(httptest.NewRecorder())
|
||||
c.Request = httptest.NewRequest(http.MethodGet, "/v1/responses", nil)
|
||||
service.SetOpenAIClientTransport(c, service.OpenAIClientTransportWS)
|
||||
|
||||
service.BeginOpsStreamTurn(c, 1)
|
||||
service.MarkOpsStreamFailure(c, "rate_limit_error", "rate_limit_exceeded", "turn one failed", http.StatusTooManyRequests)
|
||||
service.MarkOpsStreamError(c, "upstream_error", "generic duplicate for turn one", http.StatusBadGateway)
|
||||
service.BeginOpsStreamTurn(c, 2)
|
||||
service.MarkOpsStreamFailure(c, "permission_error", "permission_denied", "turn two failed", http.StatusForbidden)
|
||||
|
||||
streamErrors := service.GetOpsStreamErrors(c)
|
||||
require.Len(t, streamErrors, 2)
|
||||
require.Equal(t, 1, streamErrors[0].Turn)
|
||||
require.Equal(t, 2, streamErrors[1].Turn)
|
||||
ops := service.NewOpsService(nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil)
|
||||
logOpsStreamError(c, ops, http.StatusSwitchingProtocols)
|
||||
|
||||
require.Equal(t, int64(2), OpsErrorLogQueueLength())
|
||||
first := <-opsErrorLogQueue
|
||||
second := <-opsErrorLogQueue
|
||||
require.Equal(t, "turn one failed", first.entry.ErrorMessage)
|
||||
require.Equal(t, http.StatusTooManyRequests, first.entry.StatusCode)
|
||||
require.Equal(t, "turn two failed", second.entry.ErrorMessage)
|
||||
require.Equal(t, http.StatusForbidden, second.entry.StatusCode)
|
||||
}
|
||||
|
||||
func TestIsKnownOpsErrorType(t *testing.T) {
|
||||
known := []string{
|
||||
"invalid_request_error",
|
||||
@@ -784,7 +878,11 @@ func TestClassifyOpsAuthClientErrorsExcludedFromSLA(t *testing.T) {
|
||||
errType := normalizeOpsErrorType(tt.errType, tt.code)
|
||||
phase, isBusinessLimited, errorOwner, errorSource := classifyOpsErrorLog(c, errType, tt.message, tt.code, tt.status)
|
||||
|
||||
require.Equal(t, "api_error", errType)
|
||||
wantErrType := "api_error"
|
||||
if tt.errType == "permission_error" {
|
||||
wantErrType = "permission_error"
|
||||
}
|
||||
require.Equal(t, wantErrType, errType)
|
||||
require.Equal(t, "auth", phase)
|
||||
require.True(t, isBusinessLimited)
|
||||
require.Equal(t, "client", errorOwner)
|
||||
@@ -935,7 +1033,7 @@ func TestClassifyOpsLocalBusinessLimitErrorsExcludedFromSLA(t *testing.T) {
|
||||
message: "This group is restricted to Claude Code clients (/v1/messages only)",
|
||||
code: "",
|
||||
status: http.StatusForbidden,
|
||||
wantErrType: "api_error",
|
||||
wantErrType: "permission_error",
|
||||
wantPhase: "request",
|
||||
},
|
||||
{
|
||||
@@ -944,7 +1042,7 @@ func TestClassifyOpsLocalBusinessLimitErrorsExcludedFromSLA(t *testing.T) {
|
||||
message: "Image generation is not enabled for this group",
|
||||
code: "",
|
||||
status: http.StatusForbidden,
|
||||
wantErrType: "api_error",
|
||||
wantErrType: "permission_error",
|
||||
wantPhase: "request",
|
||||
},
|
||||
{
|
||||
@@ -971,7 +1069,7 @@ func TestClassifyOpsLocalBusinessLimitErrorsExcludedFromSLA(t *testing.T) {
|
||||
message: "model claude-3-5-sonnet not in whitelist",
|
||||
code: "",
|
||||
status: http.StatusForbidden,
|
||||
wantErrType: "api_error",
|
||||
wantErrType: "permission_error",
|
||||
wantPhase: "request",
|
||||
},
|
||||
{
|
||||
@@ -998,7 +1096,7 @@ func TestClassifyOpsLocalBusinessLimitErrorsExcludedFromSLA(t *testing.T) {
|
||||
message: "openai service_tier=priority is not allowed for model gpt-5.5",
|
||||
code: "",
|
||||
status: http.StatusForbidden,
|
||||
wantErrType: "api_error",
|
||||
wantErrType: "permission_error",
|
||||
wantPhase: "request",
|
||||
},
|
||||
{
|
||||
@@ -1309,6 +1407,254 @@ func TestParseOpsErrorResponsePreservesNestedStringCode(t *testing.T) {
|
||||
require.Equal(t, "API Key 所属分组已删除", parsed.Message)
|
||||
}
|
||||
|
||||
func TestParseOpsErrorResponsePreservesStructuredTopLevelSemantics(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
body string
|
||||
wantType string
|
||||
wantCode string
|
||||
wantMsg string
|
||||
}{
|
||||
{
|
||||
name: "model not found",
|
||||
body: `{"type":"model_not_found","code":404,"message":"model unavailable"}`,
|
||||
wantType: "model_not_found",
|
||||
wantCode: "404",
|
||||
wantMsg: "model unavailable",
|
||||
},
|
||||
{
|
||||
name: "string error",
|
||||
body: `{"type":"service_unavailable","code":"temporarily_unavailable","error":"capacity exhausted"}`,
|
||||
wantType: "service_unavailable",
|
||||
wantCode: "temporarily_unavailable",
|
||||
wantMsg: "capacity exhausted",
|
||||
},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
parsed := parseOpsErrorResponse([]byte(tt.body))
|
||||
require.Equal(t, tt.wantType, normalizeOpsErrorType(parsed.ErrorType, parsed.Code))
|
||||
require.Equal(t, tt.wantCode, parsed.Code)
|
||||
require.Equal(t, tt.wantMsg, parsed.Message)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestApplyOpsUpstreamFieldsUsesLastNonNilAttempt(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
c, _ := gin.CreateTestContext(httptest.NewRecorder())
|
||||
service.SetOpsUpstreamError(c, http.StatusUnauthorized, "stale context", "stale detail")
|
||||
c.Set(service.OpsUpstreamErrorsKey, []*service.OpsUpstreamErrorEvent{
|
||||
{UpstreamStatusCode: http.StatusTooManyRequests, Message: "first attempt", Detail: "first detail"},
|
||||
nil,
|
||||
{UpstreamStatusCode: http.StatusServiceUnavailable, Message: "final attempt", Detail: "final detail"},
|
||||
nil,
|
||||
})
|
||||
entry := &service.OpsInsertErrorLogInput{}
|
||||
|
||||
applyOpsUpstreamFieldsFromContext(c, entry)
|
||||
|
||||
require.NotNil(t, entry.UpstreamStatusCode)
|
||||
require.Equal(t, http.StatusServiceUnavailable, *entry.UpstreamStatusCode)
|
||||
require.NotNil(t, entry.UpstreamErrorMessage)
|
||||
require.Equal(t, "final attempt", *entry.UpstreamErrorMessage)
|
||||
require.NotNil(t, entry.UpstreamErrorDetail)
|
||||
require.Equal(t, "final detail", *entry.UpstreamErrorDetail)
|
||||
require.Len(t, entry.UpstreamErrors, 4)
|
||||
}
|
||||
|
||||
func TestApplyOpsUpstreamFieldsFinalStatuslessAttemptClearsStaleContext(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
c, _ := gin.CreateTestContext(httptest.NewRecorder())
|
||||
service.SetOpsUpstreamError(c, http.StatusBadGateway, "stale response", "stale body")
|
||||
c.Set(service.OpsUpstreamErrorsKey, []*service.OpsUpstreamErrorEvent{
|
||||
{UpstreamStatusCode: http.StatusBadGateway, Message: "first response"},
|
||||
{Kind: "request_error", Message: "final transport failure", Detail: "connection reset"},
|
||||
})
|
||||
entry := &service.OpsInsertErrorLogInput{}
|
||||
|
||||
applyOpsUpstreamFieldsFromContext(c, entry)
|
||||
|
||||
require.Nil(t, entry.UpstreamStatusCode)
|
||||
require.NotNil(t, entry.UpstreamErrorMessage)
|
||||
require.Equal(t, "final transport failure", *entry.UpstreamErrorMessage)
|
||||
require.NotNil(t, entry.UpstreamErrorDetail)
|
||||
require.Equal(t, "connection reset", *entry.UpstreamErrorDetail)
|
||||
}
|
||||
|
||||
func TestOpsCaptureWriter_ProtocolLevelTerminalFrameDetection(t *testing.T) {
|
||||
state := &opsCaptureWriterState{limit: opsCaptureWriterLimit}
|
||||
chunks := []string{
|
||||
"event : response.failed\r\n",
|
||||
"data: { \"response\" : { \"error\" : { \"message\" : \"busy\", \"code\" : \"service_unavailable\" } },",
|
||||
" \"type\" : \"response.failed\" }\r\n\r\n",
|
||||
}
|
||||
for _, chunk := range chunks {
|
||||
state.captureResponseChunk([]byte(chunk), http.StatusOK)
|
||||
}
|
||||
|
||||
parsed := parseOpsErrorResponse(state.buf.Bytes())
|
||||
require.True(t, state.sseCapturing)
|
||||
require.True(t, parsed.StreamFailure)
|
||||
require.Equal(t, "service_unavailable_error", parsed.ErrorType)
|
||||
require.Equal(t, "service_unavailable", parsed.Code)
|
||||
require.Equal(t, "busy", parsed.Message)
|
||||
require.Equal(t, http.StatusServiceUnavailable, inferStreamFailureStatus(nil, parsed))
|
||||
}
|
||||
|
||||
func TestParseOpsSSEFailure_TopLevelErrorsAndUnknownStatus(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
body string
|
||||
wantType string
|
||||
wantStatus int
|
||||
}{
|
||||
{
|
||||
name: "top-level permission",
|
||||
body: "event: error\ndata: {\"message\":\"denied\",\"code\":\"permission_denied\",\"type\":\"error\"}\n\n",
|
||||
wantType: "permission_error",
|
||||
wantStatus: http.StatusForbidden,
|
||||
},
|
||||
{
|
||||
name: "top-level unavailable",
|
||||
body: "data: {\"message\":\"busy\",\"type\":\"error\",\"code\":\"service_unavailable\"}\n\n",
|
||||
wantType: "service_unavailable_error",
|
||||
wantStatus: http.StatusServiceUnavailable,
|
||||
},
|
||||
{
|
||||
name: "unknown terminal",
|
||||
body: "event: response.failed\ndata: {\"type\":\"response.failed\",\"error\":{\"code\":\"new_provider_code\",\"message\":\"failed\"}}\n\n",
|
||||
wantType: "upstream_error",
|
||||
wantStatus: http.StatusBadGateway,
|
||||
},
|
||||
{
|
||||
name: "explicit terminal status",
|
||||
body: "event: error\ndata: {\"type\":\"error\",\"status_code\":429,\"code\":\"new_rate_code\",\"message\":\"slow down\"}\n\n",
|
||||
wantType: "api_error",
|
||||
wantStatus: http.StatusTooManyRequests,
|
||||
},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
parsed := parseOpsErrorResponse([]byte(tt.body))
|
||||
require.True(t, parsed.StreamFailure)
|
||||
require.Equal(t, tt.wantType, parsed.ErrorType)
|
||||
gin.SetMode(gin.TestMode)
|
||||
c, _ := gin.CreateTestContext(httptest.NewRecorder())
|
||||
service.SetOpsUpstreamError(c, http.StatusUnauthorized, "old attempt", "")
|
||||
require.Equal(t, tt.wantStatus, inferStreamFailureStatus(c, parsed), "terminal status must not inherit an earlier attempt")
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestOpsCaptureWriter_OversizedNonTerminalFrameRemainsBounded(t *testing.T) {
|
||||
state := &opsCaptureWriterState{limit: opsCaptureWriterLimit}
|
||||
state.captureResponseChunk([]byte("data: "+strings.Repeat("x", opsTerminalSSEFrameProbeLimit*2)+"\n\n"), http.StatusOK)
|
||||
require.Empty(t, state.buf.Bytes())
|
||||
require.LessOrEqual(t, cap(state.probe), opsTerminalSSEFrameProbeLimit)
|
||||
|
||||
state.captureResponseChunk([]byte("event: error\ndata: {\"type\":\"error\",\"code\":\"permission_denied\",\"message\":\"denied\"}\n\n"), http.StatusOK)
|
||||
require.True(t, state.sseCapturing)
|
||||
require.NotEmpty(t, state.buf.Bytes())
|
||||
}
|
||||
|
||||
func TestOpsCaptureWriter_TerminalMetadataSurvivesBodyCaptureTruncation(t *testing.T) {
|
||||
state := &opsCaptureWriterState{limit: opsCaptureWriterLimit}
|
||||
frame := "event: response.failed\ndata: {\"padding\":\"" + strings.Repeat("x", opsCaptureWriterLimit) + "\",\"type\":\"response.failed\",\"error\":{\"code\":\"service_unavailable\",\"message\":\"busy\"}}\n\n"
|
||||
state.captureResponseChunk([]byte(frame), http.StatusOK)
|
||||
state.finalizeResponseCapture()
|
||||
|
||||
require.Len(t, state.buf.Bytes(), opsCaptureWriterLimit)
|
||||
require.True(t, parseOpsErrorResponse(state.buf.Bytes()).StreamFailure, "the bounded parser must fail closed from the terminal event line")
|
||||
require.True(t, state.terminalFound)
|
||||
require.Equal(t, "service_unavailable_error", state.terminalError.ErrorType)
|
||||
require.Equal(t, "busy", state.terminalError.Message)
|
||||
}
|
||||
|
||||
func TestOpsErrorLoggerMiddleware_LargeTerminalFrameUsesEventFallback(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.Status(http.StatusOK)
|
||||
_, _ = c.Writer.WriteString("event: response.failed\n")
|
||||
_, _ = c.Writer.WriteString("data: {\"authorization\":\"Bearer must-not-persist\",\"padding\":\"" + strings.Repeat("x", opsTerminalSSEFrameProbeLimit*2) + "\"}")
|
||||
})
|
||||
|
||||
recorder := httptest.NewRecorder()
|
||||
router.ServeHTTP(recorder, httptest.NewRequest(http.MethodPost, "/v1/responses", nil))
|
||||
require.Equal(t, http.StatusOK, recorder.Code)
|
||||
require.Equal(t, int64(1), OpsErrorLogQueueLength())
|
||||
job := <-opsErrorLogQueue
|
||||
require.Equal(t, http.StatusBadGateway, job.entry.StatusCode)
|
||||
require.Equal(t, "upstream_error", job.entry.ErrorType)
|
||||
require.Equal(t, "upstream stream failed", job.entry.ErrorMessage)
|
||||
require.NotContains(t, job.entry.ErrorMessage, "must-not-persist")
|
||||
require.NotContains(t, job.entry.ErrorBody, "must-not-persist")
|
||||
require.Contains(t, job.entry.ErrorBody, `"payload_truncated":true`)
|
||||
}
|
||||
|
||||
func TestOpsErrorLoggerMiddleware_DetectsTerminalDataAtEOFWithoutBlankLine(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.Status(http.StatusOK)
|
||||
_, _ = c.Writer.WriteString(`data: {"message":"denied","code":"permission_denied","type":"error"}`)
|
||||
})
|
||||
|
||||
recorder := httptest.NewRecorder()
|
||||
router.ServeHTTP(recorder, httptest.NewRequest(http.MethodPost, "/v1/responses", nil))
|
||||
require.Equal(t, http.StatusOK, recorder.Code)
|
||||
require.Equal(t, int64(1), OpsErrorLogQueueLength())
|
||||
job := <-opsErrorLogQueue
|
||||
require.Equal(t, http.StatusForbidden, job.entry.StatusCode)
|
||||
require.Equal(t, "permission_error", job.entry.ErrorType)
|
||||
require.Equal(t, "denied", job.entry.ErrorMessage)
|
||||
}
|
||||
|
||||
func TestOpsCaptureWriter_DetectsCROnlySSEFrame(t *testing.T) {
|
||||
state := &opsCaptureWriterState{limit: opsCaptureWriterLimit}
|
||||
state.captureResponseChunk([]byte("data: {\"type\":\"error\",\"code\":\"service_unavailable\",\"message\":\"busy\"}\r\r"), http.StatusOK)
|
||||
|
||||
require.True(t, state.sseCapturing)
|
||||
parsed := parseOpsErrorResponse(state.buf.Bytes())
|
||||
require.True(t, parsed.StreamFailure)
|
||||
require.Equal(t, "service_unavailable_error", parsed.ErrorType)
|
||||
}
|
||||
|
||||
func TestSanitizeOpsSSEDataForPersistence_RedactsJSONFields(t *testing.T) {
|
||||
body := []byte("event: error\ndata: {\"type\":\"error\",\"authorization\":\"Bearer secret\",\ndata: \"nested\":{\"api_key\":\"sk-secret\"}}\n\n")
|
||||
sanitized := sanitizeOpsSSEDataForPersistence(body)
|
||||
require.NotContains(t, sanitized, "Bearer secret")
|
||||
require.NotContains(t, sanitized, "sk-secret")
|
||||
require.Contains(t, sanitized, `"authorization":"[REDACTED]"`)
|
||||
require.Contains(t, sanitized, `"api_key":"[REDACTED]"`)
|
||||
}
|
||||
|
||||
func TestSanitizeOpsSSEDataForPersistence_DropsTruncatedJSONFragment(t *testing.T) {
|
||||
body := []byte("event: error\ndata: {\"type\":\"error\",\"authorization\":\"Bearer leaked")
|
||||
sanitized := sanitizeOpsSSEDataForPersistence(body)
|
||||
require.NotContains(t, sanitized, "Bearer leaked")
|
||||
require.Contains(t, sanitized, `data: {"payload_truncated":true}`)
|
||||
}
|
||||
|
||||
func BenchmarkOpsCaptureWriterSuccessfulSSEFrames(b *testing.B) {
|
||||
frame := []byte("event: response.output_text.delta\ndata: {\"type\":\"response.output_text.delta\",\"delta\":\"ok\"}\n\n")
|
||||
state := &opsCaptureWriterState{limit: opsCaptureWriterLimit}
|
||||
b.ReportAllocs()
|
||||
for i := 0; i < b.N; i++ {
|
||||
state.captureResponseChunk(frame, http.StatusOK)
|
||||
}
|
||||
if state.buf.Len() != 0 {
|
||||
b.Fatal("successful frames must not be captured")
|
||||
}
|
||||
}
|
||||
|
||||
func TestSetOpsEndpointContext_SetsContextKeys(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
rec := httptest.NewRecorder()
|
||||
|
||||
@@ -15,11 +15,11 @@ func TestOpsCaptureWriterDoesNotCopyIngressRejectBody(t *testing.T) {
|
||||
context, _ := gin.CreateTestContext(httptest.NewRecorder())
|
||||
writer := acquireOpsCaptureWriter(context.Writer)
|
||||
defer releaseOpsCaptureWriter(writer)
|
||||
writer.ctx = context
|
||||
writer.setContext(context)
|
||||
context.Writer = writer
|
||||
middleware2.MarkIngressRejected(context, middleware2.IngressRejectInvalidAPIKey)
|
||||
context.Status(http.StatusUnauthorized)
|
||||
_, err := context.Writer.WriteString(`{"code":"INVALID_API_KEY","message":"Invalid API key"}`)
|
||||
require.NoError(t, err)
|
||||
require.Zero(t, writer.buf.Len())
|
||||
require.Empty(t, writer.capturedBytes())
|
||||
}
|
||||
|
||||
@@ -121,9 +121,11 @@ func (h *UserMsgQueueHelper) waitForLockWithPing(
|
||||
c.Header("X-Accel-Buffering", "no")
|
||||
*streamStarted = true
|
||||
}
|
||||
if _, err := fmt.Fprint(c.Writer, string(h.pingFormat)); err != nil {
|
||||
written, err := fmt.Fprint(c.Writer, string(h.pingFormat))
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
recordGatewayStreamHeartbeat(c, written)
|
||||
flusher.Flush()
|
||||
|
||||
case <-timer.C:
|
||||
@@ -226,9 +228,11 @@ func (h *UserMsgQueueHelper) ThrottleWithPing(
|
||||
c.Header("X-Accel-Buffering", "no")
|
||||
*streamStarted = true
|
||||
}
|
||||
if _, err := fmt.Fprint(c.Writer, string(h.pingFormat)); err != nil {
|
||||
written, err := fmt.Fprint(c.Writer, string(h.pingFormat))
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
recordGatewayStreamHeartbeat(c, written)
|
||||
flusher.Flush()
|
||||
case <-timer.C:
|
||||
return nil
|
||||
|
||||
Reference in New Issue
Block a user