Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
112 changes: 98 additions & 14 deletions controller/relay.go
Original file line number Diff line number Diff line change
Expand Up @@ -15,36 +15,98 @@ import (
"one-api/middleware"
"one-api/model"
"one-api/relay"
relaycommon "one-api/relay/common"
"one-api/relay/constant"
relayconstant "one-api/relay/constant"
"one-api/relay/helper"
"one-api/service"
"strconv"
"strings"
"time"
)

func relayHandler(c *gin.Context, relayMode int) *dto.OpenAIErrorWithStatusCode {
func relayInfoHandler(c *gin.Context, relayMode int) (*relaycommon.RelayInfo, interface{}, string, *dto.OpenAIErrorWithStatusCode) {
switch relayMode {
case relayconstant.RelayModeImagesGenerations:
relayInfo, request, err := relay.ImageInfo(c)
if err != nil {
return nil, nil, "", err
}
return relayInfo, request, request.Model, nil
case relayconstant.RelayModeAudioSpeech:
fallthrough
case relayconstant.RelayModeAudioTranslation:
fallthrough
case relayconstant.RelayModeAudioTranscription:
relayInfo, request, err := relay.AudioInfo(c)
if err != nil {
return nil, nil, "", err
}
return relayInfo, request, request.Model, nil
case relayconstant.RelayModeRerank:
relayInfo, request, err := relay.EmbeddingInfo(c)
if err != nil {
return nil, nil, "", err
}
return relayInfo, request, request.Model, nil
case relayconstant.RelayModeEmbeddings:
relayInfo, request, err := relay.EmbeddingInfo(c)
if err != nil {
return nil, nil, "", err
}
return relayInfo, request, request.Model, nil
default:
relayInfo, request, err := relay.TextInfo(c)
if err != nil {
return nil, nil, "", err
}
return relayInfo, request, request.Model, nil
}
}

func relayExecuteHandler(c *gin.Context, relayMode int, relayInfo *relaycommon.RelayInfo, request interface{}) *dto.OpenAIErrorWithStatusCode {
var err *dto.OpenAIErrorWithStatusCode
switch relayMode {
case relayconstant.RelayModeImagesGenerations:
err = relay.ImageHelper(c)
imageRequest, ok := request.(*dto.ImageRequest)
if !ok {
return service.OpenAIErrorWrapperLocal(fmt.Errorf("failed assert request: %d", relayMode), "invalid_request_type", http.StatusInternalServerError)
}
err = relay.ImageHelper(c, relayInfo, imageRequest)
case relayconstant.RelayModeAudioSpeech:
fallthrough
case relayconstant.RelayModeAudioTranslation:
fallthrough
case relayconstant.RelayModeAudioTranscription:
err = relay.AudioHelper(c)
audioRequest, ok := request.(*dto.AudioRequest)
if !ok {
return service.OpenAIErrorWrapperLocal(fmt.Errorf("failed assert request: %d", relayMode), "invalid_request_type", http.StatusInternalServerError)
}
err = relay.AudioHelper(c, relayInfo, audioRequest)
case relayconstant.RelayModeRerank:
err = relay.RerankHelper(c, relayMode)
rerankRequest, ok := request.(*dto.RerankRequest)
if !ok {
return service.OpenAIErrorWrapperLocal(fmt.Errorf("failed assert request: %d", relayMode), "invalid_request_type", http.StatusInternalServerError)
}
err = relay.RerankHelper(c, relayInfo, rerankRequest)
case relayconstant.RelayModeEmbeddings:
err = relay.EmbeddingHelper(c)
embeddingRequest, ok := request.(*dto.EmbeddingRequest)
if !ok {
return service.OpenAIErrorWrapperLocal(fmt.Errorf("failed assert request: %d", relayMode), "invalid_request_type", http.StatusInternalServerError)
}
err = relay.EmbeddingHelper(c, relayInfo, embeddingRequest)
default:
err = relay.TextHelper(c)
textRequest, ok := request.(*dto.GeneralOpenAIRequest)
if !ok {
return service.OpenAIErrorWrapperLocal(fmt.Errorf("failed assert request: %d", relayMode), "invalid_request_type", http.StatusInternalServerError)
}
err = relay.TextHelper(c, relayInfo, textRequest)
}
return err
}

func Relay(c *gin.Context) {
startTime := time.Now()
relayMode := constant.Path2RelayMode(c.Request.URL.Path)
requestId := c.GetString(common.RequestIdKey)
group := c.GetString("group")
Expand All @@ -57,19 +119,34 @@ func Relay(c *gin.Context) {
openaiErr = service.OpenAIErrorWrapperLocal(err, "get_channel_failed", http.StatusInternalServerError)
break
}
if i > 0 {
metrics.IncrementRelayRetryCounter(strconv.Itoa(channel.Id), group, 1)
fillRelayRequest(c, channel)
var (
relayInfo *relaycommon.RelayInfo
request interface{}
requestModel string
)
relayInfo, request, requestModel, openaiErr = relayInfoHandler(c, relayMode)
if i == 0 {
// e2e 用户请求计数
metrics.IncrementRelayRequestE2ETotalCounter(strconv.Itoa(channel.Id), requestModel, group, 1)
} else {
// 重试计数
metrics.IncrementRelayRetryCounter(strconv.Itoa(channel.Id), requestModel, group, 1)
}

openaiErr = relayRequest(c, relayMode, channel)

if openaiErr == nil {
return // 成功处理请求,直接返回
openaiErr = executeRelayRequest(c, relayMode, relayInfo, request)
if openaiErr == nil {
metrics.IncrementRelayRequestE2ESuccessCounter(strconv.Itoa(channel.Id), requestModel, group, 1)
metrics.ObserveRelayRequestE2EDuration(strconv.Itoa(channel.Id), requestModel, group, time.Since(startTime).Seconds())
return
}
}

go processChannelError(c, channel.Id, channel.Type, channel.Name, channel.GetAutoBan(), openaiErr)

if !shouldRetry(c, openaiErr, common.RetryTimes-i) {
// e2e 失败计数
metrics.IncrementRelayRequestE2EFailedCounter(strconv.Itoa(channel.Id), requestModel, group, strconv.Itoa(openaiErr.StatusCode), 1)
break
}
}
Expand Down Expand Up @@ -152,11 +229,14 @@ func WssRelay(c *gin.Context) {
}
}

func relayRequest(c *gin.Context, relayMode int, channel *model.Channel) *dto.OpenAIErrorWithStatusCode {
func fillRelayRequest(c *gin.Context, channel *model.Channel) {
addUsedChannel(c, channel.Id)
requestBody, _ := common.GetRequestBody(c)
c.Request.Body = io.NopCloser(bytes.NewBuffer(requestBody))
return relayHandler(c, relayMode)
}

func executeRelayRequest(c *gin.Context, relayMode int, relayInfo *relaycommon.RelayInfo, request interface{}) *dto.OpenAIErrorWithStatusCode {
return relayExecuteHandler(c, relayMode, relayInfo, request)
}

func wssRequest(c *gin.Context, ws *websocket.Conn, relayMode int, channel *model.Channel) *dto.OpenAIErrorWithStatusCode {
Expand Down Expand Up @@ -234,6 +314,10 @@ func shouldRetry(c *gin.Context, openaiErr *dto.OpenAIErrorWithStatusCode, retry
if openaiErr.StatusCode/100 == 2 {
return false
}
if strings.Contains(openaiErr.Error.Message, "deadline exceeded") || strings.Contains(openaiErr.Error.Message, "request canceled") {
common.LogInfo(c, "客户端请求下游超时,不再重试")
return false
}
return true
}

Expand Down
61 changes: 55 additions & 6 deletions metrics/metrics.go
Original file line number Diff line number Diff line change
Expand Up @@ -10,11 +10,17 @@ const (
)

func RegisterMetrics(registry prometheus.Registerer) {
// channel
registry.MustRegister(relayRequestTotalCounter)
registry.MustRegister(relayRequestSuccessCounter)
registry.MustRegister(relayRequestFailedCounter)
registry.MustRegister(relayRequestRetryCounter)
registry.MustRegister(relayRequestDurationObsever)
// e2e
registry.MustRegister(relayRequestE2ETotalCounter)
registry.MustRegister(relayRequestE2ESuccessCounter)
registry.MustRegister(relayRequestE2EFailedCounter)
registry.MustRegister(relayRequestE2EDurationObsever)
}

var (
Expand All @@ -35,13 +41,13 @@ var (
Subsystem: Namespace,
Name: "relay_request_failed",
Help: "Total number of relay request failed",
}, []string{"channel", "model", "group", "code", "msg"})
}, []string{"channel", "model", "group", "code"})
relayRequestRetryCounter = prometheus.NewCounterVec(
prometheus.CounterOpts{
Subsystem: Namespace,
Name: "relay_request_retry",
Help: "Total number of relay request retry",
}, []string{"channel", "group"})
}, []string{"channel", "model", "group"})
relayRequestDurationObsever = promauto.NewHistogramVec(
prometheus.HistogramOpts{
Subsystem: Namespace,
Expand All @@ -51,6 +57,33 @@ var (
},
[]string{"channel", "model", "group"},
)
relayRequestE2ETotalCounter = prometheus.NewCounterVec(
prometheus.CounterOpts{
Subsystem: Namespace,
Name: "relay_request_e2e_total",
Help: "Total number of relay request e2e total",
}, []string{"channel", "model", "group"})
relayRequestE2ESuccessCounter = prometheus.NewCounterVec(
prometheus.CounterOpts{
Subsystem: Namespace,
Name: "relay_request_e2e_success",
Help: "Total number of relay request e2e success",
}, []string{"channel", "model", "group"})
relayRequestE2EFailedCounter = prometheus.NewCounterVec(
prometheus.CounterOpts{
Subsystem: Namespace,
Name: "relay_request_e2e_failed",
Help: "Total number of relay request e2e failed",
}, []string{"channel", "model", "group", "code"})
relayRequestE2EDurationObsever = promauto.NewHistogramVec(
prometheus.HistogramOpts{
Subsystem: Namespace,
Name: "relay_request_e2e_duration",
Help: "Duration of relay request e2e",
Buckets: prometheus.ExponentialBuckets(1, 2, 12),
},
[]string{"channel", "model", "group"},
)
)

func IncrementRelayRequestTotalCounter(channel, model, group string, add float64) {
Expand All @@ -61,14 +94,30 @@ func IncrementRelayRequestSuccessCounter(channel, model, group string, add float
relayRequestSuccessCounter.WithLabelValues(channel, model, group).Add(add)
}

func IncrementRelayRequestFailedCounter(channel, model, group, code, msg string, add float64) {
relayRequestFailedCounter.WithLabelValues(channel, model, group, code, msg).Add(add)
func IncrementRelayRequestFailedCounter(channel, model, group, code string, add float64) {
relayRequestFailedCounter.WithLabelValues(channel, model, group, code).Add(add)
}

func IncrementRelayRetryCounter(channel, group string, add float64) {
relayRequestRetryCounter.WithLabelValues(channel, group).Add(add)
func IncrementRelayRetryCounter(channel, model, group string, add float64) {
relayRequestRetryCounter.WithLabelValues(channel, model, group).Add(add)
}

func ObserveRelayRequestDuration(channel, model, group string, duration float64) {
relayRequestDurationObsever.WithLabelValues(channel, model, group).Observe(duration)
}

func IncrementRelayRequestE2ETotalCounter(channel, model, group string, add float64) {
relayRequestE2ETotalCounter.WithLabelValues(channel, model, group).Add(add)
}

func IncrementRelayRequestE2ESuccessCounter(channel, model, group string, add float64) {
relayRequestE2ESuccessCounter.WithLabelValues(channel, model, group).Add(add)
}

func IncrementRelayRequestE2EFailedCounter(channel, model, group, code string, add float64) {
relayRequestE2EFailedCounter.WithLabelValues(channel, model, group, code).Add(add)
}

func ObserveRelayRequestE2EDuration(channel, model, group string, duration float64) {
relayRequestE2EDurationObsever.WithLabelValues(channel, model, group).Observe(duration)
}
20 changes: 14 additions & 6 deletions relay/relay-audio.go
Original file line number Diff line number Diff line change
Expand Up @@ -57,25 +57,33 @@ func getAndValidAudioRequest(c *gin.Context, info *relaycommon.RelayInfo) (*dto.
return audioRequest, nil
}

func AudioHelper(c *gin.Context) (openaiErr *dto.OpenAIErrorWithStatusCode) {
startTime := time.Now()
func AudioInfo(c *gin.Context) (*relaycommon.RelayInfo, *dto.AudioRequest, *dto.OpenAIErrorWithStatusCode) {
relayInfo := relaycommon.GenRelayInfo(c)
audioRequest, err := getAndValidAudioRequest(c, relayInfo)
if err != nil {
common.LogError(c, fmt.Sprintf("getAndValidAudioRequest failed: %s", err.Error()))
return service.OpenAIErrorWrapper(err, "invalid_audio_request", http.StatusBadRequest)
return nil, nil, service.OpenAIErrorWrapper(err, "invalid_audio_request", http.StatusBadRequest)
}

return relayInfo, audioRequest, nil
}

func AudioHelper(c *gin.Context, relayInfo *relaycommon.RelayInfo, audioRequest *dto.AudioRequest) (openaiErr *dto.OpenAIErrorWithStatusCode) {
startTime := time.Now()
var funcErr *dto.OpenAIErrorWithStatusCode
metrics.IncrementRelayRequestTotalCounter(strconv.Itoa(relayInfo.ChannelId), audioRequest.Model, relayInfo.Group, 1)
defer func() {
if openaiErr != nil {
metrics.IncrementRelayRequestFailedCounter(strconv.Itoa(relayInfo.ChannelId), audioRequest.Model, relayInfo.Group, strconv.Itoa(openaiErr.StatusCode), openaiErr.Error.Message, 1)
if funcErr != nil {
metrics.IncrementRelayRequestFailedCounter(strconv.Itoa(relayInfo.ChannelId), audioRequest.Model, relayInfo.Group, strconv.Itoa(funcErr.StatusCode), 1)
} else {
metrics.IncrementRelayRequestSuccessCounter(strconv.Itoa(relayInfo.ChannelId), audioRequest.Model, relayInfo.Group, 1)
metrics.ObserveRelayRequestDuration(strconv.Itoa(relayInfo.ChannelId), audioRequest.Model, relayInfo.Group, time.Since(startTime).Seconds())
}
}()
promptTokens := 0
var (
err error
promptTokens = 0
)
preConsumedTokens := common.PreConsumedQuota
if relayInfo.RelayMode == relayconstant.RelayModeAudioSpeech {
promptTokens, err = service.CountTTSToken(audioRequest.Input, audioRequest.Model)
Expand Down
16 changes: 10 additions & 6 deletions relay/relay-image.go
Original file line number Diff line number Diff line change
Expand Up @@ -73,26 +73,30 @@ func getAndValidImageRequest(c *gin.Context, info *relaycommon.RelayInfo) (*dto.
return imageRequest, nil
}

func ImageHelper(c *gin.Context) (openaiErr *dto.OpenAIErrorWithStatusCode) {
startTime := time.Now()
func ImageInfo(c *gin.Context) (*relaycommon.RelayInfo, *dto.ImageRequest, *dto.OpenAIErrorWithStatusCode) {
relayInfo := relaycommon.GenRelayInfo(c)

imageRequest, err := getAndValidImageRequest(c, relayInfo)
if err != nil {
common.LogError(c, fmt.Sprintf("getAndValidImageRequest failed: %s", err.Error()))
return service.OpenAIErrorWrapper(err, "invalid_image_request", http.StatusBadRequest)
return nil, nil, service.OpenAIErrorWrapper(err, "invalid_image_request", http.StatusBadRequest)
}
return relayInfo, imageRequest, nil
}

func ImageHelper(c *gin.Context, relayInfo *relaycommon.RelayInfo, imageRequest *dto.ImageRequest) (openaiErr *dto.OpenAIErrorWithStatusCode) {
startTime := time.Now()
var funcErr *dto.OpenAIErrorWithStatusCode
metrics.IncrementRelayRequestTotalCounter(strconv.Itoa(relayInfo.ChannelId), imageRequest.Model, relayInfo.Group, 1)
defer func() {
if openaiErr != nil {
metrics.IncrementRelayRequestFailedCounter(strconv.Itoa(relayInfo.ChannelId), imageRequest.Model, relayInfo.Group, strconv.Itoa(openaiErr.StatusCode), openaiErr.Error.Message, 1)
if funcErr != nil {
metrics.IncrementRelayRequestFailedCounter(strconv.Itoa(relayInfo.ChannelId), imageRequest.Model, relayInfo.Group, strconv.Itoa(funcErr.StatusCode), 1)
} else {
metrics.IncrementRelayRequestSuccessCounter(strconv.Itoa(relayInfo.ChannelId), imageRequest.Model, relayInfo.Group, 1)
metrics.ObserveRelayRequestDuration(strconv.Itoa(relayInfo.ChannelId), imageRequest.Model, relayInfo.Group, time.Since(startTime).Seconds())
}
}()
err = helper.ModelMappedHelper(c, relayInfo)
err := helper.ModelMappedHelper(c, relayInfo)
if err != nil {
funcErr = service.OpenAIErrorWrapperLocal(err, "model_mapped_error", http.StatusInternalServerError)
return funcErr
Expand Down
14 changes: 9 additions & 5 deletions relay/relay-text.go
Original file line number Diff line number Diff line change
Expand Up @@ -68,21 +68,25 @@ func getAndValidateTextRequest(c *gin.Context, relayInfo *relaycommon.RelayInfo)
return textRequest, nil
}

func TextHelper(c *gin.Context) (openaiErr *dto.OpenAIErrorWithStatusCode) {
startTime := time.Now()
func TextInfo(c *gin.Context) (*relaycommon.RelayInfo, *dto.GeneralOpenAIRequest, *dto.OpenAIErrorWithStatusCode) {
relayInfo := relaycommon.GenRelayInfo(c)

// get & validate textRequest 获取并验证文本请求
textRequest, err := getAndValidateTextRequest(c, relayInfo)
if err != nil {
common.LogError(c, fmt.Sprintf("getAndValidateTextRequest failed: %s", err.Error()))
return service.OpenAIErrorWrapperLocal(err, "invalid_text_request", http.StatusBadRequest)
return nil, nil, service.OpenAIErrorWrapperLocal(err, "invalid_text_request", http.StatusBadRequest)
}
return relayInfo, textRequest, nil
}

func TextHelper(c *gin.Context, relayInfo *relaycommon.RelayInfo, textRequest *dto.GeneralOpenAIRequest) (openaiErr *dto.OpenAIErrorWithStatusCode) {
startTime := time.Now()
var funcErr *dto.OpenAIErrorWithStatusCode
metrics.IncrementRelayRequestTotalCounter(strconv.Itoa(relayInfo.ChannelId), textRequest.Model, relayInfo.Group, 1)
defer func() {
if funcErr != nil {
metrics.IncrementRelayRequestFailedCounter(strconv.Itoa(relayInfo.ChannelId), textRequest.Model, relayInfo.Group, strconv.Itoa(openaiErr.StatusCode), funcErr.Error.Message, 1)
metrics.IncrementRelayRequestFailedCounter(strconv.Itoa(relayInfo.ChannelId), textRequest.Model, relayInfo.Group, strconv.Itoa(openaiErr.StatusCode), 1)
} else {
metrics.IncrementRelayRequestSuccessCounter(strconv.Itoa(relayInfo.ChannelId), textRequest.Model, relayInfo.Group, 1)
metrics.ObserveRelayRequestDuration(strconv.Itoa(relayInfo.ChannelId), textRequest.Model, relayInfo.Group, time.Since(startTime).Seconds())
Expand All @@ -98,7 +102,7 @@ func TextHelper(c *gin.Context) (openaiErr *dto.OpenAIErrorWithStatusCode) {
}
}

err = helper.ModelMappedHelper(c, relayInfo)
err := helper.ModelMappedHelper(c, relayInfo)
if err != nil {
funcErr = service.OpenAIErrorWrapperLocal(err, "model_mapped_error", http.StatusInternalServerError)
return funcErr
Expand Down
Loading