Files
new-api/controller/channel_pin_retry_test.go
T

96 lines
3.0 KiB
Go

package controller
import (
"errors"
"net/http"
"net/http/httptest"
"testing"
"github.com/QuantumNous/new-api/dto"
"github.com/QuantumNous/new-api/service"
"github.com/QuantumNous/new-api/relaykit/types"
"github.com/gin-gonic/gin"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)
func TestShouldRetryHonorsPinRetryMode(t *testing.T) {
openaiErr := types.NewOpenAIError(errors.New("upstream"), types.ErrorCodeBadResponseStatusCode, http.StatusInternalServerError)
c := newPinRetryContext()
assert.True(t, shouldRetry(c, openaiErr, 1))
origin := newPinRetryContext()
service.GetChannelConstraints(origin).AddPin(dto.ChannelPin{
ChannelId: 2,
Source: dto.PinSourceOriginTask,
Rank: dto.PinRankOriginTask,
RetryMode: dto.PinRetrySameChannel,
})
assert.True(t, shouldRetry(origin, openaiErr, 1), "origin pin retries on the same channel")
token := newPinRetryContext()
service.GetChannelConstraints(token).AddPin(dto.ChannelPin{
ChannelId: 1,
Source: dto.PinSourceToken,
Rank: dto.PinRankToken,
RetryMode: dto.PinRetrySingleAttempt,
})
assert.False(t, shouldRetry(token, openaiErr, 1), "token pin suppresses retry")
}
func TestShouldRetryTaskRelayHonorsPinRetryMode(t *testing.T) {
taskErr := &dto.TaskError{StatusCode: http.StatusInternalServerError}
c := newPinRetryContext()
assert.True(t, shouldRetryTaskRelay(c, 1, taskErr, 1))
origin := newPinRetryContext()
service.GetChannelConstraints(origin).AddPin(dto.ChannelPin{
ChannelId: 2,
Source: dto.PinSourceOriginTask,
Rank: dto.PinRankOriginTask,
RetryMode: dto.PinRetrySameChannel,
})
assert.True(t, shouldRetryTaskRelay(origin, 2, taskErr, 1))
token := newPinRetryContext()
service.GetChannelConstraints(token).AddPin(dto.ChannelPin{
ChannelId: 1,
Source: dto.PinSourceToken,
Rank: dto.PinRankToken,
RetryMode: dto.PinRetrySingleAttempt,
})
assert.False(t, shouldRetryTaskRelay(token, 1, taskErr, 1))
}
func TestSameChannelPinsMergeToStricterRetryMode(t *testing.T) {
c := newPinRetryContext()
constraints := service.GetChannelConstraints(c)
constraints.AddPin(dto.ChannelPin{
ChannelId: 7,
Source: dto.PinSourceOriginTask,
Rank: dto.PinRankOriginTask,
RetryMode: dto.PinRetrySameChannel,
})
constraints.AddPin(dto.ChannelPin{
ChannelId: 7,
Source: dto.PinSourceToken,
Rank: dto.PinRankToken,
RetryMode: dto.PinRetrySingleAttempt,
})
pin, found, overridden := constraints.ResolvedPin()
require.True(t, found)
assert.Equal(t, 7, pin.ChannelId)
assert.Equal(t, dto.PinRetrySingleAttempt, pin.RetryMode)
assert.Empty(t, overridden)
assert.False(t, shouldRetry(c, types.NewOpenAIError(errors.New("upstream"), types.ErrorCodeBadResponseStatusCode, http.StatusInternalServerError), 1))
}
func newPinRetryContext() *gin.Context {
recorder := httptest.NewRecorder()
c, _ := gin.CreateTestContext(recorder)
c.Request = httptest.NewRequest(http.MethodPost, "/v1/chat/completions", nil)
return c
}