mirror of
https://github.com/QuantumNous/new-api.git
synced 2026-09-11 06:30:21 +00:00
96 lines
3.0 KiB
Go
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
|
|
}
|