mirror of
https://github.com/Wei-Shaw/sub2api.git
synced 2026-10-07 11:43:23 +08:00
修复 PR 5888 审查发现的兼容性与竞态问题
This commit is contained in:
@@ -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)
|
||||
}
|
||||
|
||||
@@ -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) {
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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 超时计数器缓存接口
|
||||
|
||||
Reference in New Issue
Block a user