修复 PR 5888 审查发现的兼容性与竞态问题

This commit is contained in:
IanShaw
2026-08-20 21:06:48 -07:00
parent 16b15e870d
commit b2b2adcf8d
50 changed files with 1832 additions and 333 deletions
@@ -326,12 +326,15 @@ func TestOpsRecoveredCredentialFailoverDoesNotCreateRequestError(t *testing.T) {
router.ServeHTTP(recorder, httptest.NewRequest(http.MethodGet, "/openai/v1/responses", nil))
require.Equal(t, http.StatusOK, recorder.Code)
require.Zero(t, OpsErrorLogQueueLength())
select {
case job := <-opsErrorLogQueue:
t.Fatalf("successful failover must not create ops error row: %+v", job.entry)
default:
}
require.Equal(t, int64(1), OpsErrorLogQueueLength())
job := <-opsErrorLogQueue
require.Equal(t, http.StatusOK, job.entry.StatusCode)
require.Equal(t, string(service.GatewayFailureStageAccountAuth), job.entry.ErrorPhase)
require.NotNil(t, job.entry.UpstreamErrorsJSON)
events, err := service.ParseOpsUpstreamErrors(*job.entry.UpstreamErrorsJSON)
require.NoError(t, err)
require.Len(t, events, 2)
require.Equal(t, string(service.GatewayFailureStageAccountAuth), events[1].Stage)
}
func TestOpsWebSocketCredentialFailoverSuccessDoesNotCreateRequestError(t *testing.T) {
@@ -354,12 +357,15 @@ func TestOpsWebSocketCredentialFailoverSuccessDoesNotCreateRequestError(t *testi
router.ServeHTTP(recorder, request)
require.Equal(t, http.StatusOK, recorder.Code)
require.Zero(t, OpsErrorLogQueueLength())
select {
case job := <-opsErrorLogQueue:
t.Fatalf("successful websocket failover must not create ops error row: %+v", job.entry)
default:
}
require.Equal(t, int64(1), OpsErrorLogQueueLength())
job := <-opsErrorLogQueue
require.Equal(t, http.StatusOK, job.entry.StatusCode)
require.Equal(t, string(service.GatewayFailureStageAccountAuth), job.entry.ErrorPhase)
require.NotNil(t, job.entry.UpstreamErrorsJSON)
events, err := service.ParseOpsUpstreamErrors(*job.entry.UpstreamErrorsJSON)
require.NoError(t, err)
require.Len(t, events, 1)
require.Equal(t, string(service.GatewayFailureStageAccountAuth), events[0].Stage)
}
func TestOpsWebSocketCredentialFailoverExhaustedIsRecorded(t *testing.T) {
@@ -409,7 +409,27 @@ 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
}
groupID := int64(0)
if apiKey.GroupID != nil {
groupID = *apiKey.GroupID
}
owned, ownershipErr := h.gatewayService.ValidateOpenAIHTTPResponseOwner(
c.Request.Context(),
groupID,
previousResponseID,
subject.UserID,
apiKey.ID,
)
if ownershipErr != nil {
reqLog.Warn("openai.previous_response_owner_lookup_failed", zap.Error(ownershipErr))
}
if !owned {
reqLog.Warn("openai.request_validation_failed", zap.String("reason", "previous_response_owner_mismatch"))
h.errorResponse(c, http.StatusBadRequest, "invalid_request_error", "previous_response_id is not available for this user")
return
}
}
service.SetOpenAIHTTPResponseOwner(c, subject.UserID, apiKey.ID)
setOpsRequestContext(c, reqModel, reqStream)
setOpsEndpointContext(c, "", int16(service.RequestTypeFromLegacy(reqStream, false)))
@@ -877,12 +877,61 @@ func TestOpenAIResponses_AcceptsHTTPContinuationPreviousResponseIDBeforeRouting(
})
h := newOpenAIHandlerForPreviousResponseIDValidation(t, nil)
require.NoError(t, h.gatewayService.BindOpenAIHTTPResponseOwner(context.Background(), groupID, "resp_123456", 1, 101))
h.Responses(c)
require.NotEqual(t, http.StatusBadRequest, w.Code)
require.NotContains(t, w.Body.String(), "Responses WebSocket v2")
}
func TestOpenAIResponses_RejectsHTTPContinuationOwnedByAnotherUser(t *testing.T) {
gin.SetMode(gin.TestMode)
w := httptest.NewRecorder()
c, _ := gin.CreateTestContext(w)
c.Request = httptest.NewRequest(http.MethodPost, "/openai/v1/responses", strings.NewReader(
`{"model":"gpt-5.1","stream":false,"previous_response_id":"resp_other_tenant","input":"hello"}`,
))
c.Request.Header.Set("Content-Type", "application/json")
groupID := int64(2)
c.Set(string(middleware.ContextKeyAPIKey), &service.APIKey{
ID: 202,
UserID: 2,
GroupID: &groupID,
User: &service.User{ID: 2},
})
c.Set(string(middleware.ContextKeyUser), middleware.AuthSubject{UserID: 2, Concurrency: 1})
h := newOpenAIHandlerForPreviousResponseIDValidation(t, nil)
require.NoError(t, h.gatewayService.BindOpenAIHTTPResponseOwner(context.Background(), groupID, "resp_other_tenant", 1, 101))
h.Responses(c)
require.Equal(t, http.StatusBadRequest, w.Code)
require.Contains(t, w.Body.String(), "previous_response_id is not available for this user")
}
func TestOpenAIResponses_RejectsUnownedHTTPContinuation(t *testing.T) {
gin.SetMode(gin.TestMode)
w := httptest.NewRecorder()
c, _ := gin.CreateTestContext(w)
c.Request = httptest.NewRequest(http.MethodPost, "/openai/v1/responses", strings.NewReader(
`{"model":"gpt-5.1","stream":false,"previous_response_id":"resp_unknown","input":"hello"}`,
))
c.Request.Header.Set("Content-Type", "application/json")
groupID := int64(2)
c.Set(string(middleware.ContextKeyAPIKey), &service.APIKey{ID: 101, UserID: 1, GroupID: &groupID, User: &service.User{ID: 1}})
c.Set(string(middleware.ContextKeyUser), middleware.AuthSubject{UserID: 1, Concurrency: 1})
h := newOpenAIHandlerForPreviousResponseIDValidation(t, nil)
h.Responses(c)
require.Equal(t, http.StatusBadRequest, w.Code)
require.Contains(t, w.Body.String(), "previous_response_id is not available for this user")
}
func TestOpenAIResponses_FunctionCallOutputHTTPGuidanceDoesNotSuggestPreviousResponseReuse(t *testing.T) {
gin.SetMode(gin.TestMode)
@@ -12,6 +12,18 @@ import (
"github.com/stretchr/testify/require"
)
type blockingOpsResponseWriter struct {
gin.ResponseWriter
writeStarted chan struct{}
writeRelease chan struct{}
}
func (w *blockingOpsResponseWriter) WriteString(s string) (int, error) {
close(w.writeStarted)
<-w.writeRelease
return w.ResponseWriter.WriteString(s)
}
type deterministicOpsCaptureWriterStatePool struct {
states []*opsCaptureWriterState
}
@@ -144,3 +156,56 @@ func TestOpsCaptureWriter_StaleLeaseCannotReachReacquiredState(t *testing.T) {
defer releaseOpsCaptureWriter(other)
require.NotSame(t, current.state, other.state)
}
func TestOpsCaptureWriter_ReleaseWaitsForDelegatedWriteWithoutHoldingStateMutex(t *testing.T) {
gin.SetMode(gin.TestMode)
pool := &deterministicOpsCaptureWriterStatePool{}
recorder := httptest.NewRecorder()
ctx, _ := gin.CreateTestContext(recorder)
inner := &blockingOpsResponseWriter{
ResponseWriter: ctx.Writer,
writeStarted: make(chan struct{}),
writeRelease: make(chan struct{}),
}
w := acquireOpsCaptureWriterFromPool(pool, inner)
writeDone := make(chan struct{})
go func() {
defer close(writeDone)
_, _ = w.WriteString("body")
}()
<-inner.writeStarted
mutexAvailable := make(chan struct{})
go func() {
w.state.mu.Lock()
w.state.mu.Unlock()
close(mutexAvailable)
}()
select {
case <-mutexAvailable:
case <-time.After(time.Second):
t.Fatal("state mutex remained held across the delegated network write")
}
releaseDone := make(chan struct{})
go func() {
releaseOpsCaptureWriter(w)
close(releaseDone)
}()
select {
case <-releaseDone:
t.Fatal("release returned while a delegated write was still active")
case <-time.After(20 * time.Millisecond):
}
require.Empty(t, pool.states)
close(inner.writeRelease)
<-writeDone
select {
case <-releaseDone:
case <-time.After(time.Second):
t.Fatal("release did not finish after the delegated write returned")
}
require.Len(t, pool.states, 1)
}
+195 -16
View File
@@ -528,6 +528,7 @@ type opsCaptureWriter struct {
type opsCaptureWriterState struct {
mu sync.RWMutex
inFlight sync.WaitGroup
generation uint64
responseWriter gin.ResponseWriter
limit int
@@ -605,9 +606,15 @@ func releaseOpsCaptureWriter(w *opsCaptureWriter) {
state.mu.Unlock()
return
}
// Invalidate the lease before waiting. No new delegated calls can start for
// this handle, while calls that already copied the writer keep it alive via
// inFlight until their network operation returns.
state.generation++
state.responseWriter = nil
state.ctx = nil
state.mu.Unlock()
state.inFlight.Wait()
state.mu.Lock()
state.limit = opsCaptureWriterLimit
state.probe = state.probe[:0]
state.lineProbe = state.lineProbe[:0]
@@ -657,6 +664,27 @@ func (w *opsCaptureWriter) lockActiveWrite() (*opsCaptureWriterState, gin.Respon
return state, state.responseWriter
}
func (w *opsCaptureWriter) beginDelegatedCall() (*opsCaptureWriterState, gin.ResponseWriter) {
if w == nil || w.state == nil {
return nil, nil
}
state := w.state
state.mu.Lock()
if state.generation != w.generation || state.responseWriter == nil {
state.mu.Unlock()
return nil, nil
}
rw := state.responseWriter
state.inFlight.Add(1)
return state, rw
}
func finishDelegatedCall(state *opsCaptureWriterState) {
if state != nil {
state.inFlight.Done()
}
}
func (w *opsCaptureWriter) setContext(ctx *gin.Context) {
state, _ := w.lockActiveWrite()
if state == nil {
@@ -702,19 +730,21 @@ func (w *opsCaptureWriter) Header() http.Header {
return rw.Header()
}
func (w *opsCaptureWriter) WriteHeader(code int) {
state, rw := w.lockActive()
state, rw := w.beginDelegatedCall()
if state == nil {
return
}
defer state.mu.RUnlock()
state.mu.Unlock()
defer finishDelegatedCall(state)
rw.WriteHeader(code)
}
func (w *opsCaptureWriter) WriteHeaderNow() {
state, rw := w.lockActive()
state, rw := w.beginDelegatedCall()
if state == nil {
return
}
defer state.mu.RUnlock()
state.mu.Unlock()
defer finishDelegatedCall(state)
rw.WriteHeaderNow()
}
func (w *opsCaptureWriter) Status() int {
@@ -742,19 +772,21 @@ func (w *opsCaptureWriter) Written() bool {
return rw.Written()
}
func (w *opsCaptureWriter) Flush() {
state, rw := w.lockActive()
state, rw := w.beginDelegatedCall()
if state == nil {
return
}
defer state.mu.RUnlock()
state.mu.Unlock()
defer finishDelegatedCall(state)
rw.Flush()
}
func (w *opsCaptureWriter) Hijack() (net.Conn, *bufio.ReadWriter, error) {
state, rw := w.lockActive()
state, rw := w.beginDelegatedCall()
if state == nil {
return nil, nil, errors.New("response writer released")
}
defer state.mu.RUnlock()
state.mu.Unlock()
defer finishDelegatedCall(state)
return rw.Hijack()
}
func (w *opsCaptureWriter) CloseNotify() <-chan bool {
@@ -777,26 +809,28 @@ func (w *opsCaptureWriter) Pusher() http.Pusher {
}
func (w *opsCaptureWriter) Write(b []byte) (int, error) {
state, rw := w.lockActiveWrite()
state, rw := w.beginDelegatedCall()
if state == nil {
return 0, nil
}
defer state.mu.Unlock()
if state.shouldCapture() {
state.captureResponseChunk(b, rw.Status())
}
state.mu.Unlock()
defer finishDelegatedCall(state)
return rw.Write(b)
}
func (w *opsCaptureWriter) WriteString(s string) (int, error) {
state, rw := w.lockActiveWrite()
state, rw := w.beginDelegatedCall()
if state == nil {
return 0, nil
}
defer state.mu.Unlock()
if state.shouldCapture() {
state.captureResponseChunk([]byte(s), rw.Status())
}
state.mu.Unlock()
defer finishDelegatedCall(state)
return rw.WriteString(s)
}
@@ -887,6 +921,13 @@ func (state *opsCaptureWriterState) captureResponseChunk(chunk []byte, status in
state.appendCapturedResponse(chunk)
return
}
// Most stream writes contain one or more complete successful SSE frames.
// Skip the byte-wise frame parser when the chunk cannot contain a terminal
// event and leaves no split frame to carry into the next write.
if len(state.probe) == 0 && len(state.lineProbe) == 0 && endsAtOpsSSEFrameBoundary(chunk) &&
!mayContainOpsTerminalSSE(chunk) {
return
}
for i, b := range chunk {
if state.skipLF {
state.skipLF = false
@@ -952,6 +993,20 @@ func (state *opsCaptureWriterState) captureResponseChunk(chunk []byte, status in
}
}
func endsAtOpsSSEFrameBoundary(chunk []byte) bool {
return bytes.HasSuffix(chunk, []byte("\n\n")) ||
bytes.HasSuffix(chunk, []byte("\r\n\r\n")) ||
bytes.HasSuffix(chunk, []byte("\r\r"))
}
func mayContainOpsTerminalSSE(chunk []byte) bool {
if bytes.Contains(chunk, []byte("response.failed")) {
return true
}
return bytes.Contains(chunk, []byte("error")) &&
(bytes.Contains(chunk, []byte("event")) || bytes.Contains(chunk, []byte(`"type"`)))
}
func isOpsTerminalSSEEventLine(line []byte) bool {
line = bytes.TrimSpace(line)
field, value, found := bytes.Cut(line, []byte{':'})
@@ -1066,10 +1121,14 @@ func OpsErrorLoggerMiddleware(ops *service.OpsService) gin.HandlerFunc {
if parsed.StreamFailure {
status = inferStreamFailureStatus(c, parsed)
} else {
// Locally generated in-band errors use an explicit context marker and
// may not have a capturable terminal frame. Preserve that fallback,
// but never turn recovered upstream attempts into request errors.
logOpsStreamError(c, ops, status)
// A marked in-band error is a visible request failure even though its
// wire status is already 200. Otherwise retain recovered attempts as a
// provider-health row whose 2xx status keeps it outside request SLA.
if len(service.GetOpsStreamErrors(c)) > 0 {
logOpsStreamError(c, ops, status)
} else {
logOpsRecoveredUpstream(c, ops, status)
}
return
}
}
@@ -1214,6 +1273,126 @@ func OpsErrorLoggerMiddleware(ops *service.OpsService) gin.HandlerFunc {
}
}
func logOpsRecoveredUpstream(c *gin.Context, ops *service.OpsService, finalStatus int) {
if c == nil || ops == nil || finalStatus >= 400 {
return
}
entry := &service.OpsInsertErrorLogInput{StatusCode: finalStatus}
applyOpsUpstreamFieldsFromContext(c, entry)
if entry.UpstreamStatusCode == nil && entry.UpstreamErrorMessage == nil &&
entry.UpstreamErrorDetail == nil && len(entry.UpstreamErrors) == 0 {
return
}
lastStatus := 0
if entry.UpstreamStatusCode != nil {
lastStatus = *entry.UpstreamStatusCode
}
lastStage := ""
for i := len(entry.UpstreamErrors) - 1; i >= 0; i-- {
if event := entry.UpstreamErrors[i]; event != nil {
lastStage = event.Stage
if event.AccountID > 0 {
accountID := event.AccountID
entry.AccountID = &accountID
}
break
}
}
if entry.AccountID == nil {
if accountID, ok := c.Get(opsAccountIDKey); ok {
if value, ok := accountID.(int64); ok && value > 0 {
entry.AccountID = &value
}
}
}
entry.ErrorPhase = "upstream"
entry.ErrorType = "upstream_error"
entry.ErrorSource = "upstream_http"
entry.ErrorOwner = "provider"
entry.Severity = classifyOpsSeverity(entry.ErrorType, lastStatus)
entry.IsCountTokens = isCountTokensRequest(c)
entry.CreatedAt = time.Now()
entry.ErrorMessage = "Recovered upstream error"
if lastStage == string(service.GatewayFailureStageAccountAuth) {
entry.ErrorPhase = string(service.GatewayFailureStageAccountAuth)
entry.ErrorMessage = "Recovered account authentication failure"
} else if lastStatus > 0 {
entry.ErrorMessage += " " + strconv.Itoa(lastStatus)
}
if entry.UpstreamErrorMessage != nil && strings.TrimSpace(*entry.UpstreamErrorMessage) != "" {
entry.ErrorMessage += ": " + strings.TrimSpace(*entry.UpstreamErrorMessage)
}
entry.ErrorMessage = truncateString(entry.ErrorMessage, 2048)
if c.Request != nil {
entry.UserAgent = c.GetHeader("User-Agent")
if c.Request.URL != nil {
entry.RequestPath = c.Request.URL.Path
}
if c.Request.Context() != nil {
entry.ClientRequestID, _ = c.Request.Context().Value(ctxkey.ClientRequestID).(string)
entry.RequestID, _ = c.Request.Context().Value(ctxkey.RequestID).(string)
}
}
entry.RequestID = strings.TrimSpace(entry.RequestID)
if entry.RequestID == "" {
entry.RequestID = c.Writer.Header().Get("X-Request-Id")
}
entry.Model = c.GetString(opsModelKey)
entry.RequestedModel = entry.Model
entry.Stream = c.GetBool(opsStreamKey)
entry.InboundEndpoint = GetInboundEndpoint(c)
entry.UpstreamModel = c.GetString(opsUpstreamModelKey)
entry.RequestType = opsRequestTypeFromContext(c)
apiKey := getOpsAPIKey(c)
fallbackPlatform := guessPlatformFromPath(entry.RequestPath)
var requestContext context.Context = context.Background()
if c.Request != nil {
requestContext = c.Request.Context()
}
entry.Platform = resolveOpsPlatform(requestContext, apiKey, fallbackPlatform)
entry.UpstreamEndpoint = GetUpstreamEndpoint(c, entry.Platform)
if apiKey != nil {
entry.APIKeyID = &apiKey.ID
entry.APIKeyPrefix = keyPrefix(apiKey.Key, 8)
if apiKey.User != nil {
entry.UserID = &apiKey.User.ID
}
if apiKey.GroupID != nil {
entry.GroupID = apiKey.GroupID
}
if apiKey.Group != nil && apiKey.Group.Platform != "" {
entry.Platform = apiKey.Group.Platform
}
}
if clientIP := strings.TrimSpace(ip.GetClientIP(c)); clientIP != "" {
entry.ClientIP = &clientIP
}
applyOpsLatencyFieldsFromContext(c, entry)
enqueueOpsErrorLog(ops, entry)
}
func opsRequestTypeFromContext(c *gin.Context) *int16 {
if c == nil {
return nil
}
if value, ok := c.Get(opsRequestTypeKey); ok {
switch typed := value.(type) {
case int16:
result := typed
return &result
case int:
result := int16(typed)
return &result
}
}
return nil
}
// logOpsStreamError 记录一次挂在已固化 HTTP 200 SSE 流上的就地错误。
// 由于 wire 状态码停留在 200,常规的 status>=400 捕获路径永远不会触发;
// handleStreamingAwareError 通过 service.MarkOpsStreamError 标记这类错误,
@@ -38,15 +38,18 @@ func (r *ingressRejectSettingRepo) Set(context.Context, string, string) error {
type ingressRejectOpsRepo struct {
service.OpsRepository
insertCalls int
entries []*service.OpsInsertErrorLogInput
}
func (r *ingressRejectOpsRepo) InsertErrorLog(context.Context, *service.OpsInsertErrorLogInput) (int64, error) {
func (r *ingressRejectOpsRepo) InsertErrorLog(_ context.Context, entry *service.OpsInsertErrorLogInput) (int64, error) {
r.insertCalls++
r.entries = append(r.entries, entry)
return 0, nil
}
func (r *ingressRejectOpsRepo) BatchInsertErrorLogs(context.Context, []*service.OpsInsertErrorLogInput) (int64, error) {
func (r *ingressRejectOpsRepo) BatchInsertErrorLogs(_ context.Context, entries []*service.OpsInsertErrorLogInput) (int64, error) {
r.insertCalls++
r.entries = append(r.entries, entries...)
return 0, nil
}
@@ -328,11 +331,12 @@ func TestOpsErrorLoggerMiddleware_OrdinaryPermissionStillRecords(t *testing.T) {
require.Equal(t, http.StatusForbidden, job.entry.StatusCode)
}
func TestOpsErrorLoggerMiddleware_SkipsRecoveredUpstreamErrorOnSuccessfulRequest(t *testing.T) {
func TestOpsErrorLoggerMiddleware_RecordsRecoveredUpstreamTelemetryOutsideFailureSLA(t *testing.T) {
setupOpsErrorLogTestQueue(t, 2)
gin.SetMode(gin.TestMode)
ops := service.NewOpsService(nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil)
repo := &ingressRejectOpsRepo{}
ops := service.NewOpsService(repo, 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) {
@@ -347,7 +351,50 @@ func TestOpsErrorLoggerMiddleware_SkipsRecoveredUpstreamErrorOnSuccessfulRequest
router.ServeHTTP(recorder, httptest.NewRequest(http.MethodPost, "/v1/responses", nil))
require.Equal(t, http.StatusOK, recorder.Code)
require.Zero(t, OpsErrorLogQueueLength())
require.Equal(t, int64(1), OpsErrorLogQueueLength())
job := <-opsErrorLogQueue
require.Nil(t, job.entry.UpstreamErrors, "raw attempts must be released before async queueing")
require.NotNil(t, job.entry.UpstreamErrorsJSON)
queuedEvents, err := service.ParseOpsUpstreamErrors(*job.entry.UpstreamErrorsJSON)
require.NoError(t, err)
require.Len(t, queuedEvents, 1)
require.Equal(t, http.StatusTooManyRequests, queuedEvents[0].UpstreamStatusCode)
flushOpsErrorLogBatch([]opsErrorLogJob{job})
require.Equal(t, 1, repo.insertCalls)
require.Len(t, repo.entries, 1)
persisted := repo.entries[0]
require.Equal(t, http.StatusOK, persisted.StatusCode, "recovered telemetry must remain outside failed-request SLA")
require.Equal(t, "upstream", persisted.ErrorPhase)
require.Equal(t, "upstream_error", persisted.ErrorType)
require.Equal(t, "Recovered upstream error 429: earlier attempt was rate limited", persisted.ErrorMessage)
require.NotNil(t, persisted.UpstreamErrorsJSON)
persistedEvents, err := service.ParseOpsUpstreamErrors(*persisted.UpstreamErrorsJSON)
require.NoError(t, err)
require.Len(t, persistedEvents, 1)
require.Equal(t, http.StatusTooManyRequests, persistedEvents[0].UpstreamStatusCode)
}
func TestOpsErrorLoggerMiddleware_IntermediateSkipMonitoringDoesNotHideFinalVisibleFailure(t *testing.T) {
setupOpsErrorLogTestQueue(t, 2)
gin.SetMode(gin.TestMode)
ops := service.NewOpsService(nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil)
router := gin.New()
router.Use(OpsErrorLoggerMiddleware(ops))
router.POST("/v1/responses", func(c *gin.Context) {
c.Set(service.OpsUpstreamErrorsKey, []*service.OpsUpstreamErrorEvent{
{UpstreamStatusCode: http.StatusBadGateway, Message: "hidden retry", SkipMonitoring: true},
{UpstreamStatusCode: http.StatusServiceUnavailable, Message: "visible final"},
})
c.JSON(http.StatusServiceUnavailable, gin.H{"error": gin.H{"type": "upstream_error", "message": "visible final"}})
})
recorder := httptest.NewRecorder()
router.ServeHTTP(recorder, httptest.NewRequest(http.MethodPost, "/v1/responses", nil))
require.Equal(t, int64(1), OpsErrorLogQueueLength())
job := <-opsErrorLogQueue
require.Equal(t, http.StatusServiceUnavailable, job.entry.StatusCode)
require.Equal(t, "visible final", job.entry.ErrorMessage)
}
func TestOpsErrorLoggerMiddleware_CapturesSplitResponsesFailedSSE(t *testing.T) {
+49 -23
View File
@@ -270,27 +270,47 @@ func (c *gatewayCache) GetReasoningContent(ctx context.Context, itemID string) (
}
const (
cyberSessionBlockPrefix = "cyber_session_block:"
cyberSessionScopePrefix = "cyber_session_scope:"
cyberSessionBlockPrefix = "cyber_session_block:"
cyberSessionScopePrefix = "cyber_session_scope:"
cyberSessionRedisCommandMaxKeys = 128
)
// SetCyberSessionBlocked atomically writes all exact blocks and their optional
// coarse source scope with the same TTL.
// SetCyberSessionBlocked writes exact blocks in bounded transactions. The
// coarse scope is activated only after all exact blocks have been stored.
func (c *gatewayCache) SetCyberSessionBlocked(ctx context.Context, scopeKey string, keys []string, ttl time.Duration) error {
if len(keys) == 0 {
return nil
}
pipe := c.rdb.TxPipeline()
for _, key := range keys {
if key != "" {
exactKeys := make([]string, 0, cyberSessionRedisCommandMaxKeys)
flush := func() error {
if len(exactKeys) == 0 {
return nil
}
pipe := c.rdb.TxPipeline()
for _, key := range exactKeys {
pipe.Set(ctx, cyberSessionBlockPrefix+key, "1", ttl)
}
_, err := pipe.Exec(ctx)
exactKeys = exactKeys[:0]
return err
}
for _, key := range keys {
if key != "" {
exactKeys = append(exactKeys, key)
if len(exactKeys) == cyberSessionRedisCommandMaxKeys {
if err := flush(); err != nil {
return err
}
}
}
}
if err := flush(); err != nil {
return err
}
if scopeKey != "" {
pipe.Set(ctx, cyberSessionScopePrefix+scopeKey, "1", ttl)
return c.rdb.Set(ctx, cyberSessionScopePrefix+scopeKey, "1", ttl).Err()
}
_, err := pipe.Exec(ctx)
return err
return nil
}
func (c *gatewayCache) IsCyberSessionScopeActive(ctx context.Context, scopeKey string) (bool, error) {
@@ -301,23 +321,29 @@ func (c *gatewayCache) IsCyberSessionScopeActive(ctx context.Context, scopeKey s
return n > 0, nil
}
// FindCyberSessionBlocked checks transcript-prefix candidates in one Redis
// round trip and returns the first blocked key in caller order.
// FindCyberSessionBlocked checks bounded batches in caller order and stops at
// the first blocked key, preserving the original earliest-match behavior.
func (c *gatewayCache) FindCyberSessionBlocked(ctx context.Context, keys []string) (string, error) {
if len(keys) == 0 {
return "", nil
}
redisKeys := make([]string, len(keys))
for i, key := range keys {
redisKeys[i] = cyberSessionBlockPrefix + key
}
values, err := c.rdb.MGet(ctx, redisKeys...).Result()
if err != nil {
return "", err
}
for i, value := range values {
if value != nil {
return keys[i], nil
for start := 0; start < len(keys); start += cyberSessionRedisCommandMaxKeys {
end := start + cyberSessionRedisCommandMaxKeys
if end > len(keys) {
end = len(keys)
}
redisKeys := make([]string, end-start)
for i, key := range keys[start:end] {
redisKeys[i] = cyberSessionBlockPrefix + key
}
values, err := c.rdb.MGet(ctx, redisKeys...).Result()
if err != nil {
return "", err
}
for i, value := range values {
if value != nil {
return keys[start+i], nil
}
}
}
return "", nil
@@ -2,6 +2,8 @@ package repository
import (
"context"
"strconv"
"sync"
"testing"
"time"
@@ -11,6 +13,42 @@ import (
"github.com/stretchr/testify/require"
)
type cyberRedisCommandHook struct {
mu sync.Mutex
mgetKeyCounts []int
setBatchSizes []int
}
func (h *cyberRedisCommandHook) DialHook(next redis.DialHook) redis.DialHook { return next }
func (h *cyberRedisCommandHook) ProcessHook(next redis.ProcessHook) redis.ProcessHook {
return func(ctx context.Context, cmd redis.Cmder) error {
if cmd.Name() == "mget" {
h.mu.Lock()
h.mgetKeyCounts = append(h.mgetKeyCounts, len(cmd.Args())-1)
h.mu.Unlock()
}
return next(ctx, cmd)
}
}
func (h *cyberRedisCommandHook) ProcessPipelineHook(next redis.ProcessPipelineHook) redis.ProcessPipelineHook {
return func(ctx context.Context, cmds []redis.Cmder) error {
setCount := 0
for _, cmd := range cmds {
if cmd.Name() == "set" {
setCount++
}
}
if setCount > 0 {
h.mu.Lock()
h.setBatchSizes = append(h.setBatchSizes, setCount)
h.mu.Unlock()
}
return next(ctx, cmds)
}
}
func TestGatewayCacheCyberBlockWritesScopeAndExactKeysTogether(t *testing.T) {
server := miniredis.RunT(t)
client := redis.NewClient(&redis.Options{Addr: server.Addr()})
@@ -29,3 +67,31 @@ func TestGatewayCacheCyberBlockWritesScopeAndExactKeysTogether(t *testing.T) {
require.Greater(t, server.TTL(cyberSessionScopePrefix+"scope-1"), time.Duration(0))
require.Equal(t, server.TTL(cyberSessionBlockPrefix+"block-1"), server.TTL(cyberSessionBlockPrefix+"block-2"))
}
func TestGatewayCacheCyberBlockCommandsAreBoundedAndLookupShortCircuits(t *testing.T) {
server := miniredis.RunT(t)
client := redis.NewClient(&redis.Options{Addr: server.Addr()})
t.Cleanup(func() { _ = client.Close() })
hook := &cyberRedisCommandHook{}
client.AddHook(hook)
store, ok := NewGatewayCache(client).(service.CyberSessionBlockStore)
require.True(t, ok)
keys := make([]string, cyberSessionRedisCommandMaxKeys*2+44)
for i := range keys {
keys[i] = "block-" + strconv.Itoa(i)
}
ctx := context.Background()
require.NoError(t, store.SetCyberSessionBlocked(ctx, "large-scope", keys, time.Minute))
require.Equal(t, []int{cyberSessionRedisCommandMaxKeys, cyberSessionRedisCommandMaxKeys, 44}, hook.setBatchSizes)
lookup := make([]string, len(keys))
for i := range lookup {
lookup[i] = "missing-" + strconv.Itoa(i)
}
lookup[cyberSessionRedisCommandMaxKeys+3] = keys[cyberSessionRedisCommandMaxKeys+3]
matched, err := store.FindCyberSessionBlocked(ctx, lookup)
require.NoError(t, err)
require.Equal(t, keys[cyberSessionRedisCommandMaxKeys+3], matched)
require.Equal(t, []int{cyberSessionRedisCommandMaxKeys, cyberSessionRedisCommandMaxKeys}, hook.mgetKeyCounts)
}
@@ -141,8 +141,3 @@ func (c *tempUnschedCache) RecordOpenAIAPIKeyHealthFailure(ctx context.Context,
}
return count, tripped == 1, nil
}
func (c *tempUnschedCache) ResetOpenAIAPIKeyHealthFailures(ctx context.Context, accountID int64) error {
key := c.openAIAPIKeyHealthKey(accountID)
return c.rdb.Del(ctx, key, key+":sequence").Err()
}
@@ -3,6 +3,7 @@ package repository
import (
"context"
"testing"
"time"
"github.com/Wei-Shaw/sub2api/internal/service"
"github.com/alicebob/miniredis/v2"
@@ -10,7 +11,7 @@ import (
"github.com/stretchr/testify/require"
)
func TestOpenAIAPIKeyHealthCacheTripsAndSuccessResetsWindow(t *testing.T) {
func TestOpenAIAPIKeyHealthCacheTripsWithinRollingWindow(t *testing.T) {
server := miniredis.RunT(t)
client := redis.NewClient(&redis.Options{Addr: server.Addr()})
t.Cleanup(func() { _ = client.Close() })
@@ -18,14 +19,6 @@ func TestOpenAIAPIKeyHealthCacheTripsAndSuccessResetsWindow(t *testing.T) {
require.True(t, ok)
ctx := context.Background()
for attempt := 1; attempt <= 2; attempt++ {
count, tripped, err := store.RecordOpenAIAPIKeyHealthFailure(ctx, 42, 1, 3)
require.NoError(t, err)
require.EqualValues(t, attempt, count)
require.False(t, tripped)
}
require.NoError(t, store.ResetOpenAIAPIKeyHealthFailures(ctx, 42))
for attempt := 1; attempt <= 3; attempt++ {
count, tripped, err := store.RecordOpenAIAPIKeyHealthFailure(ctx, 42, 1, 3)
require.NoError(t, err)
@@ -33,3 +26,25 @@ func TestOpenAIAPIKeyHealthCacheTripsAndSuccessResetsWindow(t *testing.T) {
require.Equal(t, attempt == 3, tripped)
}
}
func TestOpenAIAPIKeyHealthCacheDropsFailuresOutsideRollingWindow(t *testing.T) {
server := miniredis.RunT(t)
now := time.Unix(1_700_000_000, 0)
server.SetTime(now)
client := redis.NewClient(&redis.Options{Addr: server.Addr()})
t.Cleanup(func() { _ = client.Close() })
store, ok := NewTempUnschedCache(client).(service.OpenAIAPIKeyHealthCache)
require.True(t, ok)
ctx := context.Background()
count, tripped, err := store.RecordOpenAIAPIKeyHealthFailure(ctx, 42, 1, 3)
require.NoError(t, err)
require.EqualValues(t, 1, count)
require.False(t, tripped)
server.SetTime(now.Add(61 * time.Second))
count, tripped, err = store.RecordOpenAIAPIKeyHealthFailure(ctx, 42, 1, 3)
require.NoError(t, err)
require.EqualValues(t, 1, count)
require.False(t, tripped)
}
@@ -24,22 +24,59 @@ func (r *openAIStream403AccountRepo) SetError(context.Context, int64, string) er
return nil
}
type openAIAuthPolicyAccountRepo struct {
AccountRepository
tempCalls int
setErrorCalls int
}
func (r *openAIAuthPolicyAccountRepo) SetTempUnschedulable(context.Context, int64, time.Time, string) error {
r.tempCalls++
return nil
}
func (r *openAIAuthPolicyAccountRepo) SetError(context.Context, int64, string) error {
r.setErrorCalls++
return nil
}
type openAIAuthPolicy403Counter struct {
counts []int64
}
func (s *openAIAuthPolicy403Counter) IncrementOpenAI403Count(context.Context, int64, int) (int64, error) {
if len(s.counts) == 0 {
return 1, nil
}
count := s.counts[0]
s.counts = s.counts[1:]
return count, nil
}
func (*openAIAuthPolicy403Counter) ResetOpenAI403Count(context.Context, int64) error {
return nil
}
func TestOpenAIUpstreamAccessStateClassification(t *testing.T) {
tests := []struct {
name string
body string
want bool
}{
{"workspace_code", `{"detail":{"code":"deactivated_workspace"}}`},
{"disabled_account", `{"error":{"message":"Your account is disabled"}}`},
{"suspended_workspace", `{"response":{"error":{"message":"This workspace has been suspended"}}}`},
{"deactivated_organization", `{"detail":{"message":"The organization is deactivated"}}`},
{"scalar_detail", `{"detail":"This workspace has been disabled"}`},
{"suspended_org_code", `{"error":{"code":"org_suspended"}}`},
{"workspace_code", `{"detail":{"code":"deactivated_workspace"}}`, true},
{"disabled_account_message", `{"error":{"message":"Your account is disabled"}}`, false},
{"suspended_workspace_message", `{"response":{"error":{"message":"This workspace has been suspended"}}}`, false},
{"deactivated_organization_message", `{"detail":{"message":"The organization is deactivated"}}`, false},
{"scalar_detail", `{"detail":"This workspace has been disabled"}`, false},
{"suspended_org_code", `{"error":{"code":"org_suspended"}}`, true},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
body := []byte(tt.body)
require.True(t, isOpenAIUpstreamAccessStateError("", body))
require.Equal(t, tt.want, isOpenAIUpstreamAccessStateError("", body))
if !tt.want {
return
}
require.True(t, (&OpenAIGatewayService{}).shouldFailoverOpenAIUpstreamResponse(http.StatusForbidden, "", body))
require.True(t, shouldFailoverOpenAIPassthroughResponse(&Account{Type: AccountTypeOAuth}, http.StatusForbidden, body))
@@ -62,6 +99,93 @@ func TestOpenAIUpstreamAccessStateDoesNotScanEchoedJSON(t *testing.T) {
require.False(t, shouldFailoverOpenAIPassthroughResponse(&Account{Type: AccountTypeOAuth}, http.StatusBadRequest, body))
}
func TestOpenAIHTTPAccessStateDoesNotTrustBadRequestMessage(t *testing.T) {
body := []byte(`{"error":{"type":"invalid_request_error","code":"unknown_parameter","message":"Unknown parameter: account disabled"}}`)
svc := &OpenAIGatewayService{}
require.False(t, isOpenAIUpstreamAccessStateError("", body), "free-form stream messages are not durable account evidence")
require.False(t, isOpenAIHTTPUpstreamAccessStateError(http.StatusBadRequest, "", body))
require.False(t, svc.shouldFailoverOpenAIUpstreamResponse(http.StatusBadRequest, "", body))
require.False(t, shouldFailoverOpenAIPassthroughResponse(&Account{Type: AccountTypeOAuth}, http.StatusBadRequest, body))
err := newOpenAIUpstreamFailoverError(http.StatusBadRequest, nil, body, "", false)
require.False(t, err.IsCredentialFailure())
}
func TestOpenAIHTTPAccessStateBadRequestDoesNotDisableAccount(t *testing.T) {
repo := &openAIStream403AccountRepo{}
svc := &OpenAIGatewayService{rateLimitService: &RateLimitService{accountRepo: repo}}
account := &Account{ID: 925, Platform: PlatformOpenAI, Type: AccountTypeOAuth, Status: StatusActive, Schedulable: true}
body := []byte(`{"error":{"code":"unknown_parameter","message":"Unknown parameter: account disabled"}}`)
disabled := svc.handleOpenAIAccountUpstreamError(context.Background(), account, http.StatusBadRequest, nil, body)
require.False(t, disabled)
require.Zero(t, repo.setErrorCalls)
require.False(t, svc.isOpenAIAccountRuntimeBlocked(account))
}
func TestOpenAIStreamEchoedAccessStateMessageDoesNotDisableOrFailover(t *testing.T) {
repo := &openAIStream403AccountRepo{}
svc := &OpenAIGatewayService{rateLimitService: &RateLimitService{accountRepo: repo}}
account := &Account{ID: 926, Platform: PlatformOpenAI, Type: AccountTypeOAuth, Status: StatusActive, Schedulable: true}
payload := []byte(`{"type":"response.failed","response":{"error":{"type":"invalid_request_error","code":"unknown_parameter","message":"Unknown parameter: account disabled"}}}`)
message := extractOpenAISSEErrorMessage(payload)
require.False(t, isOpenAIUpstreamAccessStateError(message, payload))
require.False(t, openAIStreamFailedEventShouldFailover(payload, message))
status, disabled := svc.handleOpenAIStreamTerminalAccountSideEffects(nil, account, payload, message, nil)
require.Equal(t, http.StatusBadGateway, status)
require.False(t, disabled)
require.Zero(t, repo.setErrorCalls)
require.False(t, svc.isOpenAIAccountRuntimeBlocked(account))
}
func TestOpenAIHTTPAccessStateTrustsStructuredCode(t *testing.T) {
repo := &openAIStream403AccountRepo{}
svc := &OpenAIGatewayService{rateLimitService: &RateLimitService{accountRepo: repo}}
account := &Account{ID: 930, Platform: PlatformOpenAI, Type: AccountTypeOAuth, Status: StatusActive, Schedulable: true}
body := []byte(`{"error":{"code":"organization_deactivated","message":"request rejected"}}`)
require.True(t, isOpenAIHTTPUpstreamAccessStateError(http.StatusBadRequest, "", body))
require.True(t, svc.shouldFailoverOpenAIUpstreamResponse(http.StatusBadRequest, "", body))
require.True(t, svc.handleOpenAIAccountUpstreamError(context.Background(), account, http.StatusBadRequest, nil, body))
require.Equal(t, 1, repo.setErrorCalls)
require.True(t, svc.isOpenAIAccountRuntimeBlocked(account))
}
func TestOpenAIHTTPAuthMessagesUseExistingStatusPolicies(t *testing.T) {
t.Run("oauth 401 remains recoverable", func(t *testing.T) {
repo := &openAIAuthPolicyAccountRepo{}
rateLimits := NewRateLimitService(repo, nil, &config.Config{}, nil, nil)
svc := &OpenAIGatewayService{rateLimitService: rateLimits}
account := &Account{ID: 931, Platform: PlatformOpenAI, Type: AccountTypeOAuth, Status: StatusActive, Schedulable: true,
Credentials: map[string]any{"refresh_token": "refreshable"}}
body := []byte(`{"error":{"message":"account is disabled"}}`)
require.False(t, isOpenAIHTTPUpstreamAccessStateError(http.StatusUnauthorized, "", body))
require.True(t, svc.handleOpenAIAccountUpstreamError(context.Background(), account, http.StatusUnauthorized, nil, body))
require.Zero(t, repo.setErrorCalls)
require.Equal(t, 1, repo.tempCalls)
})
t.Run("403 uses counter cooldown", func(t *testing.T) {
repo := &openAIAuthPolicyAccountRepo{}
counter := &openAIAuthPolicy403Counter{counts: []int64{1}}
rateLimits := NewRateLimitService(repo, nil, &config.Config{}, nil, nil)
rateLimits.openAI403CounterCache = counter
svc := &OpenAIGatewayService{rateLimitService: rateLimits}
rateLimits.SetAccountRuntimeBlocker(svc)
account := &Account{ID: 932, Platform: PlatformOpenAI, Type: AccountTypeOAuth, Status: StatusActive, Schedulable: true}
body := []byte(`{"error":{"message":"workspace has been suspended"}}`)
require.False(t, isOpenAIHTTPUpstreamAccessStateError(http.StatusForbidden, "", body))
require.True(t, svc.handleOpenAIAccountUpstreamError(context.Background(), account, http.StatusForbidden, nil, body))
require.Zero(t, repo.setErrorCalls)
require.Equal(t, 1, repo.tempCalls)
})
}
func TestOpenAICyberPolicyWrapped5xxNeverFailsOver(t *testing.T) {
body := []byte(`{"error":{"code":"cyber_policy","message":"blocked"}}`)
svc := &OpenAIGatewayService{}
@@ -65,7 +65,7 @@ func (s *OpenAIGatewayService) handleOpenAIAccountUpstreamError(ctx context.Cont
}
stateCtx, cancel := openAIAccountStateContext(ctx)
defer cancel()
if account != nil && account.Platform == PlatformOpenAI && isOpenAIUpstreamAccessStateError("", responseBody) {
if account != nil && account.Platform == PlatformOpenAI && isOpenAIHTTPUpstreamAccessStateError(statusCode, "", responseBody) {
message := "OpenAI upstream account or workspace is unavailable"
if upstreamMsg := strings.TrimSpace(extractUpstreamErrorMessage(responseBody)); upstreamMsg != "" {
message = upstreamMsg
@@ -15,11 +15,13 @@ import (
type oauth429RateLimitRepo struct {
AccountRepository
setRateLimitedCalls int
setRateLimitedCalls int
lastRateLimitedUntil time.Time
}
func (r *oauth429RateLimitRepo) SetRateLimited(context.Context, int64, time.Time) error {
func (r *oauth429RateLimitRepo) SetRateLimited(_ context.Context, _ int64, until time.Time) error {
r.setRateLimitedCalls++
r.lastRateLimitedUntil = until
return nil
}
@@ -60,6 +62,52 @@ func TestOpenAI429FastPath_BlocksOAuthOnlyAfterRetryWindow(t *testing.T) {
require.False(t, svc.shouldRetryOpenAIOAuth429OnSameAccount(account, http.StatusTooManyRequests, false))
}
func TestOpenAIStream429IgnoresSuccessfulQuotaSnapshotHeaders(t *testing.T) {
repo := &oauth429RateLimitRepo{}
rateLimits := NewRateLimitService(repo, nil, &config.Config{}, nil, nil)
svc := &OpenAIGatewayService{rateLimitService: rateLimits}
rateLimits.SetAccountRuntimeBlocker(svc)
account := &Account{ID: 421, Platform: PlatformOpenAI, Type: AccountTypeOAuth}
svc.openaiOAuth429RetryStartedAt.Store(account.ID, time.Now().Add(-openAIOAuth429RetryWindow-time.Second))
headers := http.Header{}
headers.Set("x-codex-primary-used-percent", "37")
headers.Set("x-codex-primary-reset-after-seconds", "604800")
headers.Set("x-codex-primary-window-minutes", "10080")
payload := []byte(`{"type":"error","error":{"type":"rate_limit_error","code":"rate_limit_exceeded","message":"slow down"}}`)
status, disabled := svc.handleOpenAIStreamTerminalAccountSideEffects(nil, account, payload, "slow down", headers)
require.Equal(t, http.StatusTooManyRequests, status)
require.False(t, disabled)
require.True(t, svc.isOpenAIAccountRuntimeBlocked(account))
value, ok := svc.openaiAccountRuntimeBlockUntil.Load(account.ID)
require.True(t, ok)
blockedUntil, ok := value.(time.Time)
require.True(t, ok)
require.Less(t, time.Until(blockedUntil), time.Minute, "stream 429 must not inherit the normal seven-day quota snapshot")
if !repo.lastRateLimitedUntil.IsZero() {
require.Less(t, time.Until(repo.lastRateLimitedUntil), time.Minute)
}
}
func TestOpenAIHTTP429StillUsesQuotaResetHeaders(t *testing.T) {
svc := &OpenAIGatewayService{rateLimitService: &RateLimitService{}}
account := &Account{ID: 422, Platform: PlatformOpenAI, Type: AccountTypeOAuth}
svc.openaiOAuth429RetryStartedAt.Store(account.ID, time.Now().Add(-openAIOAuth429RetryWindow-time.Second))
headers := http.Header{}
headers.Set("x-codex-primary-used-percent", "37")
headers.Set("x-codex-primary-reset-after-seconds", "604800")
headers.Set("x-codex-primary-window-minutes", "10080")
svc.markOpenAIOAuth429RateLimited(context.Background(), account, headers, nil)
value, ok := svc.openaiAccountRuntimeBlockUntil.Load(account.ID)
require.True(t, ok)
blockedUntil, ok := value.(time.Time)
require.True(t, ok)
require.Greater(t, time.Until(blockedUntil), 6*24*time.Hour, "real HTTP 429 must retain the upstream quota reset")
}
func TestOpenAI429RetryDelayHonorsBoundedRetryAfter(t *testing.T) {
deadline := time.Now().Add(openAIOAuth429RetryWindow)
require.Equal(t, openAIOAuth429RetryDelay, openAIOAuth429SameAccountRetryDelay(nil, deadline))
@@ -104,7 +104,7 @@ func (s *OpenAIGatewayService) ForwardAlphaSearch(ctx context.Context, c *gin.Co
if account.IsOpenAIOAuthLike() && resp.StatusCode == http.StatusTooManyRequests {
return nil, s.newOpenAIAccountFailoverError(account, resp.StatusCode, resp.Header, respBody, upstreamMessage, shouldDisable, retryableOnSameAccount)
}
if isOpenAIUpstreamAccessStateError(upstreamMessage, respBody) {
if isOpenAIHTTPUpstreamAccessStateError(resp.StatusCode, upstreamMessage, respBody) {
return nil, newOpenAIUpstreamFailoverError(resp.StatusCode, resp.Header, respBody, upstreamMessage, retryableOnSameAccount)
}
return nil, &UpstreamFailoverError{StatusCode: resp.StatusCode, ResponseBody: respBody, RetryableOnSameAccount: retryableOnSameAccount}
@@ -182,7 +182,7 @@ func (s *OpenAIGatewayService) forwardAlphaSearchViaResponsesWebSearch(
if account.IsOpenAIOAuthLike() && resp.StatusCode == http.StatusTooManyRequests {
return nil, s.newOpenAIAccountFailoverError(account, resp.StatusCode, resp.Header, respBody, upstreamMessage, shouldDisable, retryableOnSameAccount)
}
if isOpenAIUpstreamAccessStateError(upstreamMessage, respBody) {
if isOpenAIHTTPUpstreamAccessStateError(resp.StatusCode, upstreamMessage, respBody) {
return nil, newOpenAIUpstreamFailoverError(resp.StatusCode, resp.Header, respBody, upstreamMessage, retryableOnSameAccount)
}
return nil, &UpstreamFailoverError{StatusCode: resp.StatusCode, ResponseBody: respBody, RetryableOnSameAccount: retryableOnSameAccount}
@@ -124,15 +124,7 @@ func (s *RateLimitService) ObserveOpenAIAPIKeyHealthFailure(ctx context.Context,
return true
}
func (s *RateLimitService) ObserveOpenAIAPIKeyHealthSuccess(ctx context.Context, account *Account) {
if s == nil || s.openAIAPIKeyHealth == nil || s.settingService == nil || !isOpenAIAPIKeyHealthBreakerAccount(account) {
return
}
settings, err := s.settingService.GetOpenAIAPIKeyHealthBreakerSettings(ctx)
if err != nil || settings == nil || !settings.Enabled {
return
}
if err := s.openAIAPIKeyHealth.ResetOpenAIAPIKeyHealthFailures(ctx, account.ID); err != nil {
logger.L().Warn("openai.apikey_health_breaker_reset_failed", zap.Int64("account_id", account.ID), zap.Error(err))
}
func (s *RateLimitService) ObserveOpenAIAPIKeyHealthSuccess(context.Context, *Account) {
// Health failures are accumulated in a rolling time window. A success does
// not reset that window and must not add a Redis round trip to the hot path.
}
@@ -13,10 +13,12 @@ import (
type openAIAPIKeyHealthSettingRepo struct {
SettingRepository
value string
value string
getCalls int
}
func (r *openAIAPIKeyHealthSettingRepo) GetValue(context.Context, string) (string, error) {
r.getCalls++
return r.value, nil
}
@@ -35,7 +37,6 @@ func (r *openAIAPIKeyHealthAccountRepo) SetTempUnschedulable(_ context.Context,
type openAIAPIKeyHealthCacheStub struct {
TempUnschedCache
recordCalls int
resetCalls int
setCalls int
tripped bool
}
@@ -45,11 +46,6 @@ func (c *openAIAPIKeyHealthCacheStub) RecordOpenAIAPIKeyHealthFailure(context.Co
return 3, c.tripped, nil
}
func (c *openAIAPIKeyHealthCacheStub) ResetOpenAIAPIKeyHealthFailures(context.Context, int64) error {
c.resetCalls++
return nil
}
func (c *openAIAPIKeyHealthCacheStub) SetTempUnsched(context.Context, int64, *TempUnschedState) error {
c.setCalls++
return nil
@@ -127,10 +123,11 @@ func TestOpenAIAPIKeyHealthBreakerTripsPersistedAndRuntimeState(t *testing.T) {
require.Contains(t, repo.reason, openAIAPIKeyHealthBreakerReason)
}
func TestOpenAIAPIKeyHealthSuccessResetsOnlyEligiblePoolAccount(t *testing.T) {
func TestOpenAIAPIKeyHealthSuccessDoesNotTouchSettingsOrCache(t *testing.T) {
encoded, err := json.Marshal(OpenAIAPIKeyHealthBreakerSettings{Enabled: true, WindowMinutes: 1, FailureThreshold: 3, CooldownMinutes: 5})
require.NoError(t, err)
settings := NewSettingService(&openAIAPIKeyHealthSettingRepo{value: string(encoded)}, &config.Config{})
settingRepo := &openAIAPIKeyHealthSettingRepo{value: string(encoded)}
settings := NewSettingService(settingRepo, &config.Config{})
cache := &openAIAPIKeyHealthCacheStub{}
svc := NewRateLimitService(&openAIAPIKeyHealthAccountRepo{}, nil, &config.Config{}, nil, cache)
svc.SetSettingService(settings)
@@ -138,5 +135,6 @@ func TestOpenAIAPIKeyHealthSuccessResetsOnlyEligiblePoolAccount(t *testing.T) {
svc.ObserveOpenAIAPIKeyHealthSuccess(context.Background(), openAIHealthPoolAccount())
svc.ObserveOpenAIAPIKeyHealthSuccess(context.Background(), &Account{ID: 43, Platform: PlatformOpenAI, Type: AccountTypeAPIKey})
require.Equal(t, 1, cache.resetCalls)
require.Zero(t, settingRepo.getCalls)
require.Zero(t, cache.recordCalls)
}
@@ -151,6 +151,28 @@ func TestFilterCodexInput_ExistingItemIDWinsOverLegacyCallIDMapping(t *testing.T
require.Equal(t, "call_shared", filtered[2].(map[string]any)["id"])
}
func TestFilterCodexInput_NormalizesCrossTurnLegacyCallReference(t *testing.T) {
input := []any{
map[string]any{"type": "item_reference", "id": "call_previous_turn"},
}
filtered := filterCodexInputWithOptions(input, codexInputFilterOptions{PreserveReferences: true})
require.Equal(t, "fc_previous_turn", filtered[0].(map[string]any)["id"])
}
func TestFilterCodexInput_PreservesNativeRemoteItemReferences(t *testing.T) {
for _, id := range []string{"fc_remote", "ctc_remote", "tsc_remote", "msg_remote", "rs_remote", "vendor_remote"} {
t.Run(id, func(t *testing.T) {
input := []any{map[string]any{"type": "item_reference", "id": id}}
filtered := filterCodexInputWithOptions(input, codexInputFilterOptions{PreserveReferences: true})
require.Equal(t, id, filtered[0].(map[string]any)["id"])
})
}
}
// TestFilterCodexInput_StripsItemIDFromAllToolCallInputTypes verifies that
// item_* ids are stripped from all call-input types (not output types).
func TestFilterCodexInput_StripsItemIDFromAllToolCallInputTypes(t *testing.T) {
@@ -1562,10 +1562,25 @@ func codexInputItemIDs(input []any) map[string]struct{} {
return itemIDs
}
func codexInputCallIDs(input []any) map[string]struct{} {
callIDs := make(map[string]struct{})
for _, rawItem := range input {
item, ok := rawItem.(map[string]any)
if !ok || !isCodexToolCallItemType(strings.TrimSpace(firstNonEmptyString(item["type"]))) {
continue
}
if callID := strings.TrimSpace(firstNonEmptyString(item["call_id"])); callID != "" {
callIDs[callID] = struct{}{}
}
}
return callIDs
}
func filterCodexInputWithOptions(input []any, opts codexInputFilterOptions) []any {
filtered := make([]any, 0, len(input))
referenceIDMappings := codexItemReferenceIDMappings(input, opts.PreserveCallIDs)
inputItemIDs := codexInputItemIDs(input)
inputCallIDs := codexInputCallIDs(input)
for _, item := range input {
m, ok := item.(map[string]any)
if !ok {
@@ -1629,8 +1644,14 @@ func filterCodexInputWithOptions(input []any, opts codexInputFilterOptions) []an
if id, ok := newItem["id"].(string); ok && strings.HasPrefix(strings.TrimSpace(id), "call_") {
trimmedID := strings.TrimSpace(id)
_, referencesExistingItem := inputItemIDs[trimmedID]
if normalizedID, mapped := referenceIDMappings[trimmedID]; mapped && !referencesExistingItem {
newItem["id"] = normalizedID
if !referencesExistingItem {
if normalizedID, mapped := referenceIDMappings[trimmedID]; mapped {
newItem["id"] = normalizedID
} else if _, hasSameTurnCall := inputCallIDs[trimmedID]; !hasSameTurnCall {
// A bare call_* reference is a legacy function-call identifier.
// Normalize it even when its call item lives in an earlier turn.
newItem["id"] = normalizeCodexCallID(trimmedID)
}
}
}
filtered = append(filtered, newItem)
@@ -127,6 +127,26 @@ func TestApplyCodexOAuthTransform_ToolContinuationNormalizesToolReferenceIDsOnly
require.Equal(t, "fc_1", second["call_id"])
}
func TestApplyCodexOAuthTransform_NormalizesIsolatedLegacyReferenceAcrossTurns(t *testing.T) {
reqBody := map[string]any{
"model": "gpt-5.2",
"input": []any{
map[string]any{"type": "item_reference", "id": "call_previous_turn"},
map[string]any{"type": "item_reference", "id": "fc_remote_item"},
map[string]any{"type": "item_reference", "id": "vendor_remote_item"},
},
}
applyCodexOAuthTransform(reqBody, false, false)
input, ok := reqBody["input"].([]any)
require.True(t, ok)
require.Len(t, input, 3)
require.Equal(t, "fc_previous_turn", input[0].(map[string]any)["id"])
require.Equal(t, "fc_remote_item", input[1].(map[string]any)["id"])
require.Equal(t, "vendor_remote_item", input[2].(map[string]any)["id"])
}
func TestApplyCodexOAuthTransform_BoundsLongCallIDsAndPreservesPairing(t *testing.T) {
suffix := strings.Repeat("z", 62)
for _, tc := range []struct {
@@ -1,7 +1,10 @@
package service
import (
"bytes"
"encoding/json"
"errors"
"io"
"net/http"
"strings"
@@ -89,11 +92,7 @@ func isOpenAICompactModelFailure(statusCode int, upstreamMsg string, upstreamBod
case "model_not_found", "model_not_available", "unsupported_model", "invalid_model":
return true
}
if strings.Contains(value, "model") && (strings.Contains(value, "not found") ||
strings.Contains(value, "does not exist") ||
strings.Contains(value, "unavailable") ||
strings.Contains(value, "unsupported") ||
strings.Contains(value, "not supported")) {
if isExplicitOpenAIModelAvailabilityMessage(value) {
return true
}
}
@@ -115,6 +114,121 @@ func isOpenAICompactModelFailure(statusCode int, upstreamMsg string, upstreamBod
return false
}
func isExplicitOpenAIModelAvailabilityMessage(value string) bool {
value = strings.ToLower(strings.TrimSpace(value))
if value == "" {
return false
}
for _, phrase := range []string{
"model not found",
"model does not exist",
"model is unavailable",
"model is not available",
"model is unsupported",
"model is not supported",
"unsupported model",
} {
if strings.Contains(value, phrase) {
return true
}
}
// OpenAI commonly identifies the missing model between the word "model"
// and the terminal availability phrase, for example: "The model `x` does
// not exist". Requiring the message to start with the model subject avoids
// treating unrelated feature errors such as "model output is not supported"
// as a signal to change models.
if strings.HasPrefix(value, "the model ") || strings.HasPrefix(value, "model ") {
return strings.Contains(value, " does not exist") ||
strings.Contains(value, " was not found") ||
strings.Contains(value, " is unavailable") ||
strings.Contains(value, " is not available")
}
return false
}
func openAICompactFallbackErrorResponse(resp *http.Response, signal *openAICompactFallbackSignal) (*http.Response, []byte) {
headers := make(http.Header)
if resp != nil {
headers = resp.Header.Clone()
}
if headers.Get("Content-Type") == "" {
headers.Set("Content-Type", "application/json")
}
payload := normalizeOpenAICompactFallbackHTTPErrorPayload(signal)
return &http.Response{
StatusCode: http.StatusBadRequest,
Header: headers,
Body: io.NopCloser(bytes.NewReader(payload)),
}, payload
}
func normalizeOpenAICompactFallbackHTTPErrorPayload(signal *openAICompactFallbackSignal) []byte {
if signal == nil {
return nil
}
payload := append([]byte(nil), signal.payload...)
var terminal struct {
Error json.RawMessage `json:"error"`
Response struct {
Error json.RawMessage `json:"error"`
} `json:"response"`
}
if json.Unmarshal(payload, &terminal) != nil || len(bytes.TrimSpace(terminal.Response.Error)) == 0 ||
bytes.Equal(bytes.TrimSpace(terminal.Response.Error), []byte("null")) {
return payload
}
// Standard HTTP error handlers consume error.message/type/code. A streamed
// response.failed terminal nests the same object under response.error, so
// normalize only that envelope at the stream-to-HTTP boundary.
normalized, err := json.Marshal(struct {
Error json.RawMessage `json:"error"`
}{Error: terminal.Response.Error})
if err != nil {
return payload
}
return normalized
}
func (s *OpenAIGatewayService) appendOpenAICompactFallbackRetryOps(
c *gin.Context,
account *Account,
resp *http.Response,
payload []byte,
message string,
passthrough bool,
) {
if account == nil {
return
}
statusCode := http.StatusBadRequest
requestID := ""
if resp != nil {
statusCode = resp.StatusCode
requestID = resp.Header.Get("x-request-id")
}
detail := ""
if s != nil && s.cfg != nil && s.cfg.Gateway.LogUpstreamErrorBody {
maxBytes := s.cfg.Gateway.LogUpstreamErrorBodyMaxBytes
if maxBytes <= 0 {
maxBytes = 2048
}
detail = truncateString(string(payload), maxBytes)
}
appendOpsUpstreamError(c, OpsUpstreamErrorEvent{
Platform: account.Platform,
AccountID: account.ID,
AccountName: account.Name,
UpstreamStatusCode: statusCode,
UpstreamRequestID: requestID,
Passthrough: passthrough,
Kind: "retry",
Reason: "compact_model_fallback",
Message: sanitizeUpstreamErrorMessage(strings.TrimSpace(message)),
Detail: detail,
UpstreamResponseBody: detail,
})
}
// prepareOpenAICompactFallbackRetry returns a body for one safe, same-account
// retry. Callers invoke it only before any downstream response has been
// written; it changes the model and deliberately leaves path, trigger, and
@@ -164,6 +278,7 @@ func (s *OpenAIGatewayService) applyOpenAIPassthroughCompactFallbackFromSignal(
if !retry {
return body, "", false
}
s.appendOpenAICompactFallbackRetryOps(c, account, resp, signal.payload, signal.message, true)
if resp != nil && resp.Body != nil {
_ = resp.Body.Close()
}
@@ -3,11 +3,14 @@ package service
import (
"bytes"
"context"
"errors"
"io"
"net/http"
"net/http/httptest"
"strconv"
"strings"
"testing"
"time"
"github.com/Wei-Shaw/sub2api/internal/config"
"github.com/gin-gonic/gin"
@@ -138,6 +141,29 @@ func TestPrepareOpenAICompactFallbackRetryDoesNotHideSpecificBusinessFailure(t *
require.Equal(t, body, retryBody)
}
func TestIsOpenAICompactModelFailureRequiresExplicitModelAvailabilityMessage(t *testing.T) {
tests := []struct {
name string
message string
want bool
}{
{name: "explicit unsupported model", message: "The requested model is not supported", want: true},
{name: "named missing model", message: "The model `gpt-5.5` does not exist", want: true},
{name: "unsupported model code-like message", message: "unsupported model: gpt-5.5", want: true},
{name: "unsupported model feature", message: "This model output format is not supported", want: false},
{name: "unsupported parameter for model", message: "Parameter tools is not supported for this model", want: false},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
require.Equal(t, tt.want, isOpenAICompactModelFailure(
http.StatusBadRequest,
tt.message,
[]byte(`{"error":{"message":`+strconv.Quote(tt.message)+`}}`),
))
})
}
}
func TestPrepareOpenAICompactFallbackRetrySkipsSameModel(t *testing.T) {
gin.SetMode(gin.TestMode)
svc := &OpenAIGatewayService{cfg: &config.Config{Gateway: config.GatewayConfig{OpenAICompactModel: "gpt-5.5"}}}
@@ -191,6 +217,14 @@ func TestOpenAIGatewayForwardRetriesExplicitNativeCompactHTTPFailureOnce(t *test
require.True(t, HasCompactionTriggerInInput(upstream.bodies[1]))
require.Equal(t, upstream.requests[0].URL.Path, upstream.requests[1].URL.Path)
require.NotContains(t, upstream.requests[1].URL.Path, "/compact")
rawEvents, ok := c.Get(OpsUpstreamErrorsKey)
require.True(t, ok)
events, ok := rawEvents.([]*OpsUpstreamErrorEvent)
require.True(t, ok)
require.Len(t, events, 1)
require.Equal(t, "retry", events[0].Kind)
require.Equal(t, "compact_model_fallback", events[0].Reason)
require.Equal(t, http.StatusBadRequest, events[0].UpstreamStatusCode)
}
func TestOpenAIGatewayForwardRetriesExplicitNativeCompactSSEFailureBeforeOutput(t *testing.T) {
@@ -309,4 +343,63 @@ func TestOpenAIGatewayForwardDoesNotRecurseWhenCompactFallbackAlsoFails(t *testi
require.Len(t, upstream.bodies, 2)
require.Equal(t, "gpt-5.5", gjson.GetBytes(upstream.bodies[0], "model").String())
require.Equal(t, "gpt-5.4", gjson.GetBytes(upstream.bodies[1], "model").String())
var compactSignal *openAICompactFallbackSignal
require.False(t, errors.As(err, &compactSignal))
require.Equal(t, http.StatusBadRequest, recorder.Code)
require.Contains(t, recorder.Body.String(), "model not found")
rawEvents, ok := c.Get(OpsUpstreamErrorsKey)
require.True(t, ok)
events, ok := rawEvents.([]*OpsUpstreamErrorEvent)
require.True(t, ok)
require.Len(t, events, 2)
require.Equal(t, "retry", events[0].Kind)
require.Equal(t, "compact_model_fallback", events[0].Reason)
require.Equal(t, "http_error", events[1].Kind)
}
func TestOpenAIPassthroughCompactFallbackSecondStreamFailureUsesStandardErrorPath(t *testing.T) {
gin.SetMode(gin.TestMode)
body := []byte(`{"model":"gpt-5.5","stream":true,"input":[{"type":"compaction_trigger"}]}`)
recorder := httptest.NewRecorder()
c, _ := gin.CreateTestContext(recorder)
c.Request = httptest.NewRequest(http.MethodPost, "/v1/responses", nil)
c.Request.Header.Set("Content-Type", "application/json")
MarkOpenAINativeCompactionV2(c)
failed := "event: response.failed\n" +
`data: {"type":"response.failed","response":{"status":"failed","error":{"code":"context_length_exceeded","message":"context window exceeded"}}}` + "\n\n"
upstream := &httpUpstreamRecorder{responses: []*http.Response{
{StatusCode: http.StatusOK, Header: http.Header{"Content-Type": []string{"text/event-stream"}}, Body: io.NopCloser(strings.NewReader(failed))},
{StatusCode: http.StatusOK, Header: http.Header{"Content-Type": []string{"text/event-stream"}}, Body: io.NopCloser(strings.NewReader(failed))},
}}
svc := &OpenAIGatewayService{
cfg: &config.Config{Gateway: config.GatewayConfig{OpenAICompactModel: "gpt-5.4"}},
httpUpstream: upstream,
}
account := &Account{
ID: 1, Name: "openai-oauth", Platform: PlatformOpenAI, Type: AccountTypeOAuth, Concurrency: 1,
Credentials: map[string]any{"access_token": "oauth-token", "chatgpt_account_id": "chatgpt-account"},
Status: StatusActive, Schedulable: true,
}
result, err := svc.forwardOpenAIPassthrough(
context.Background(), c, account, body, body, "gpt-5.5", false, nil, true, time.Now(),
)
require.Error(t, err)
require.Nil(t, result)
require.Len(t, upstream.bodies, 2)
var compactSignal *openAICompactFallbackSignal
require.False(t, errors.As(err, &compactSignal))
require.Equal(t, http.StatusBadRequest, recorder.Code)
require.Contains(t, recorder.Body.String(), "context window exceeded")
rawEvents, ok := c.Get(OpsUpstreamErrorsKey)
require.True(t, ok)
events, ok := rawEvents.([]*OpsUpstreamErrorEvent)
require.True(t, ok)
require.Len(t, events, 2)
require.Equal(t, "retry", events[0].Kind)
require.Equal(t, "compact_model_fallback", events[0].Reason)
require.Equal(t, "http_error", events[1].Kind)
require.True(t, events[1].Passthrough)
}
@@ -21,6 +21,8 @@ type CyberSessionBlockStore interface {
FindCyberSessionBlocked(ctx context.Context, keys []string) (string, error)
}
const cyberSessionTranscriptLookupOverflowBlockKey = "transcript_lookup_limit_exceeded"
// CyberSessionExplicitBlockKey returns an inexpensive exact key when the
// client supplies a stable session signal.
func CyberSessionExplicitBlockKey(apiKeyID int64, c *gin.Context, body []byte) string {
@@ -142,7 +144,13 @@ func (s *OpenAIGatewayService) FindCyberSessionBlockedForRequest(ctx context.Con
if !active {
return ""
}
keys := CyberSessionTranscriptLookupKeys(apiKeyID, body)
transcript := deriveOpenAICyberTranscriptBlockKeys(apiKeyID, body)
if transcript.lookupKeysTruncated {
// Once the coarse scope is active, silently dropping old candidates would
// let a blocked client evade prefix matching by appending dummy items.
return cyberSessionTranscriptLookupOverflowBlockKey
}
keys := transcript.lookupKeys
if len(keys) == 0 {
return ""
}
@@ -2,8 +2,10 @@ package service
import (
"context"
"encoding/json"
"errors"
"net/http/httptest"
"strconv"
"strings"
"testing"
"time"
@@ -75,11 +77,32 @@ func TestCyberTranscriptBlockKeysWebSocketResponseCreate(t *testing.T) {
require.Len(t, CyberSessionTranscriptBlockKeys(88, body), 2)
}
func TestCyberTranscriptLookupKeysAreBoundedAndKeepNewestOrder(t *testing.T) {
messages := make([]map[string]string, maxOpenAICyberTranscriptLookupKeys+44)
for i := range messages {
messages[i] = map[string]string{"role": "user", "content": "message-" + strconv.Itoa(i)}
}
body, err := json.Marshal(map[string]any{"messages": messages})
require.NoError(t, err)
keys := CyberSessionTranscriptLookupKeys(77, body)
require.Len(t, keys, maxOpenAICyberTranscriptLookupKeys)
firstRetainedBody, err := json.Marshal(map[string]any{"messages": messages[:45]})
require.NoError(t, err)
firstRetainedPrefix := CyberSessionTranscriptLookupKeys(77, firstRetainedBody)
require.Equal(t, firstRetainedPrefix[len(firstRetainedPrefix)-1], keys[0])
fullKey := CyberSessionTranscriptBlockKeys(77, body)[0]
require.Equal(t, fullKey, keys[len(keys)-1])
}
// --- fakes ---
type fakeCyberBlockStore struct {
blocked map[string]bool
scopes map[string]bool
blocked map[string]bool
scopes map[string]bool
findCalls int
}
var _ CyberSessionBlockStore = (*fakeCyberBlockStore)(nil)
@@ -105,6 +128,7 @@ func (f *fakeCyberBlockStore) IsCyberSessionScopeActive(_ context.Context, scope
}
func (f *fakeCyberBlockStore) FindCyberSessionBlocked(_ context.Context, keys []string) (string, error) {
f.findCalls++
for _, key := range keys {
if f.blocked[key] {
return key, nil
@@ -271,6 +295,34 @@ func TestFindCyberSessionBlockedForRequestUsesScopeForTranscript(t *testing.T) {
require.Equal(t, blockKey, svc.FindCyberSessionBlockedForRequest(ctx, 9, nextCtx, nextBody, clientIP, "Codex CLI 1.2.4"))
}
func TestFindCyberSessionBlockedForRequestFailsClosedOnScopedTranscriptOverflow(t *testing.T) {
settingSvc := &SettingService{settingRepo: &fakeSettingRepo{vals: map[string]string{
SettingKeyCyberSessionBlockEnabled: "true",
SettingKeyCyberSessionBlockTTLSeconds: "60",
}}}
combo := &comboCacheAndStore{}
svc := &OpenAIGatewayService{cache: combo, settingService: settingSvc}
ctx := context.Background()
const apiKeyID = int64(9)
const clientIP = "203.0.113.20"
const userAgent = "Codex CLI 1.2.3"
messages := make([]map[string]string, maxOpenAICyberTranscriptLookupKeys+1)
for i := range messages {
messages[i] = map[string]string{"role": "user", "content": "message-" + strconv.Itoa(i)}
}
body, err := json.Marshal(map[string]any{"messages": messages})
require.NoError(t, err)
c, _ := newCyberBlockTestCtx(nil, string(body))
require.Empty(t, svc.FindCyberSessionBlockedForRequest(ctx, apiKeyID, c, body, clientIP, userAgent),
"overflow alone must not bypass the scope gate")
combo.store.scopes = map[string]bool{CyberSessionScopeKey(apiKeyID, clientIP, userAgent): true}
require.Equal(t, cyberSessionTranscriptLookupOverflowBlockKey,
svc.FindCyberSessionBlockedForRequest(ctx, apiKeyID, c, body, clientIP, userAgent))
require.Zero(t, combo.store.findCalls, "overflow must not issue an unbounded Redis lookup")
}
func TestCyberSessionScopeKeyNormalizesUserAgentVersion(t *testing.T) {
base := CyberSessionScopeKey(7, "203.0.113.10", "Codex CLI 1.2.3")
require.NotEmpty(t, base)
@@ -11,10 +11,15 @@ import (
)
type openAICyberTranscriptBlockKeys struct {
lookupKeys []string
preLatestUserKey string
lookupKeys []string
preLatestUserKey string
lookupKeysTruncated bool
}
// Bound the Redis lookup work for a single request while retaining the most
// recent transcript prefixes, where a continuation is most likely to match.
const maxOpenAICyberTranscriptLookupKeys = 256
// deriveOpenAICyberTranscriptBlockKeys returns cumulative semantic-history
// hashes plus the context key immediately before the latest user turn. The
// context key requires model-generated history so shared first-turn templates
@@ -53,8 +58,11 @@ func deriveOpenAICyberTranscriptBlockKeys(apiKeyID int64, body []byte) openAICyb
return openAICyberTranscriptBlockKeys{}
}
result := openAICyberTranscriptBlockKeys{
lookupKeys: make([]string, 0, int(sequence.Get("#").Int())),
lookupKeys: make([]string, 0, maxOpenAICyberTranscriptLookupKeys),
}
nextLookupKey := 0
lookupKeysRotated := false
lastLookupKey := ""
// This is an entropy heuristic, not provenance proof: authenticated
// server-side history would be required to distinguish fixed few-shot
// assistant items perfectly.
@@ -71,17 +79,31 @@ func deriveOpenAICyberTranscriptBlockKeys(apiKeyID int64, body []byte) openAICyb
if strings.TrimSpace(canonical) == "" {
return true
}
if openAICyberTranscriptItemStartsUserTurn(item) && hasModelGeneratedItem && len(result.lookupKeys) > 0 {
result.preLatestUserKey = result.lookupKeys[len(result.lookupKeys)-1]
if openAICyberTranscriptItemStartsUserTurn(item) && hasModelGeneratedItem && lastLookupKey != "" {
result.preLatestUserKey = lastLookupKey
}
_, _ = h.Write([]byte("|item="))
_, _ = h.Write([]byte(canonical))
result.lookupKeys = append(result.lookupKeys, hex.EncodeToString(h.Sum(nil)))
lastLookupKey = hex.EncodeToString(h.Sum(nil))
if len(result.lookupKeys) < maxOpenAICyberTranscriptLookupKeys {
result.lookupKeys = append(result.lookupKeys, lastLookupKey)
} else {
result.lookupKeys[nextLookupKey] = lastLookupKey
nextLookupKey = (nextLookupKey + 1) % maxOpenAICyberTranscriptLookupKeys
lookupKeysRotated = true
result.lookupKeysTruncated = true
}
if openAICyberTranscriptItemIsModelGenerated(item) {
hasModelGeneratedItem = true
}
return true
})
if lookupKeysRotated {
ordered := make([]string, 0, len(result.lookupKeys))
ordered = append(ordered, result.lookupKeys[nextLookupKey:]...)
ordered = append(ordered, result.lookupKeys[:nextLookupKey]...)
result.lookupKeys = ordered
}
return result
}
@@ -140,7 +140,7 @@ func (s *OpenAIGatewayService) ForwardEmbeddings(
if account.IsOpenAIOAuth() && resp.StatusCode == http.StatusTooManyRequests {
return nil, s.newOpenAIAccountFailoverError(account, resp.StatusCode, resp.Header, respBody, upstreamMsg, shouldDisable, retryableOnSameAccount)
}
if isOpenAIUpstreamAccessStateError(upstreamMsg, respBody) {
if isOpenAIHTTPUpstreamAccessStateError(resp.StatusCode, upstreamMsg, respBody) {
return nil, newOpenAIUpstreamFailoverError(resp.StatusCode, resp.Header, respBody, upstreamMsg, retryableOnSameAccount)
}
return nil, &UpstreamFailoverError{StatusCode: resp.StatusCode, ResponseBody: respBody, RetryableOnSameAccount: retryableOnSameAccount}
@@ -258,7 +258,7 @@ func TestSanitizeOpenAIResponsesInputItemIDs_AllocationGrowthIsLinear(t *testing
"10x more input items must not cause quadratic whole-body allocation growth")
}
func TestNormalizeOpenAIResponsesWebSocketCompatibilityBodyClosesInvalidIDReferences(t *testing.T) {
func TestNormalizeOpenAIResponsesWebSocketCompatibilityBodyPreservesOpaqueReferences(t *testing.T) {
body := []byte(`{"type":"response.create","input":[
{"type":"custom_tool_call","id":"ctc_call","call_id":"call_custom","name":"apply_patch","input":"patch"},
{"type":"custom_tool_call_output","id":"ctco_bad","call_id":"call_custom","output":"done"},
@@ -275,11 +275,12 @@ func TestNormalizeOpenAIResponsesWebSocketCompatibilityBodyClosesInvalidIDRefere
require.NoError(t, err)
require.True(t, changed)
require.Len(t, gjson.GetBytes(normalized, "input").Array(), 3)
require.Len(t, gjson.GetBytes(normalized, "input").Array(), 4)
require.Equal(t, "ctc_call", gjson.GetBytes(normalized, "input.0.id").String())
require.Equal(t, "call_custom", gjson.GetBytes(normalized, "input.1.call_id").String())
require.False(t, gjson.GetBytes(normalized, "input.1.id").Exists())
require.Equal(t, "item_future", gjson.GetBytes(normalized, "input.2.id").String())
require.Equal(t, "ctco_bad", gjson.GetBytes(normalized, "input.2.id").String())
require.Equal(t, "item_future", gjson.GetBytes(normalized, "input.3.id").String())
second, changedAgain, err := normalizeOpenAIResponsesWebSocketCompatibilityBody(normalized, &Account{
Platform: PlatformOpenAI,
@@ -59,21 +59,10 @@ func (s *OpenAIGatewayService) Forward(ctx context.Context, c *gin.Context, acco
// 在分流到 passthrough / Codex transform / 原生 ChatCompletions 之前统一修正
// 显式为 null 的工具 Schema type,否则 upstream 的 400 会被归一成可重试的 502,
// 同一份坏定义在账号池里反复重放。
if shouldSanitizeOpenAIResponsesToolSchemas(account.Platform) {
sanitizedToolBody, toolSchemaSanitized, toolSchemaErr := sanitizeOpenAIResponsesToolParameterTypes(body)
if toolSchemaErr != nil {
return nil, fmt.Errorf("sanitize OpenAI Responses tool parameters: %w", toolSchemaErr)
}
if toolSchemaSanitized {
body = sanitizedToolBody
}
patternSanitizedBody, patternSanitized, patternErr := sanitizeOpenAIResponsesToolSchemaPatterns(body)
if patternErr != nil {
return nil, fmt.Errorf("sanitize OpenAI Responses tool schema patterns: %w", patternErr)
}
if patternSanitized {
body = patternSanitizedBody
}
if sanitizedToolBody, toolSchemaSanitized, toolSchemaErr := sanitizeOpenAIResponsesToolSchemasForPlatform(body, account.Platform); toolSchemaErr != nil {
return nil, toolSchemaErr
} else if toolSchemaSanitized {
body = sanitizedToolBody
}
if account.IsOpenAIOAuthLike() {
reasoningBody, reasoningChanged, reasoningErr := normalizeOpenAIResponsesReasoningMode(body)
@@ -1015,6 +1004,7 @@ func (s *OpenAIGatewayService) Forward(ctx context.Context, c *gin.Context, acco
if retryBody, fallbackModel, retry := s.prepareOpenAICompactFallbackRetry(
c, account, requestedModel, body, resp.StatusCode, upstreamMsg, respBody, compactModelFallbackRetried,
); retry {
s.appendOpenAICompactFallbackRetryOps(c, account, resp, respBody, upstreamMsg, false)
fromModel := strings.TrimSpace(gjson.GetBytes(body, "model").String())
body = retryBody
requestView = newOpenAIRequestView(body)
@@ -1082,6 +1072,7 @@ func (s *OpenAIGatewayService) Forward(ctx context.Context, c *gin.Context, acco
if retryBody, fallbackModel, retry := s.prepareOpenAICompactFallbackRetry(
c, account, requestedModel, body, http.StatusBadRequest, signal.message, signal.payload, compactModelFallbackRetried,
); retry {
s.appendOpenAICompactFallbackRetryOps(c, account, resp, signal.payload, signal.message, false)
body = retryBody
requestView = newOpenAIRequestView(body)
upstreamModel = fallbackModel
@@ -1089,6 +1080,27 @@ func (s *OpenAIGatewayService) Forward(ctx context.Context, c *gin.Context, acco
SetOpsUpstreamModel(c, fallbackModel)
continue
}
if resp.Body != nil {
_ = resp.Body.Close()
}
compactResp, compactBody := openAICompactFallbackErrorResponse(resp, signal)
if s.shouldFailoverOpenAIUpstreamResponse(compactResp.StatusCode, signal.message, compactBody) {
appendOpsUpstreamError(c, OpsUpstreamErrorEvent{
Platform: account.Platform,
AccountID: account.ID,
AccountName: account.Name,
UpstreamStatusCode: compactResp.StatusCode,
UpstreamRequestID: compactResp.Header.Get("x-request-id"),
Kind: "failover",
Message: signal.message,
})
shouldDisable := s.handleFailoverSideEffects(ctx, compactResp, account, compactBody, upstreamModel)
return nil, s.newOpenAIAccountFailoverError(
account, compactResp.StatusCode, compactResp.Header, compactBody, signal.message, shouldDisable,
!shouldDisable && account.IsPoolMode() && (account.IsPoolModeRetryableStatus(compactResp.StatusCode) || isOpenAITransientProcessingError(compactResp.StatusCode, signal.message, compactBody)),
)
}
return s.handleErrorResponse(ctx, compactResp, c, account, body, resolveOpenAIErrorSchedulingModel(billingModel, upstreamModel))
}
return nil, err
}
@@ -377,6 +377,7 @@ func (s *OpenAIGatewayService) forwardOpenAIPassthrough(
if retryBody, fallbackModel, retry := s.prepareOpenAICompactFallbackRetry(
c, account, requestedModel, body, resp.StatusCode, upstreamMsg, probeBody, compactModelFallbackRetried,
); retry {
s.appendOpenAICompactFallbackRetryOps(c, account, resp, probeBody, upstreamMsg, true)
fromModel := strings.TrimSpace(gjson.GetBytes(body, "model").String())
body = retryBody
upstreamPassthroughModel = fallbackModel
@@ -425,6 +426,14 @@ func (s *OpenAIGatewayService) forwardOpenAIPassthrough(
compactModelFallbackRetried = true
continue
}
if signal, ok := asOpenAICompactFallbackSignal(handleErr); ok {
_ = resp.Body.Close()
compactResp, compactBody := openAICompactFallbackErrorResponse(resp, signal)
if shouldFailoverOpenAIPassthroughResponse(account, compactResp.StatusCode, compactBody) {
return nil, s.handleFailoverErrorResponsePassthrough(ctx, compactResp, c, account, body, compactBody)
}
return nil, s.handleErrorResponsePassthrough(ctx, compactResp, c, account, body, compactBody)
}
_ = resp.Body.Close()
return nil, handleErr
}
@@ -444,6 +453,14 @@ func (s *OpenAIGatewayService) forwardOpenAIPassthrough(
compactModelFallbackRetried = true
continue
}
if signal, ok := asOpenAICompactFallbackSignal(handleErr); ok {
_ = resp.Body.Close()
compactResp, compactBody := openAICompactFallbackErrorResponse(resp, signal)
if shouldFailoverOpenAIPassthroughResponse(account, compactResp.StatusCode, compactBody) {
return nil, s.handleFailoverErrorResponsePassthrough(ctx, compactResp, c, account, body, compactBody)
}
return nil, s.handleErrorResponsePassthrough(ctx, compactResp, c, account, body, compactBody)
}
_ = resp.Body.Close()
return nil, handleErr
}
@@ -713,7 +730,7 @@ func shouldFailoverOpenAIPassthroughResponse(account *Account, statusCode int, r
if isOpenAIContextWindowError("", responseBody) {
return false
}
if isOpenAIUpstreamAccessStateError("", responseBody) {
if isOpenAIHTTPUpstreamAccessStateError(statusCode, "", responseBody) {
return true
}
if isOpenAIRequestBodyTooLargeError(statusCode, "", responseBody) {
@@ -1520,7 +1537,14 @@ func (s *OpenAIGatewayService) handleOpenAIStreamTerminalAccountSideEffects(
if c != nil && c.Request != nil {
ctx = c.Request.Context()
}
return statusCode, s.handleOpenAIAccountUpstreamError(ctx, account, statusCode, headers, payload)
accountHeaders := headers
if statusCode == http.StatusTooManyRequests {
// The enclosing HTTP response succeeded. Its quota snapshot describes
// normal account state and must not become the reset for a semantic 429
// carried by a stream terminal event.
accountHeaders = nil
}
return statusCode, s.handleOpenAIAccountUpstreamError(ctx, account, statusCode, accountHeaders, payload)
default:
return statusCode, false
}
@@ -1033,17 +1033,11 @@ func normalizeOpenAIResponsesWebSocketCompatibilityBody(body []byte, account *Ac
changed = true
}
}
if account != nil && shouldSanitizeOpenAIResponsesToolSchemas(account.Platform) {
if toolBody, toolChanged, err := sanitizeOpenAIResponsesToolParameterTypes(normalized); err != nil {
return body, false, fmt.Errorf("normalize websocket tool parameter types: %w", err)
} else if toolChanged {
normalized = toolBody
changed = true
}
if patternBody, patternChanged, err := sanitizeOpenAIResponsesToolSchemaPatterns(normalized); err != nil {
return body, false, fmt.Errorf("normalize websocket tool schema patterns: %w", err)
} else if patternChanged {
normalized = patternBody
if account != nil {
if schemaBody, schemaChanged, err := sanitizeOpenAIResponsesToolSchemasForPlatform(normalized, account.Platform); err != nil {
return body, false, fmt.Errorf("normalize websocket tool schemas: %w", err)
} else if schemaChanged {
normalized = schemaBody
changed = true
}
}
@@ -441,6 +441,62 @@ func TestOpenAIResponseFlush_FailedAndErrorEventsFlushAtBoundaries(t *testing.T)
})
}
func TestOpenAIResponseFlush_BareErrorFollowedByCompletedUsesCompletedTerminal(t *testing.T) {
body := "data: {\"type\":\"error\",\"error\":{\"code\":\"transient\",\"message\":\"retrying\"}}\n\n" +
"data: {\"type\":\"response.completed\",\"response\":{\"id\":\"resp_recovered\",\"status\":\"completed\",\"output\":[],\"usage\":{\"input_tokens\":7,\"output_tokens\":3}}}\n\n"
recorder := newOpenAIResponseFlushRecorder()
result, err := runOpenAIResponseFlushTest(recorder, io.NopCloser(strings.NewReader(body)), config.GatewayConfig{})
require.NoError(t, err)
require.NotNil(t, result)
require.Equal(t, 7, result.usage.InputTokens)
require.Equal(t, 3, result.usage.OutputTokens)
gotBody, _ := recorder.snapshot()
require.NotContains(t, gotBody, `"type":"error"`)
require.NotContains(t, gotBody, `"type":"response.failed"`)
require.Contains(t, gotBody, `"type":"response.completed"`)
}
func TestOpenAIResponseFlush_CompatibleAPIKeyDoesNotUseCodexBareErrorSynthesis(t *testing.T) {
body := "data: {\"type\":\"error\",\"error\":{\"code\":\"provider_error\",\"message\":\"provider failed\"}}\n\n"
recorder := newOpenAIResponseFlushRecorder()
account := &Account{ID: 2, Platform: PlatformOpenAI, Type: AccountTypeAPIKey}
result, err := runOpenAIResponseFlushTestWithAccount(recorder, io.NopCloser(strings.NewReader(body)), config.GatewayConfig{}, account)
require.Error(t, err)
require.NotNil(t, result)
gotBody, _ := recorder.snapshot()
require.Contains(t, gotBody, `"type":"error"`)
require.NotContains(t, gotBody, `"type":"response.failed"`)
}
func TestOpenAIResponseFlush_RecentBareErrorAllowsCompletedBeforeIdleTimeout(t *testing.T) {
reader, writer := io.Pipe()
defer func() { _ = writer.Close() }()
recorder := newOpenAIResponseFlushRecorder()
resultCh, errCh := runOpenAIResponseFlushTestAsync(recorder, reader, config.GatewayConfig{StreamDataIntervalTimeout: 1})
// Place the bare error shortly before the first ticker firing. It is fresh
// data, so the ticker must leave the stream open for an authoritative event.
time.Sleep(700 * time.Millisecond)
_, err := io.WriteString(writer, "data: {\"type\":\"error\",\"error\":{\"code\":\"transient\",\"message\":\"retrying\"}}\n\n")
require.NoError(t, err)
time.Sleep(500 * time.Millisecond)
_, err = io.WriteString(writer, "data: {\"type\":\"response.completed\",\"response\":{\"id\":\"resp_late\",\"status\":\"completed\",\"output\":[],\"usage\":{\"input_tokens\":5,\"output_tokens\":1}}}\n\n")
require.NoError(t, err)
require.NoError(t, writer.Close())
require.NoError(t, <-errCh)
result := <-resultCh
require.NotNil(t, result)
require.Equal(t, 5, result.usage.InputTokens)
gotBody, _ := recorder.snapshot()
require.Contains(t, gotBody, `"type":"response.completed"`)
require.NotContains(t, gotBody, `"type":"response.failed"`)
}
func TestOpenAIResponseFlush_BareErrorTimeoutSynthesizesFailed(t *testing.T) {
tests := []struct {
name string
@@ -534,6 +590,10 @@ func TestOpenAIResponseFlush_ClientDisconnectStillDrainsUsage(t *testing.T) {
}
func runOpenAIResponseFlushTest(recorder *openAIResponseFlushRecorder, body io.ReadCloser, gatewayCfg config.GatewayConfig) (*openaiStreamingResult, error) {
return runOpenAIResponseFlushTestWithAccount(recorder, body, gatewayCfg, &Account{ID: 1, Platform: PlatformOpenAI, Type: AccountTypeOAuth})
}
func runOpenAIResponseFlushTestWithAccount(recorder *openAIResponseFlushRecorder, body io.ReadCloser, gatewayCfg config.GatewayConfig, account *Account) (*openaiStreamingResult, error) {
gin.SetMode(gin.TestMode)
c, _ := gin.CreateTestContext(recorder)
c.Request = httptest.NewRequest(http.MethodPost, "/v1/responses", nil)
@@ -546,7 +606,7 @@ func runOpenAIResponseFlushTest(recorder *openAIResponseFlushRecorder, body io.R
Header: http.Header{"Content-Type": []string{"text/event-stream"}},
Body: body,
}
return svc.handleStreamingResponse(context.Background(), resp, c, &Account{ID: 1, Platform: PlatformOpenAI}, time.Now(), "gpt-5", "gpt-5")
return svc.handleStreamingResponse(context.Background(), resp, c, account, time.Now(), "gpt-5", "gpt-5")
}
func runOpenAIResponseFlushTestAsync(recorder *openAIResponseFlushRecorder, body io.ReadCloser, gatewayCfg config.GatewayConfig) (<-chan *openaiStreamingResult, <-chan error) {
@@ -251,7 +251,7 @@ func (s *OpenAIGatewayService) handleStreamingResponseWithReasoning(ctx context.
capacityFailoverSuppressedLogged := false
failedMessage := ""
clientOutputStarted := false
codexFailureTerminal := account != nil && account.Platform == PlatformOpenAI
codexFailureTerminal := account != nil && account.IsOpenAIOAuthLike()
upstreamRequestID := strings.TrimSpace(resp.Header.Get("x-request-id"))
var streamEarlyErr error
terminalFailurePending := false
@@ -478,6 +478,18 @@ func (s *OpenAIGatewayService) handleStreamingResponseWithReasoning(ctx context.
if data, ok := extractOpenAISSEDataLine(line); ok {
dataBytes := []byte(data)
eventType := effectiveOpenAISSEEventType(dataBytes, pendingSSEEventType)
if codexFailureTerminal && sawBareError && !sawResponseFailed &&
(eventType == "response.completed" || eventType == "response.done") {
// A later successful terminal is authoritative over a pending bare
// error. Keep its usage and terminal visible to the client.
sawBareError = false
sawFailedEvent = false
terminalFailurePending = false
suppressCurrentEvent = false
bareErrorPayload = nil
bareErrorAccountSideEffectsPending = false
failedMessage = ""
}
if codexFailureTerminal && sawBareError && !sawResponseFailed && eventType != "response.failed" {
suppressCurrentEvent = true
}
@@ -862,10 +874,6 @@ func (s *OpenAIGatewayService) handleStreamingResponseWithReasoning(ctx context.
}
case <-intervalCh:
if codexFailureTerminal && sawBareError && !sawResponseFailed {
_ = resp.Body.Close()
return finalizeStream()
}
if failureDelivered {
return resultWithUsage(), fmt.Errorf("upstream response failed: %s", failedMessage)
}
@@ -873,6 +881,10 @@ func (s *OpenAIGatewayService) handleStreamingResponseWithReasoning(ctx context.
if time.Since(lastRead) < streamInterval {
continue
}
if codexFailureTerminal && sawBareError && !sawResponseFailed {
_ = resp.Body.Close()
return finalizeStream()
}
if clientDisconnected {
return resultWithUsage(), fmt.Errorf("stream usage incomplete after timeout")
}
@@ -899,7 +911,7 @@ func (s *OpenAIGatewayService) handleStreamingResponseWithReasoning(ctx context.
stopFirstOutputTimer()
continue
}
if codexFailureTerminal && sawBareError && !sawResponseFailed {
if codexFailureTerminal && sawBareError && !sawResponseFailed && len(events) == 0 {
_ = resp.Body.Close()
return finalizeStream()
}
@@ -1181,6 +1193,12 @@ func (s *OpenAIGatewayService) parseSSEUsageBytesWithType(data []byte, eventType
if usage == nil || len(data) == 0 || bytes.Equal(bytes.TrimSpace(data), []byte("[DONE]")) {
return
}
// Usage is absent from nearly every delta event. Avoid full JSON validation
// and four gjson path scans on that hot path while retaining progressive
// usage from compatible upstreams on any event type.
if !bytes.Contains(data, []byte(`"usage"`)) {
return
}
parsedUsage, ok := extractOpenAIUsageFromJSONBytes(data)
if !ok {
return
@@ -1368,6 +1386,57 @@ func extractOpenAIResponseIDFromJSONBytes(body []byte) string {
return strings.TrimSpace(gjson.GetBytes(body, "response.id").String())
}
const openAIHTTPResponseOwnerContextKey = "openai_http_response_owner"
type openAIHTTPResponseOwner struct {
userID int64
apiKeyID int64
}
// SetOpenAIHTTPResponseOwner marks the authenticated downstream owner whose
// successful Responses IDs may be used for later HTTP continuations.
func SetOpenAIHTTPResponseOwner(c *gin.Context, userID, apiKeyID int64) {
if c == nil || userID <= 0 || apiKeyID <= 0 {
return
}
c.Set(openAIHTTPResponseOwnerContextKey, openAIHTTPResponseOwner{userID: userID, apiKeyID: apiKeyID})
}
// ValidateOpenAIHTTPResponseOwner authorizes a continuation by downstream
// tenant. API key identity is retained in the binding, while keys owned by the
// same user remain interoperable.
func (s *OpenAIGatewayService) ValidateOpenAIHTTPResponseOwner(
ctx context.Context,
groupID int64,
responseID string,
userID, apiKeyID int64,
) (bool, error) {
if s == nil || strings.TrimSpace(responseID) == "" || userID <= 0 || apiKeyID <= 0 {
return false, nil
}
ownerUserID, ownerAPIKeyID, found, err := s.getOpenAIWSStateStore().GetHTTPResponseOwner(ctx, groupID, responseID)
if err != nil || !found {
return false, err
}
return ownerUserID == userID || (ownerUserID <= 0 && ownerAPIKeyID == apiKeyID), nil
}
// BindOpenAIHTTPResponseOwner records an HTTP continuation owner independently
// from the upstream account selected for that response.
func (s *OpenAIGatewayService) BindOpenAIHTTPResponseOwner(
ctx context.Context,
groupID int64,
responseID string,
userID, apiKeyID int64,
) error {
if s == nil {
return nil
}
return s.getOpenAIWSStateStore().BindHTTPResponseOwner(
ctx, groupID, responseID, userID, apiKeyID, s.openAIWSResponseStickyTTL(),
)
}
func (s *OpenAIGatewayService) bindHTTPResponseAccount(ctx context.Context, c *gin.Context, account *Account, responseID string) {
if s == nil || account == nil || account.ID <= 0 {
return
@@ -1383,6 +1452,21 @@ func (s *OpenAIGatewayService) bindHTTPResponseAccount(ctx context.Context, c *g
groupID := getOpenAIGroupIDFromContext(c)
ttl := s.openAIWSResponseStickyTTL()
logOpenAIWSBindResponseAccountWarn(groupID, account.ID, responseID, store.BindResponseAccount(ctx, groupID, responseID, account.ID, ttl))
if rawOwner, ok := c.Get(openAIHTTPResponseOwnerContextKey); ok {
if owner, ok := rawOwner.(openAIHTTPResponseOwner); ok && owner.userID > 0 && owner.apiKeyID > 0 {
if err := s.BindOpenAIHTTPResponseOwner(ctx, groupID, responseID, owner.userID, owner.apiKeyID); err != nil {
logger.L().Warn(
"openai.http_bind_response_owner_failed",
zap.Int64("group_id", groupID),
zap.Int64("account_id", account.ID),
zap.Int64("user_id", owner.userID),
zap.Int64("api_key_id", owner.apiKeyID),
zap.String("response_id", truncateOpenAIWSLogValue(responseID, openAIWSIDValueMaxLen)),
zap.Error(err),
)
}
}
}
}
func openAIUsageFromGJSON(value gjson.Result) (OpenAIUsage, bool) {
@@ -1665,9 +1749,6 @@ func extractOpenAISSETerminalEvent(body string) (string, []byte, bool) {
var terminalType string
var terminalPayload []byte
forEachOpenAISSEFrame(body, func(eventType string, data []byte) {
if terminalPayload != nil {
return
}
switch eventType {
case "response.completed", "response.done", "response.failed", "response.incomplete", "response.cancelled", "response.canceled", "error":
terminalType = eventType
@@ -536,6 +536,7 @@ func TestOpenAIGatewayService_BindHTTPResponseAccount(t *testing.T) {
c.Request = httptest.NewRequest(http.MethodPost, "/openai/v1/responses", nil)
groupID := int64(4201)
c.Set("api_key", &APIKey{ID: 501, GroupID: &groupID})
SetOpenAIHTTPResponseOwner(c, 601, 501)
svc := &OpenAIGatewayService{}
account := &Account{ID: 37001, Platform: PlatformOpenAI, Type: AccountTypeAPIKey}
@@ -544,6 +545,22 @@ func TestOpenAIGatewayService_BindHTTPResponseAccount(t *testing.T) {
got, err := svc.getOpenAIWSStateStore().GetResponseAccount(context.Background(), groupID, "resp_http_001")
require.NoError(t, err)
require.Equal(t, account.ID, got)
owned, err := svc.ValidateOpenAIHTTPResponseOwner(context.Background(), groupID, "resp_http_001", 601, 501)
require.NoError(t, err)
require.True(t, owned)
owned, err = svc.ValidateOpenAIHTTPResponseOwner(context.Background(), groupID, "resp_http_001", 601, 502)
require.NoError(t, err)
require.True(t, owned, "API keys owned by the same downstream user remain interoperable")
owned, err = svc.ValidateOpenAIHTTPResponseOwner(context.Background(), groupID, "resp_http_001", 602, 501)
require.NoError(t, err)
require.False(t, owned)
owned, err = svc.ValidateOpenAIHTTPResponseOwner(context.Background(), groupID, "resp_unknown", 601, 501)
require.NoError(t, err)
require.False(t, owned)
}
func TestOpenAIGatewayService_GenerateExplicitSessionHash_SkipsContentFallback(t *testing.T) {
@@ -260,7 +260,7 @@ func (s *OpenAIGatewayService) shouldFailoverOpenAIUpstreamResponse(statusCode i
if isOpenAIContextWindowError(upstreamMsg, upstreamBody) {
return false
}
if isOpenAIUpstreamAccessStateError(upstreamMsg, upstreamBody) {
if isOpenAIHTTPUpstreamAccessStateError(statusCode, upstreamMsg, upstreamBody) {
return true
}
if isOpenAIRequestBodyTooLargeError(statusCode, upstreamMsg, upstreamBody) {
@@ -306,7 +306,7 @@ func newOpenAIUpstreamFailoverError(
failoverErr.ClientStatusCode = http.StatusRequestEntityTooLarge
failoverErr.ClientMessage = OpenAIRequestBodyTooLargeClientMessage
}
if isOpenAIUpstreamAccessStateError(upstreamMsg, responseBody) {
if isOpenAIHTTPUpstreamAccessStateError(statusCode, upstreamMsg, responseBody) {
failoverErr.RetryableOnSameAccount = false
failoverErr.RequestScopedTransient = false
failoverErr.Stage = GatewayFailureStageAccountAuth
@@ -359,59 +359,42 @@ const (
)
// isOpenAIUpstreamAccessStateError recognizes provider-side credential state
// failures from explicit structured fields. Valid JSON is never scanned as a
// blob because it may contain echoed user input with the same words.
func isOpenAIUpstreamAccessStateError(upstreamMsg string, body []byte) bool {
matchCode := func(value string) bool {
value = strings.ToLower(strings.TrimSpace(value))
if value == "deactivated_workspace" {
return true
}
for _, subject := range []string{"workspace", "account", "organization", "org"} {
for _, state := range []string{"deactivated", "disabled", "suspended"} {
if value == subject+"_"+state || value == state+"_"+subject {
return true
}
}
}
// failures only from explicit structured codes. Free-form messages may contain
// echoed user input, including inside stream terminal error.message fields.
func isOpenAIUpstreamAccessStateError(_ string, body []byte) bool {
if len(body) == 0 || !gjson.ValidBytes(body) {
return false
}
matchMessage := func(value string) bool {
value = strings.ToLower(strings.TrimSpace(value))
for _, subject := range []string{"workspace", "account", "organization", "org"} {
for _, state := range []string{"deactivated", "disabled", "suspended"} {
if strings.Contains(value, subject+" is "+state) ||
strings.Contains(value, subject+" has been "+state) ||
strings.Contains(value, subject+" "+state) {
return true
}
}
}
return false
}
if matchMessage(upstreamMsg) {
return true
}
if len(body) == 0 {
return false
}
if !gjson.ValidBytes(body) {
return matchMessage(string(body)) || matchCode(string(body))
}
for _, path := range []string{"error.code", "response.error.code", "detail.code", "code"} {
if matchCode(gjson.GetBytes(body, path).String()) {
return true
}
}
for _, path := range []string{"error.message", "response.error.message", "detail.message", "detail", "message"} {
if matchMessage(gjson.GetBytes(body, path).String()) {
if isOpenAIUpstreamAccessStateCode(gjson.GetBytes(body, path).String()) {
return true
}
}
return false
}
func isOpenAIUpstreamAccessStateCode(value string) bool {
value = strings.ToLower(strings.TrimSpace(value))
if value == "deactivated_workspace" {
return true
}
for _, subject := range []string{"workspace", "account", "organization", "org"} {
for _, state := range []string{"deactivated", "disabled", "suspended"} {
if value == subject+"_"+state || value == state+"_"+subject {
return true
}
}
}
return false
}
// isOpenAIHTTPUpstreamAccessStateError is deliberately status-independent:
// known provider codes are durable evidence, while 401/403 messages without
// such a code must flow through the existing authentication/403 policies.
func isOpenAIHTTPUpstreamAccessStateError(_ int, _ string, body []byte) bool {
return isOpenAIUpstreamAccessStateError("", body)
}
func openAICapacityShedClientMessage(upstreamMsg string, body []byte) string {
for _, candidate := range []string{
upstreamMsg,
+1 -1
View File
@@ -663,7 +663,7 @@ func (s *OpenAIGatewayService) forwardOpenAIImagesAPIKey(
if account.IsOpenAIOAuthLike() && resp.StatusCode == http.StatusTooManyRequests {
return nil, s.newOpenAIAccountFailoverError(account, resp.StatusCode, resp.Header, respBody, upstreamMsg, shouldDisable, retryableOnSameAccount)
}
if isOpenAIUpstreamAccessStateError(upstreamMsg, respBody) {
if isOpenAIHTTPUpstreamAccessStateError(resp.StatusCode, upstreamMsg, respBody) {
return nil, newOpenAIUpstreamFailoverError(resp.StatusCode, resp.Header, respBody, upstreamMsg, retryableOnSameAccount)
}
return nil, &UpstreamFailoverError{StatusCode: resp.StatusCode, ResponseBody: respBody, RetryableOnSameAccount: retryableOnSameAccount}
@@ -9,6 +9,7 @@ import (
"github.com/Wei-Shaw/sub2api/internal/pkg/apicompat"
"github.com/stretchr/testify/require"
"github.com/tidwall/gjson"
)
func TestEffectiveOpenAISSEEventTypePrefersPayload(t *testing.T) {
@@ -31,6 +32,17 @@ func TestExtractOpenAISSETerminalEventUsesEventField(t *testing.T) {
require.Equal(t, "provider failed", extractOpenAISSEErrorMessage(payload))
}
func TestExtractOpenAISSETerminalEventUsesFinalAuthoritativeTerminal(t *testing.T) {
t.Parallel()
body := "data: {\"type\":\"error\",\"error\":{\"message\":\"recovering\"}}\n\n" +
"data: {\"type\":\"response.completed\",\"response\":{\"id\":\"resp_1\",\"status\":\"completed\"}}\n\n"
eventType, payload, ok := extractOpenAISSETerminalEvent(body)
require.True(t, ok)
require.Equal(t, "response.completed", eventType)
require.Equal(t, "resp_1", gjson.GetBytes(payload, "response.id").String())
}
func TestParseSSEUsageEffectiveTerminalRules(t *testing.T) {
t.Parallel()
@@ -47,6 +59,17 @@ func TestParseSSEUsageEffectiveTerminalRules(t *testing.T) {
require.Equal(t, OpenAIUsage{InputTokens: 2}, *usage)
}
func BenchmarkParseSSEUsageNoUsageDelta(b *testing.B) {
svc := &OpenAIGatewayService{}
usage := &OpenAIUsage{}
payload := []byte(`{"type":"response.output_text.delta","delta":"hello"}`)
b.ReportAllocs()
b.ResetTimer()
for i := 0; i < b.N; i++ {
svc.parseSSEUsageBytesWithType(payload, "response.output_text.delta", usage)
}
}
func TestOpenAICompatTerminalResponseSynthesizesBareError(t *testing.T) {
t.Parallel()
@@ -46,11 +46,17 @@ func normalizeOpenAIResponsesLegacyIngress(body []byte) ([]byte, bool, error) {
}
if prompt, hasPrompt := request["prompt"]; hasPrompt {
if input, hasInput := request["input"]; (!hasInput || input == nil) && prompt != nil {
request["input"] = prompt
// Only the legacy string alias is unambiguously equivalent to Responses
// input. Objects are native reusable prompt templates and must remain in
// prompt; arrays and other shapes are left for upstream validation rather
// than being relabeled as a structurally different input value.
if promptText, isLegacyString := prompt.(string); isLegacyString {
if input, hasInput := request["input"]; !hasInput || input == nil {
request["input"] = promptText
}
delete(request, "prompt")
changed = true
}
delete(request, "prompt")
changed = true
}
if _, hasCommands := request["commands"]; hasCommands {
delete(request, "commands")
@@ -70,3 +70,25 @@ func TestNormalizeOpenAIResponsesLegacyIngressKeepsPromptAliasAndDropsCommands(t
require.False(t, gjson.GetBytes(normalized, "prompt").Exists())
require.False(t, gjson.GetBytes(normalized, "commands").Exists())
}
func TestNormalizeOpenAIResponsesLegacyIngressPreservesNativePromptTemplate(t *testing.T) {
body := []byte(`{"model":"gpt-5.4","prompt":{"id":"pmpt_abc","version":"7","variables":{"topic":"ownership"}}}`)
normalized, changed, err := normalizeOpenAIResponsesLegacyIngress(body)
require.NoError(t, err)
require.False(t, changed)
require.JSONEq(t, string(body), string(normalized))
require.Equal(t, "pmpt_abc", gjson.GetBytes(normalized, "prompt.id").String())
require.False(t, gjson.GetBytes(normalized, "input").Exists())
}
func TestNormalizeOpenAIResponsesLegacyIngressPreservesUnknownPromptShapeWhileDroppingCommands(t *testing.T) {
body := []byte(`{"model":"gpt-5.4","prompt":["one","two"],"commands":[{"name":"legacy"}]}`)
normalized, changed, err := normalizeOpenAIResponsesLegacyIngress(body)
require.NoError(t, err)
require.True(t, changed)
require.Equal(t, int64(2), gjson.GetBytes(normalized, "prompt.#").Int())
require.False(t, gjson.GetBytes(normalized, "input").Exists())
require.False(t, gjson.GetBytes(normalized, "commands").Exists())
}
@@ -70,35 +70,20 @@ func sanitizeOpenAIResponsesInputItemIDs(body []byte) ([]byte, bool, error) {
type inputItem struct {
body []byte
itemType string
id string
callID string
stripID bool
stripCallID bool
drop bool
isObject bool
}
items := make([]inputItem, 0)
strippedIDs := make(map[string]struct{})
validCallIDs := make(map[string]struct{})
input.ForEach(func(_, item gjson.Result) bool {
parsed := inputItem{body: []byte(item.Raw), isObject: item.IsObject()}
parsed := inputItem{body: []byte(item.Raw)}
if item.IsObject() {
itemType := item.Get("type")
id := item.Get("id")
parsed.itemType = strings.TrimSpace(itemType.String())
parsed.callID = strings.TrimSpace(item.Get("call_id").String())
parsed.stripCallID = item.Get("call_id").Exists() && shouldStripOpenAIResponsesNonPairCallID(parsed.itemType)
trimmedItemType := strings.TrimSpace(itemType.String())
parsed.stripCallID = item.Get("call_id").Exists() && shouldStripOpenAIResponsesNonPairCallID(trimmedItemType)
if id.Type == gjson.String {
parsed.id = id.String()
parsed.stripID = shouldStripOpenAIResponsesInputItemID(parsed.itemType, parsed.id)
if parsed.stripID && parsed.id != "" {
strippedIDs[parsed.id] = struct{}{}
}
}
if isCodexToolCallContextItemType(parsed.itemType) && parsed.callID != "" {
validCallIDs[parsed.callID] = struct{}{}
parsed.stripID = shouldStripOpenAIResponsesInputItemID(trimmedItemType, id.String())
}
}
items = append(items, parsed)
@@ -115,51 +100,8 @@ func sanitizeOpenAIResponsesInputItemIDs(body []byte) ([]byte, bool, error) {
return body, false, nil
}
// First decide which outputs become dangling because their call_id points at
// an item ID that is being removed and no call item owns that call_id. This
// must happen before computing retained IDs: a dropped output cannot keep an
// item_reference alive merely because it used to have the same id.
for index := range items {
item := &items[index]
if !item.isObject || !isCodexToolCallOutputItemType(item.itemType) {
continue
}
if _, pointsAtStrippedID := strippedIDs[item.callID]; !pointsAtStrippedID {
continue
}
if _, hasMatchingCallID := validCallIDs[item.callID]; !hasMatchingCallID {
item.drop = true
}
}
removedItemIDs := make(map[string]struct{}, len(strippedIDs))
retainedItemIDs := make(map[string]struct{}, len(items))
for _, item := range items {
if item.id == "" || item.itemType == "item_reference" {
continue
}
if item.stripID || item.drop {
removedItemIDs[item.id] = struct{}{}
continue
}
retainedItemIDs[item.id] = struct{}{}
}
for id := range retainedItemIDs {
delete(removedItemIDs, id)
}
rebuiltItems := make([][]byte, 0, len(items))
for index, item := range items {
if item.isObject {
if item.itemType == "item_reference" {
if _, dangling := removedItemIDs[item.id]; dangling {
continue
}
}
if item.drop {
continue
}
}
itemBody := item.body
if item.stripID {
var err error
@@ -45,11 +45,11 @@ func TestOpenAIResponsesInputItemIDPrefixUsesObservedOutputContracts(t *testing.
}
}
func TestSanitizeOpenAIResponsesInputItemIDsKeepsReferenceGraphConsistent(t *testing.T) {
func TestSanitizeOpenAIResponsesInputItemIDsDoesNotCascadeAcrossIDNamespaces(t *testing.T) {
body := []byte(`{"input":[
{"type":"function_call","id":"item_bad_call","call_id":"call_valid","name":"lookup","arguments":"{}"},
{"type":"function_call_output","call_id":"call_valid","output":"preserve paired output"},
{"type":"function_call_output","call_id":"item_bad_call","output":"drop dangling output"},
{"type":"function_call_output","call_id":"item_bad_call","output":"preserve opaque output"},
{"type":"item_reference","id":"item_bad_call"},
{"type":"item_reference","id":"remote_valid"},
{"type":"custom_tool_call","id":"ctc_valid","call_id":"ctco_bad_output","name":"apply_patch","input":"patch"},
@@ -61,14 +61,16 @@ func TestSanitizeOpenAIResponsesInputItemIDsKeepsReferenceGraphConsistent(t *tes
require.NoError(t, err)
require.True(t, changed)
items := gjson.GetBytes(sanitized, "input").Array()
require.Len(t, items, 5)
require.Len(t, items, 7)
require.False(t, items[0].Get("id").Exists())
require.Equal(t, "call_valid", items[0].Get("call_id").String())
require.Equal(t, "preserve paired output", items[1].Get("output").String())
require.Equal(t, "remote_valid", items[2].Get("id").String())
require.Equal(t, "ctc_valid", items[3].Get("id").String())
require.False(t, items[4].Get("id").Exists())
require.Equal(t, "ctco_bad_output", items[4].Get("call_id").String())
require.Equal(t, "preserve opaque output", items[2].Get("output").String())
require.Equal(t, "item_bad_call", items[3].Get("id").String())
require.Equal(t, "remote_valid", items[4].Get("id").String())
require.Equal(t, "ctc_valid", items[5].Get("id").String())
require.False(t, items[6].Get("id").Exists())
require.Equal(t, "ctco_bad_output", items[6].Get("call_id").String())
}
func TestSanitizeOpenAIResponsesInputItemIDsLeavesUnrelatedReferencesUntouched(t *testing.T) {
@@ -93,7 +95,7 @@ func TestSanitizeOpenAIResponsesInputItemIDsPreservesReferenceToDuplicateRetaine
require.Equal(t, "ctc_shared", gjson.GetBytes(sanitized, "input.2.id").String())
}
func TestSanitizeOpenAIResponsesInputItemIDsClosesReferencesAfterDroppingOutput(t *testing.T) {
func TestSanitizeOpenAIResponsesInputItemIDsPreservesOpaqueOutputsAndReferences(t *testing.T) {
body := []byte(`{"input":[
{"type":"function_call","id":"item_shared","call_id":"call_real"},
{"type":"function_call_output","id":"item_shared","call_id":"item_shared","output":"dangling"},
@@ -106,10 +108,12 @@ func TestSanitizeOpenAIResponsesInputItemIDsClosesReferencesAfterDroppingOutput(
require.NoError(t, err)
require.True(t, changed)
require.Len(t, gjson.GetBytes(sanitized, "input").Array(), 3)
require.Len(t, gjson.GetBytes(sanitized, "input").Array(), 5)
require.False(t, gjson.GetBytes(sanitized, "input.0.id").Exists())
require.Equal(t, "kept_output", gjson.GetBytes(sanitized, "input.1.id").String())
require.Equal(t, "kept_output", gjson.GetBytes(sanitized, "input.2.id").String())
require.Equal(t, "dangling", gjson.GetBytes(sanitized, "input.1.output").String())
require.Equal(t, "item_shared", gjson.GetBytes(sanitized, "input.2.id").String())
require.Equal(t, "kept_output", gjson.GetBytes(sanitized, "input.3.id").String())
require.Equal(t, "kept_output", gjson.GetBytes(sanitized, "input.4.id").String())
second, changedAgain, err := sanitizeOpenAIResponsesInputItemIDs(sanitized)
require.NoError(t, err)
@@ -26,13 +26,45 @@ const (
var errOpenAIResponsesToolSchemaLimit = errors.New("OpenAI Responses tool schema safety limit exceeded")
// shouldSanitizeOpenAIResponsesToolSchemas centralizes the platform boundary
// for callers. These rewrites describe OpenAI Responses constraints, not the
// behavior of every provider routed through the generic OpenAI gateway.
func shouldSanitizeOpenAIResponsesToolSchemas(platform string) bool {
// shouldRepairOpenAIResponsesNullToolSchemaType reports whether the upstream
// path requires a concrete object type at a function tool's parameter root.
// This defect is shared by the OpenAI, Anthropic, and CN-compatible paths.
func shouldRepairOpenAIResponsesNullToolSchemaType(platform string) bool {
return platform == PlatformOpenAI || platform == PlatformAnthropic || IsCNProvider(platform)
}
// shouldSanitizeOpenAIResponsesToolSchemaPatterns is intentionally narrower:
// regex lookaround rejection is an OpenAI-specific schema constraint.
func shouldSanitizeOpenAIResponsesToolSchemaPatterns(platform string) bool {
return platform == PlatformOpenAI
}
func sanitizeOpenAIResponsesToolSchemasForPlatform(body []byte, platform string) ([]byte, bool, error) {
normalized := body
changed := false
if shouldRepairOpenAIResponsesNullToolSchemaType(platform) {
next, repaired, err := sanitizeOpenAIResponsesToolParameterTypes(normalized)
if err != nil {
return body, false, fmt.Errorf("sanitize OpenAI Responses tool parameters: %w", err)
}
if repaired {
normalized = next
changed = true
}
}
if shouldSanitizeOpenAIResponsesToolSchemaPatterns(platform) {
next, sanitized, err := sanitizeOpenAIResponsesToolSchemaPatterns(normalized)
if err != nil {
return body, false, fmt.Errorf("sanitize OpenAI Responses tool schema patterns: %w", err)
}
if sanitized {
normalized = next
changed = true
}
}
return normalized, changed, nil
}
// sanitizeOpenAIResponsesToolSchemaPatterns removes only schema constraints
// containing regex lookaround, which OpenAI rejects. It deliberately does not
// descend into instance-valued keywords such as default, examples, const, or
@@ -239,11 +239,59 @@ func TestSanitizeOpenAIResponsesToolSchemas_InvalidAndTrailingJSON(t *testing.T)
}
}
func TestShouldSanitizeOpenAIResponsesToolSchemas_PlatformBoundary(t *testing.T) {
require.True(t, shouldSanitizeOpenAIResponsesToolSchemas(PlatformOpenAI))
for _, platform := range []string{PlatformGrok, PlatformKimi, PlatformZhipu, PlatformDeepseek, PlatformComposite, ""} {
require.False(t, shouldSanitizeOpenAIResponsesToolSchemas(platform), platform)
func TestOpenAIResponsesToolSchemaCapabilities_PlatformBoundary(t *testing.T) {
tests := []struct {
platform string
repairNullType bool
removeLookaround bool
}{
{PlatformOpenAI, true, true},
{PlatformAnthropic, true, false},
{PlatformKimi, true, false},
{PlatformZhipu, true, false},
{PlatformDeepseek, true, false},
{PlatformGrok, false, false},
{PlatformGemini, false, false},
{PlatformAntigravity, false, false},
{PlatformComposite, false, false},
{"", false, false},
}
for _, tt := range tests {
t.Run(tt.platform, func(t *testing.T) {
require.Equal(t, tt.repairNullType, shouldRepairOpenAIResponsesNullToolSchemaType(tt.platform))
require.Equal(t, tt.removeLookaround, shouldSanitizeOpenAIResponsesToolSchemaPatterns(tt.platform))
})
}
}
func TestSanitizeOpenAIResponsesToolSchemasForPlatform_ReplayBoundary(t *testing.T) {
body := []byte(`{"tools":[{"type":"function","parameters":{"type":null,"properties":{"query":{"type":"string","pattern":"(?=keep)"}}}}]}`)
// A malformed tool definition may be replayed after account failover. Every
// compatible account must repair it, while non-OpenAI providers retain their
// supported regex semantics.
for _, platform := range []string{PlatformAnthropic, PlatformKimi, PlatformZhipu, PlatformDeepseek} {
t.Run(platform, func(t *testing.T) {
for attempt := 0; attempt < 2; attempt++ {
normalized, changed, err := sanitizeOpenAIResponsesToolSchemasForPlatform(body, platform)
require.NoError(t, err)
require.True(t, changed)
require.Equal(t, "object", gjson.GetBytes(normalized, "tools.0.parameters.type").String())
require.Equal(t, "(?=keep)", gjson.GetBytes(normalized, "tools.0.parameters.properties.query.pattern").String())
}
})
}
openAI, changed, err := sanitizeOpenAIResponsesToolSchemasForPlatform(body, PlatformOpenAI)
require.NoError(t, err)
require.True(t, changed)
require.Equal(t, "object", gjson.GetBytes(openAI, "tools.0.parameters.type").String())
require.False(t, gjson.GetBytes(openAI, "tools.0.parameters.properties.query.pattern").Exists())
unsupported, changed, err := sanitizeOpenAIResponsesToolSchemasForPlatform(body, PlatformGrok)
require.NoError(t, err)
require.False(t, changed)
require.Equal(t, string(body), string(unsupported))
}
// 索引映射:只有坏条目被改,前后兄弟条目按原下标保持不变。
@@ -164,6 +164,7 @@ func TestHandleErrorResponse_NonDeterministicStatusesKeepGeneric502(t *testing.T
{"unprocessable", http.StatusUnprocessableEntity, `{"error":{"message":"Invalid schema for field messages"}}`,
http.StatusBadGateway, "upstream_error", "Upstream request failed"},
// 401/402/403 是网关运营方的凭据/账单问题,必须继续对客户端屏蔽上游账号状态。
// 403 的自由文本不能升级成 durable access-state typed failover;只有明确结构化 code 才可以。
{"unauthorized", http.StatusUnauthorized, `{"error":{"message":"Incorrect API key provided: sk-abc"}}`,
http.StatusBadGateway, "upstream_error", "Upstream authentication failed, please contact administrator"},
{"forbidden", http.StatusForbidden, `{"error":{"message":"Your account is deactivated"}}`,
@@ -183,15 +184,11 @@ func TestHandleErrorResponse_NonDeterministicStatusesKeepGeneric502(t *testing.T
newOpenAIUpstreamErrorResponse(tc.statusCode, tc.body),
c, newOpenAIUpstreamErrorTestAccount(), nil,
)
require.Error(t, err)
if tc.name == "forbidden" {
var failoverErr *UpstreamFailoverError
require.ErrorAs(t, err, &failoverErr)
require.Equal(t, http.StatusForbidden, failoverErr.StatusCode)
require.False(t, c.Writer.Written())
return
require.False(t, errors.As(err, &failoverErr))
}
require.Error(t, err)
require.Equal(t, tc.wantStatus, rec.Code)
require.Equal(t, tc.wantType, gjson.Get(rec.Body.String(), "error.type").String())
require.Equal(t, tc.wantMsg, gjson.Get(rec.Body.String(), "error.message").String())
@@ -489,6 +489,7 @@ func (s *OpenAIGatewayService) proxyOpenAIWSHTTPBridgeTurn(
pendingClientMessageBytes := int64(0)
capacityFailoverSuppressedLogged := false
clientDisconnected := false
officialOpenAIResponses := account != nil && account.Platform == PlatformOpenAI
bareErrorPending := false
var bareErrorPayload []byte
bareErrorMessage := ""
@@ -642,7 +643,15 @@ func (s *OpenAIGatewayService) proxyOpenAIWSHTTPBridgeTurn(
replayCollector.AddEvent(eventType, upstreamMessage)
var upstreamEventErr error
suppressClientMessage := bareErrorPending && eventType != "response.failed"
if officialOpenAIResponses && bareErrorPending && (eventType == "response.completed" || eventType == "response.done") {
// Some upstreams emit a recoverable bare error before the authoritative
// successful terminal. Do not replace that terminal with a synthetic
// failure or retain side effects from the superseded error.
bareErrorPending = false
bareErrorPayload = nil
bareErrorMessage = ""
}
suppressClientMessage := officialOpenAIResponses && bareErrorPending && eventType != "response.failed"
if eventType == "error" || eventType == "response.failed" {
errMessage := extractOpenAISSEErrorMessage(upstreamMessage)
if errMessage == "" {
@@ -677,7 +686,7 @@ func (s *OpenAIGatewayService) proxyOpenAIWSHTTPBridgeTurn(
return nil, s.newOpenAIStreamFailoverError(c, account, true, resp.Header.Get("x-request-id"), upstreamMessage, errMessage, resp.Header)
}
if account.Platform != PlatformGrok && !failureAccountSideEffectsApplied {
if eventType == "response.failed" || (shouldFailover && !requestScopedCapacity) {
if eventType == "response.failed" || (!officialOpenAIResponses && shouldFailover && !requestScopedCapacity) {
failureAccountSideEffectsApplied = s.handleOpenAIWSFailureAccountSideEffects(ctx, account, mappedModel, resp.Header, upstreamMessage)
}
}
@@ -685,7 +694,7 @@ func (s *OpenAIGatewayService) proxyOpenAIWSHTTPBridgeTurn(
logOpenAICapacityFailoverSuppressed(ctx, account, "ws_http_bridge", resp.Header.Get("x-request-id"), eventType)
capacityFailoverSuppressedLogged = true
}
if eventType == "error" && account.Platform == PlatformGrok {
if eventType == "error" && !officialOpenAIResponses {
upstreamEventErr = errors.New(errMessage)
} else if eventType == "error" {
bareErrorPending = true
@@ -758,6 +758,40 @@ func TestProxyOpenAIWSHTTPBridgeTurnBareErrorEOFSynthesizesFailed(t *testing.T)
require.Equal(t, "resp_eof", gjson.GetBytes(writes[1], "response.id").String())
}
func TestProxyOpenAIWSHTTPBridgeTurnBareErrorFollowedByCompletedUsesCompleted(t *testing.T) {
gin.SetMode(gin.TestMode)
body := strings.Join([]string{
`data: {"type":"response.created","response":{"id":"resp_recovered","status":"in_progress"}}`,
``,
`data: {"type":"error","error":{"code":"transient","message":"retrying"}}`,
``,
`data: {"type":"response.completed","response":{"id":"resp_recovered","status":"completed","output":[],"usage":{"input_tokens":8,"output_tokens":4}}}`,
``,
}, "\n")
upstream := &httpUpstreamRecorder{resp: &http.Response{StatusCode: http.StatusOK, Header: make(http.Header), Body: io.NopCloser(strings.NewReader(body))}}
svc := &OpenAIGatewayService{cfg: &config.Config{}, httpUpstream: upstream}
account := &Account{ID: 113, Platform: PlatformOpenAI, Type: AccountTypeOAuth, Concurrency: 1}
recorder := httptest.NewRecorder()
c, _ := gin.CreateTestContext(recorder)
c.Request = httptest.NewRequest(http.MethodGet, "/v1/responses", nil)
payload := []byte(`{"type":"response.create","model":"gpt-5","input":"hi"}`)
var writes [][]byte
result, err := svc.proxyOpenAIWSHTTPBridgeTurn(context.Background(), c, account, "sk-test", payload, len(payload), "gpt-5", "", "", "", "", 2, func(message []byte) error {
writes = append(writes, append([]byte(nil), message...))
return nil
})
require.NoError(t, err)
require.NotNil(t, result)
require.Equal(t, "response.completed", result.UpstreamTerminalEvent)
require.Equal(t, 8, result.Usage.InputTokens)
require.Equal(t, 4, result.Usage.OutputTokens)
require.Len(t, writes, 2)
require.Equal(t, "response.created", gjson.GetBytes(writes[0], "type").String())
require.Equal(t, "response.completed", gjson.GetBytes(writes[1], "type").String())
}
func TestProxyOpenAIWSHTTPBridgeTurnStagesMetadataBeforeCapacityFailover(t *testing.T) {
gin.SetMode(gin.TestMode)
body := strings.Join([]string{
@@ -13,6 +13,8 @@ import (
const (
openAIWSResponseAccountCachePrefix = "openai:response:"
openAIHTTPResponseOwnerUserPrefix = "openai:http-response-owner:user:"
openAIHTTPResponseOwnerKeyPrefix = "openai:http-response-owner:key:"
openAIWSStateStoreCleanupInterval = time.Minute
openAIWSStateStoreCleanupMaxPerMap = 512
openAIWSStateStoreMaxEntriesPerMap = 65536
@@ -24,6 +26,12 @@ type openAIWSAccountBinding struct {
expiresAt time.Time
}
type openAIHTTPResponseOwnerBinding struct {
userID int64
apiKeyID int64
expiresAt time.Time
}
type openAIWSConnBinding struct {
connID string
expiresAt time.Time
@@ -49,6 +57,8 @@ type OpenAIWSStateStore interface {
BindResponseAccount(ctx context.Context, groupID int64, responseID string, accountID int64, ttl time.Duration) error
GetResponseAccount(ctx context.Context, groupID int64, responseID string) (int64, error)
DeleteResponseAccount(ctx context.Context, groupID int64, responseID string) error
BindHTTPResponseOwner(ctx context.Context, groupID int64, responseID string, userID, apiKeyID int64, ttl time.Duration) error
GetHTTPResponseOwner(ctx context.Context, groupID int64, responseID string) (userID, apiKeyID int64, found bool, err error)
BindResponseConn(responseID, connID string, ttl time.Duration)
GetResponseConn(responseID string) (string, bool)
@@ -68,6 +78,8 @@ type defaultOpenAIWSStateStore struct {
responseToAccountMu sync.RWMutex
responseToAccount map[string]openAIWSAccountBinding
responseOwnerMu sync.RWMutex
responseOwners map[string]openAIHTTPResponseOwnerBinding
responseToConnMu sync.RWMutex
responseToConn map[string]openAIWSConnBinding
sessionToTurnStateMu sync.RWMutex
@@ -83,6 +95,7 @@ func NewOpenAIWSStateStore(cache GatewayCache) OpenAIWSStateStore {
store := &defaultOpenAIWSStateStore{
cache: cache,
responseToAccount: make(map[string]openAIWSAccountBinding, 256),
responseOwners: make(map[string]openAIHTTPResponseOwnerBinding, 256),
responseToConn: make(map[string]openAIWSConnBinding, 256),
sessionToTurnState: make(map[string]openAIWSTurnStateBinding, 256),
sessionToConn: make(map[string]openAIWSSessionConnBinding, 256),
@@ -91,6 +104,72 @@ func NewOpenAIWSStateStore(cache GatewayCache) OpenAIWSStateStore {
return store
}
func (s *defaultOpenAIWSStateStore) BindHTTPResponseOwner(ctx context.Context, groupID int64, responseID string, userID, apiKeyID int64, ttl time.Duration) error {
id := normalizeOpenAIWSResponseID(responseID)
if id == "" || userID <= 0 || apiKeyID <= 0 {
return nil
}
ttl = normalizeOpenAIWSTTL(ttl)
s.maybeCleanup()
mapKey := openAIWSResponseAccountMapKey(groupID, id)
s.responseOwnerMu.Lock()
ensureBindingCapacity(s.responseOwners, mapKey, openAIWSStateStoreMaxEntriesPerMap)
s.responseOwners[mapKey] = openAIHTTPResponseOwnerBinding{
userID: userID, apiKeyID: apiKeyID, expiresAt: time.Now().Add(ttl),
}
s.responseOwnerMu.Unlock()
if s.cache == nil {
return nil
}
cacheCtx, cancel := withOpenAIWSStateStoreRedisTimeout(ctx)
defer cancel()
if err := s.cache.SetSessionAccountID(cacheCtx, groupID, openAIHTTPResponseOwnerCacheKey(openAIHTTPResponseOwnerUserPrefix, id), userID, ttl); err != nil {
return err
}
return s.cache.SetSessionAccountID(cacheCtx, groupID, openAIHTTPResponseOwnerCacheKey(openAIHTTPResponseOwnerKeyPrefix, id), apiKeyID, ttl)
}
func (s *defaultOpenAIWSStateStore) GetHTTPResponseOwner(ctx context.Context, groupID int64, responseID string) (int64, int64, bool, error) {
id := normalizeOpenAIWSResponseID(responseID)
if id == "" {
return 0, 0, false, nil
}
s.maybeCleanup()
now := time.Now()
mapKey := openAIWSResponseAccountMapKey(groupID, id)
s.responseOwnerMu.RLock()
if binding, ok := s.responseOwners[mapKey]; ok && now.Before(binding.expiresAt) {
s.responseOwnerMu.RUnlock()
return binding.userID, binding.apiKeyID, true, nil
}
s.responseOwnerMu.RUnlock()
if s.cache == nil {
return 0, 0, false, nil
}
cacheCtx, cancel := withOpenAIWSStateStoreRedisTimeout(ctx)
defer cancel()
userID, err := s.cache.GetSessionAccountID(cacheCtx, groupID, openAIHTTPResponseOwnerCacheKey(openAIHTTPResponseOwnerUserPrefix, id))
if err != nil || userID <= 0 {
return 0, 0, false, err
}
apiKeyID, err := s.cache.GetSessionAccountID(cacheCtx, groupID, openAIHTTPResponseOwnerCacheKey(openAIHTTPResponseOwnerKeyPrefix, id))
if err != nil || apiKeyID <= 0 {
return 0, 0, false, err
}
s.responseOwnerMu.Lock()
ensureBindingCapacity(s.responseOwners, mapKey, openAIWSStateStoreMaxEntriesPerMap)
s.responseOwners[mapKey] = openAIHTTPResponseOwnerBinding{
userID: userID, apiKeyID: apiKeyID, expiresAt: now.Add(time.Minute),
}
s.responseOwnerMu.Unlock()
return userID, apiKeyID, true, nil
}
func (s *defaultOpenAIWSStateStore) BindResponseAccount(ctx context.Context, groupID int64, responseID string, accountID int64, ttl time.Duration) error {
id := normalizeOpenAIWSResponseID(responseID)
if id == "" || accountID <= 0 {
@@ -115,6 +194,22 @@ func (s *defaultOpenAIWSStateStore) BindResponseAccount(ctx context.Context, gro
return s.cache.SetSessionAccountID(cacheCtx, groupID, cacheKey, accountID, ttl)
}
func cleanupExpiredHTTPResponseOwnerBindings(bindings map[string]openAIHTTPResponseOwnerBinding, now time.Time, maxScan int) {
if len(bindings) == 0 || maxScan <= 0 {
return
}
scanned := 0
for key, binding := range bindings {
if now.After(binding.expiresAt) {
delete(bindings, key)
}
scanned++
if scanned >= maxScan {
break
}
}
}
func (s *defaultOpenAIWSStateStore) GetResponseAccount(ctx context.Context, groupID int64, responseID string) (int64, error) {
id := normalizeOpenAIWSResponseID(responseID)
if id == "" {
@@ -319,6 +414,10 @@ func (s *defaultOpenAIWSStateStore) maybeCleanup() {
cleanupExpiredAccountBindings(s.responseToAccount, now, openAIWSStateStoreCleanupMaxPerMap)
s.responseToAccountMu.Unlock()
s.responseOwnerMu.Lock()
cleanupExpiredHTTPResponseOwnerBindings(s.responseOwners, now, openAIWSStateStoreCleanupMaxPerMap)
s.responseOwnerMu.Unlock()
s.responseToConnMu.Lock()
cleanupExpiredConnBindings(s.responseToConn, now, openAIWSStateStoreCleanupMaxPerMap)
s.responseToConnMu.Unlock()
@@ -419,6 +518,11 @@ func openAIWSResponseAccountCacheKey(responseID string) string {
return openAIWSResponseAccountCachePrefix + hex.EncodeToString(sum[:])
}
func openAIHTTPResponseOwnerCacheKey(prefix, responseID string) string {
sum := sha256.Sum256([]byte(responseID))
return prefix + hex.EncodeToString(sum[:])
}
// openAIWSResponseAccountMapKey 本地热缓存按分组隔离的 key,与 Redis 层保持一致,避免跨组命中。
func openAIWSResponseAccountMapKey(groupID int64, responseID string) string {
return fmt.Sprintf("%d:%s", groupID, responseID)
@@ -28,6 +28,27 @@ func TestOpenAIWSStateStore_BindGetDeleteResponseAccount(t *testing.T) {
require.Zero(t, accountID)
}
func TestOpenAIWSStateStore_HTTPResponseOwnerPersistsAcrossStoreInstances(t *testing.T) {
cache := &stubGatewayCache{}
ctx := context.Background()
groupID := int64(8)
writer := NewOpenAIWSStateStore(cache)
require.NoError(t, writer.BindHTTPResponseOwner(ctx, groupID, "resp_owned", 201, 301, time.Minute))
userID, apiKeyID, found, err := writer.GetHTTPResponseOwner(ctx, groupID, "resp_owned")
require.NoError(t, err)
require.True(t, found)
require.Equal(t, int64(201), userID)
require.Equal(t, int64(301), apiKeyID)
reader := NewOpenAIWSStateStore(cache)
userID, apiKeyID, found, err = reader.GetHTTPResponseOwner(ctx, groupID, "resp_owned")
require.NoError(t, err)
require.True(t, found)
require.Equal(t, int64(201), userID)
require.Equal(t, int64(301), apiKeyID)
}
func TestOpenAIWSStateStore_ResponseConnTTL(t *testing.T) {
store := NewOpenAIWSStateStore(nil)
store.BindResponseConn("resp_conn", "conn_1", 30*time.Millisecond)
@@ -270,30 +270,34 @@ func Relay(
if !options.StartClientAfterFirstDownstream {
startClientReader()
}
go runUpstreamToClient(
relayCtx,
upstreamConn,
writeClient,
startAt,
nowFn,
state,
options.OnUsageParseFailure,
options.OnTurnComplete,
options.BeforeWriteClient,
options.BeforeClientWrite,
options.AfterClientWrite,
func(msgType coderws.MessageType, payload []byte) {
if options.StartClientAfterFirstDownstream {
startClientReader()
}
},
&dropDownstreamWrites,
upstreamToClientFrames,
droppedDownstreamFrames,
markActivity,
onTrace,
exitCh,
)
upstreamDone := make(chan struct{})
go func() {
defer close(upstreamDone)
runUpstreamToClient(
relayCtx,
upstreamConn,
writeClient,
startAt,
nowFn,
state,
options.OnUsageParseFailure,
options.OnTurnComplete,
options.BeforeWriteClient,
options.BeforeClientWrite,
options.AfterClientWrite,
func(msgType coderws.MessageType, payload []byte) {
if options.StartClientAfterFirstDownstream {
startClientReader()
}
},
&dropDownstreamWrites,
upstreamToClientFrames,
droppedDownstreamFrames,
markActivity,
onTrace,
exitCh,
)
}()
go runIdleWatchdog(relayCtx, nowFn, options.IdleTimeout, &lastActivity, onTrace, exitCh)
firstExit := <-exitCh
@@ -347,6 +351,10 @@ func Relay(
relayCancel()
_ = upstreamConn.Close()
// ReadFrame observes relayCtx cancellation and Close is the transport-level
// fallback. Join the reader before touching relayState or firing the final
// turn callback; otherwise a late read can race Relay's result settlement.
<-upstreamDone
emitTurnComplete(options.OnTurnComplete, state, finalizePendingBareError(state, nowFn()))
enrichResult(&result, state, nowFn().Sub(startAt))
@@ -41,6 +41,16 @@ type closeSpyFrameConn struct {
closeCalls atomic.Int32
}
type cancelJoinProbeFrameConn struct {
readStarted chan struct{}
readCanceled chan struct{}
allowReturn chan struct{}
readReturned chan struct{}
startOnce sync.Once
cancelOnce sync.Once
returnOnce sync.Once
}
func newPassthroughTestFrameConn(frames []passthroughTestFrame, autoClose bool) *passthroughTestFrameConn {
c := &passthroughTestFrameConn{
readCh: make(chan passthroughTestFrame, len(frames)+1),
@@ -179,6 +189,38 @@ func (c *closeSpyFrameConn) CloseCalls() int32 {
return c.closeCalls.Load()
}
func newCancelJoinProbeFrameConn() *cancelJoinProbeFrameConn {
return &cancelJoinProbeFrameConn{
readStarted: make(chan struct{}),
readCanceled: make(chan struct{}),
allowReturn: make(chan struct{}),
readReturned: make(chan struct{}),
}
}
func (c *cancelJoinProbeFrameConn) ReadFrame(ctx context.Context) (coderws.MessageType, []byte, error) {
c.startOnce.Do(func() { close(c.readStarted) })
<-ctx.Done()
c.cancelOnce.Do(func() { close(c.readCanceled) })
<-c.allowReturn
c.returnOnce.Do(func() { close(c.readReturned) })
return coderws.MessageText, nil, ctx.Err()
}
func (c *cancelJoinProbeFrameConn) WriteFrame(ctx context.Context, _ coderws.MessageType, _ []byte) error {
if ctx == nil {
ctx = context.Background()
}
select {
case <-ctx.Done():
return ctx.Err()
default:
return nil
}
}
func (c *cancelJoinProbeFrameConn) Close() error { return nil }
func TestRelay_BasicRelayAndUsage(t *testing.T) {
t.Parallel()
@@ -372,6 +414,58 @@ func TestRelay_IdleTimeoutDoesNotCloseClientOnError(t *testing.T) {
require.GreaterOrEqual(t, upstreamConn.CloseCalls(), int32(1))
}
func TestRelay_JoinsUpstreamReaderBeforeReturning(t *testing.T) {
t.Parallel()
clientConn := &closeSpyFrameConn{}
upstreamConn := newCancelJoinProbeFrameConn()
ctx, cancel := context.WithCancel(context.Background())
resultCh := make(chan *RelayExit, 1)
go func() {
_, relayExit := Relay(
ctx,
clientConn,
upstreamConn,
[]byte(`{"type":"response.create","model":"gpt-4o","input":[]}`),
RelayOptions{},
)
resultCh <- relayExit
}()
select {
case <-upstreamConn.readStarted:
case <-time.After(time.Second):
t.Fatal("upstream reader did not start")
}
cancel()
select {
case <-upstreamConn.readCanceled:
case <-time.After(time.Second):
t.Fatal("upstream reader did not observe relay cancellation")
}
select {
case <-resultCh:
t.Fatal("Relay returned before the upstream reader exited")
case <-time.After(50 * time.Millisecond):
}
close(upstreamConn.allowReturn)
select {
case relayExit := <-resultCh:
require.NotNil(t, relayExit)
require.ErrorIs(t, relayExit.Err, context.Canceled)
case <-time.After(time.Second):
t.Fatal("Relay did not return after the upstream reader exited")
}
select {
case <-upstreamConn.readReturned:
default:
t.Fatal("Relay returned before the upstream reader completion signal")
}
require.Zero(t, clientConn.CloseCalls(), "错误路径不应提前关闭客户端连接")
}
func TestRelay_NilConnections(t *testing.T) {
t.Parallel()
@@ -388,7 +388,9 @@ type OpsUpstreamErrorEvent struct {
Detail string `json:"detail,omitempty"`
// SkipMonitoring is request-local rule state. It is intentionally excluded
// from persisted attempt JSON and only lets the final attempt control Ops.
// from persisted attempt JSON. The logger consults it only when this event is
// the final client-visible failure; recovered attempts remain provider-health
// telemetry and do not count as failed requests.
SkipMonitoring bool `json:"-"`
}
@@ -427,11 +429,10 @@ func appendOpsUpstreamError(c *gin.Context, ev OpsUpstreamErrorEvent) {
checkSkipMonitoringForUpstreamEvent(c, &evCopy)
}
// checkSkipMonitoringForUpstreamEvent checks whether the upstream error event
// matches a passthrough rule with skip_monitoring=true and, if so, sets the
// OpsSkipPassthroughKey on the context. This ensures intermediate retry /
// failover errors (which never go through the final applyErrorPassthroughRule
// path) can still suppress ops_error_logs recording.
// checkSkipMonitoringForUpstreamEvent snapshots whether this attempt matches a
// skip_monitoring passthrough rule. The final failure decides whether the
// request error is hidden; an intermediate recovered attempt cannot suppress a
// later client-visible failure.
func checkSkipMonitoringForUpstreamEvent(c *gin.Context, ev *OpsUpstreamErrorEvent) {
if ev.UpstreamStatusCode == 0 {
return
-1
View File
@@ -29,7 +29,6 @@ type TempUnschedCache interface {
// aggregate pool API-key failures across gateway instances.
type OpenAIAPIKeyHealthCache interface {
RecordOpenAIAPIKeyHealthFailure(ctx context.Context, accountID int64, windowMinutes, threshold int) (count int64, tripped bool, err error)
ResetOpenAIAPIKeyHealthFailures(ctx context.Context, accountID int64) error
}
// TimeoutCounterCache 超时计数器缓存接口