完善 OpenAI 网关切换与运维错误语义

This commit is contained in:
IanShaw
2026-08-19 23:13:40 -07:00
parent acce29af27
commit c374ff2951
23 changed files with 1854 additions and 349 deletions
+15 -2
View File
@@ -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)
}
+26 -1
View File
@@ -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:
+3 -3
View File
@@ -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"))
+13 -7
View File
@@ -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
+363 -17
View File
@@ -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