Commit fc4d5933 by t0ng7u

Merge remote-tracking branch 'origin/alpha' into alpha

parents bee94700 48a6123c
package common
import (
"fmt"
"github.com/antlabs/pcopy"
)
func DeepCopy[T any](src *T) (*T, error) {
if src == nil {
return nil, fmt.Errorf("copy source cannot be nil")
}
var dst T
err := pcopy.Copy(&dst, src)
if err != nil {
return nil, err
}
if &dst == nil {
return nil, fmt.Errorf("copy result cannot be nil")
}
return &dst, nil
}
...@@ -2,12 +2,13 @@ package common ...@@ -2,12 +2,13 @@ package common
import ( import (
"bytes" "bytes"
"github.com/gin-gonic/gin"
"io" "io"
"net/http" "net/http"
"one-api/constant" "one-api/constant"
"strings" "strings"
"time" "time"
"github.com/gin-gonic/gin"
) )
const KeyRequestBody = "key_request_body" const KeyRequestBody = "key_request_body"
......
...@@ -5,6 +5,7 @@ import ( ...@@ -5,6 +5,7 @@ import (
"one-api/common" "one-api/common"
"one-api/model" "one-api/model"
"strconv" "strconv"
"strings"
"github.com/gin-gonic/gin" "github.com/gin-gonic/gin"
) )
...@@ -82,6 +83,57 @@ func GetTokenStatus(c *gin.Context) { ...@@ -82,6 +83,57 @@ func GetTokenStatus(c *gin.Context) {
}) })
} }
func GetTokenUsage(c *gin.Context) {
authHeader := c.GetHeader("Authorization")
if authHeader == "" {
c.JSON(http.StatusUnauthorized, gin.H{
"success": false,
"message": "No Authorization header",
})
return
}
parts := strings.Split(authHeader, " ")
if len(parts) != 2 || strings.ToLower(parts[0]) != "bearer" {
c.JSON(http.StatusUnauthorized, gin.H{
"success": false,
"message": "Invalid Bearer token",
})
return
}
tokenKey := parts[1]
token, err := model.GetTokenByKey(strings.TrimPrefix(tokenKey, "sk-"), false)
if err != nil {
c.JSON(http.StatusOK, gin.H{
"success": false,
"message": err.Error(),
})
return
}
expiredAt := token.ExpiredTime
if expiredAt == -1 {
expiredAt = 0
}
c.JSON(http.StatusOK, gin.H{
"code": true,
"message": "ok",
"data": gin.H{
"object": "token_usage",
"name": token.Name,
"total_granted": token.RemainQuota + token.UsedQuota,
"total_used": token.UsedQuota,
"total_available": token.RemainQuota,
"unlimited_quota": token.UnlimitedQuota,
"model_limits": token.GetModelLimitsMap(),
"model_limits_enabled": token.ModelLimitsEnabled,
"expires_at": expiredAt,
},
})
}
func AddToken(c *gin.Context) { func AddToken(c *gin.Context) {
token := model.Token{} token := model.Token{}
err := c.ShouldBindJSON(&token) err := c.ShouldBindJSON(&token)
......
...@@ -26,6 +26,12 @@ func (r *AudioRequest) IsStream(c *gin.Context) bool { ...@@ -26,6 +26,12 @@ func (r *AudioRequest) IsStream(c *gin.Context) bool {
return false return false
} }
func (r *AudioRequest) SetModelName(modelName string) {
if modelName != "" {
r.Model = modelName
}
}
type AudioResponse struct { type AudioResponse struct {
Text string `json:"text"` Text string `json:"text"`
} }
......
...@@ -321,8 +321,14 @@ func (c *ClaudeRequest) GetTokenCountMeta() *types.TokenCountMeta { ...@@ -321,8 +321,14 @@ func (c *ClaudeRequest) GetTokenCountMeta() *types.TokenCountMeta {
return &tokenCountMeta return &tokenCountMeta
} }
func (claudeRequest *ClaudeRequest) IsStream(c *gin.Context) bool { func (c *ClaudeRequest) IsStream(ctx *gin.Context) bool {
return claudeRequest.Stream return c.Stream
}
func (c *ClaudeRequest) SetModelName(modelName string) {
if modelName != "" {
c.Model = modelName
}
} }
func (c *ClaudeRequest) SearchToolNameByToolCallId(toolCallId string) string { func (c *ClaudeRequest) SearchToolNameByToolCallId(toolCallId string) string {
......
...@@ -48,6 +48,12 @@ func (r *EmbeddingRequest) IsStream(c *gin.Context) bool { ...@@ -48,6 +48,12 @@ func (r *EmbeddingRequest) IsStream(c *gin.Context) bool {
return false return false
} }
func (r *EmbeddingRequest) SetModelName(modelName string) {
if modelName != "" {
r.Model = modelName
}
}
func (r *EmbeddingRequest) ParseInput() []string { func (r *EmbeddingRequest) ParseInput() []string {
if r.Input == nil { if r.Input == nil {
return make([]string, 0) return make([]string, 0)
......
...@@ -73,6 +73,10 @@ func (r *GeminiChatRequest) IsStream(c *gin.Context) bool { ...@@ -73,6 +73,10 @@ func (r *GeminiChatRequest) IsStream(c *gin.Context) bool {
return false return false
} }
func (r *GeminiChatRequest) SetModelName(modelName string) {
// GeminiChatRequest does not have a model field, so this method does nothing.
}
func (r *GeminiChatRequest) GetTools() []GeminiChatTool { func (r *GeminiChatRequest) GetTools() []GeminiChatTool {
var tools []GeminiChatTool var tools []GeminiChatTool
if strings.HasSuffix(string(r.Tools), "[") { if strings.HasSuffix(string(r.Tools), "[") {
...@@ -312,10 +316,61 @@ type GeminiEmbeddingRequest struct { ...@@ -312,10 +316,61 @@ type GeminiEmbeddingRequest struct {
OutputDimensionality int `json:"outputDimensionality,omitempty"` OutputDimensionality int `json:"outputDimensionality,omitempty"`
} }
func (r *GeminiEmbeddingRequest) IsStream(c *gin.Context) bool {
// Gemini embedding requests are not streamed
return false
}
func (r *GeminiEmbeddingRequest) GetTokenCountMeta() *types.TokenCountMeta {
var inputTexts []string
for _, part := range r.Content.Parts {
if part.Text != "" {
inputTexts = append(inputTexts, part.Text)
}
}
inputText := strings.Join(inputTexts, "\n")
return &types.TokenCountMeta{
CombineText: inputText,
}
}
func (r *GeminiEmbeddingRequest) SetModelName(modelName string) {
if modelName != "" {
r.Model = modelName
}
}
type GeminiBatchEmbeddingRequest struct { type GeminiBatchEmbeddingRequest struct {
Requests []*GeminiEmbeddingRequest `json:"requests"` Requests []*GeminiEmbeddingRequest `json:"requests"`
} }
func (r *GeminiBatchEmbeddingRequest) IsStream(c *gin.Context) bool {
// Gemini batch embedding requests are not streamed
return false
}
func (r *GeminiBatchEmbeddingRequest) GetTokenCountMeta() *types.TokenCountMeta {
var inputTexts []string
for _, request := range r.Requests {
meta := request.GetTokenCountMeta()
if meta != nil && meta.CombineText != "" {
inputTexts = append(inputTexts, meta.CombineText)
}
}
inputText := strings.Join(inputTexts, "\n")
return &types.TokenCountMeta{
CombineText: inputText,
}
}
func (r *GeminiBatchEmbeddingRequest) SetModelName(modelName string) {
if modelName != "" {
for _, req := range r.Requests {
req.SetModelName(modelName)
}
}
}
type GeminiEmbeddingResponse struct { type GeminiEmbeddingResponse struct {
Embedding ContentEmbedding `json:"embedding"` Embedding ContentEmbedding `json:"embedding"`
} }
......
...@@ -12,10 +12,10 @@ type ImageRequest struct { ...@@ -12,10 +12,10 @@ type ImageRequest struct {
Model string `json:"model"` Model string `json:"model"`
Prompt string `json:"prompt" binding:"required"` Prompt string `json:"prompt" binding:"required"`
N uint `json:"n,omitempty"` N uint `json:"n,omitempty"`
Size string `json:"size,omitempty"` Size string `json:"size,omitempty"`
Quality string `json:"quality,omitempty"` Quality string `json:"quality,omitempty"`
ResponseFormat string `json:"response_format,omitempty"` ResponseFormat string `json:"response_format,omitempty"`
Style json.RawMessage `json:"style,omitempty"` Style json.RawMessage `json:"style,omitempty"`
User json.RawMessage `json:"user,omitempty"` User json.RawMessage `json:"user,omitempty"`
ExtraFields json.RawMessage `json:"extra_fields,omitempty"` ExtraFields json.RawMessage `json:"extra_fields,omitempty"`
Background json.RawMessage `json:"background,omitempty"` Background json.RawMessage `json:"background,omitempty"`
...@@ -63,6 +63,12 @@ func (i *ImageRequest) IsStream(c *gin.Context) bool { ...@@ -63,6 +63,12 @@ func (i *ImageRequest) IsStream(c *gin.Context) bool {
return false return false
} }
func (i *ImageRequest) SetModelName(modelName string) {
if modelName != "" {
i.Model = modelName
}
}
type ImageResponse struct { type ImageResponse struct {
Data []ImageData `json:"data"` Data []ImageData `json:"data"`
Created int64 `json:"created"` Created int64 `json:"created"`
......
...@@ -183,6 +183,12 @@ func (r *GeneralOpenAIRequest) IsStream(c *gin.Context) bool { ...@@ -183,6 +183,12 @@ func (r *GeneralOpenAIRequest) IsStream(c *gin.Context) bool {
return r.Stream return r.Stream
} }
func (r *GeneralOpenAIRequest) SetModelName(modelName string) {
if modelName != "" {
r.Model = modelName
}
}
func (r *GeneralOpenAIRequest) ToMap() map[string]any { func (r *GeneralOpenAIRequest) ToMap() map[string]any {
result := make(map[string]any) result := make(map[string]any)
data, _ := common.Marshal(r) data, _ := common.Marshal(r)
...@@ -841,6 +847,12 @@ func (r *OpenAIResponsesRequest) IsStream(c *gin.Context) bool { ...@@ -841,6 +847,12 @@ func (r *OpenAIResponsesRequest) IsStream(c *gin.Context) bool {
return r.Stream return r.Stream
} }
func (r *OpenAIResponsesRequest) SetModelName(modelName string) {
if modelName != "" {
r.Model = modelName
}
}
type Reasoning struct { type Reasoning struct {
Effort string `json:"effort,omitempty"` Effort string `json:"effort,omitempty"`
Summary string `json:"summary,omitempty"` Summary string `json:"summary,omitempty"`
......
...@@ -8,6 +8,7 @@ import ( ...@@ -8,6 +8,7 @@ import (
type Request interface { type Request interface {
GetTokenCountMeta() *types.TokenCountMeta GetTokenCountMeta() *types.TokenCountMeta
IsStream(c *gin.Context) bool IsStream(c *gin.Context) bool
SetModelName(modelName string)
} }
type BaseRequest struct { type BaseRequest struct {
...@@ -18,7 +19,7 @@ func (b *BaseRequest) GetTokenCountMeta() *types.TokenCountMeta { ...@@ -18,7 +19,7 @@ func (b *BaseRequest) GetTokenCountMeta() *types.TokenCountMeta {
TokenType: types.TokenTypeTokenizer, TokenType: types.TokenTypeTokenizer,
} }
} }
func (b *BaseRequest) IsStream(c *gin.Context) bool { func (b *BaseRequest) IsStream(c *gin.Context) bool {
return false return false
} }
func (b *BaseRequest) SetModelName(modelName string) {}
...@@ -37,6 +37,12 @@ func (r *RerankRequest) GetTokenCountMeta() *types.TokenCountMeta { ...@@ -37,6 +37,12 @@ func (r *RerankRequest) GetTokenCountMeta() *types.TokenCountMeta {
} }
} }
func (r *RerankRequest) SetModelName(modelName string) {
if modelName != "" {
r.Model = modelName
}
}
func (r *RerankRequest) GetReturnDocuments() bool { func (r *RerankRequest) GetReturnDocuments() bool {
if r.ReturnDocuments == nil { if r.ReturnDocuments == nil {
return false return false
......
...@@ -44,7 +44,11 @@ require ( ...@@ -44,7 +44,11 @@ require (
) )
require ( require (
github.com/Masterminds/goutils v1.1.1 // indirect
github.com/Masterminds/semver/v3 v3.2.0 // indirect
github.com/Masterminds/sprig/v3 v3.2.3 // indirect
github.com/anknown/darts v0.0.0-20151216065714-83ff685239e6 // indirect github.com/anknown/darts v0.0.0-20151216065714-83ff685239e6 // indirect
github.com/antlabs/pcopy v0.1.5 // indirect
github.com/aws/aws-sdk-go-v2/aws/protocol/eventstream v1.7.0 // indirect github.com/aws/aws-sdk-go-v2/aws/protocol/eventstream v1.7.0 // indirect
github.com/aws/aws-sdk-go-v2/internal/configsources v1.4.2 // indirect github.com/aws/aws-sdk-go-v2/internal/configsources v1.4.2 // indirect
github.com/aws/aws-sdk-go-v2/internal/endpoints/v2 v2.7.2 // indirect github.com/aws/aws-sdk-go-v2/internal/endpoints/v2 v2.7.2 // indirect
...@@ -69,6 +73,8 @@ require ( ...@@ -69,6 +73,8 @@ require (
github.com/gorilla/context v1.1.1 // indirect github.com/gorilla/context v1.1.1 // indirect
github.com/gorilla/securecookie v1.1.1 // indirect github.com/gorilla/securecookie v1.1.1 // indirect
github.com/gorilla/sessions v1.2.1 // indirect github.com/gorilla/sessions v1.2.1 // indirect
github.com/huandu/xstrings v1.3.3 // indirect
github.com/imdario/mergo v0.3.11 // indirect
github.com/jackc/pgpassfile v1.0.0 // indirect github.com/jackc/pgpassfile v1.0.0 // indirect
github.com/jackc/pgservicefile v0.0.0-20240606120523-5a60cdf6a761 // indirect github.com/jackc/pgservicefile v0.0.0-20240606120523-5a60cdf6a761 // indirect
github.com/jackc/pgx/v5 v5.7.1 // indirect github.com/jackc/pgx/v5 v5.7.1 // indirect
...@@ -79,11 +85,14 @@ require ( ...@@ -79,11 +85,14 @@ require (
github.com/klauspost/cpuid/v2 v2.2.9 // indirect github.com/klauspost/cpuid/v2 v2.2.9 // indirect
github.com/leodido/go-urn v1.4.0 // indirect github.com/leodido/go-urn v1.4.0 // indirect
github.com/mattn/go-isatty v0.0.20 // indirect github.com/mattn/go-isatty v0.0.20 // indirect
github.com/mitchellh/copystructure v1.0.0 // indirect
github.com/mitchellh/mapstructure v1.5.0 // indirect github.com/mitchellh/mapstructure v1.5.0 // indirect
github.com/mitchellh/reflectwalk v1.0.0 // indirect
github.com/modern-go/concurrent v0.0.0-20180306012644-bacd9c7ef1dd // indirect github.com/modern-go/concurrent v0.0.0-20180306012644-bacd9c7ef1dd // indirect
github.com/modern-go/reflect2 v1.0.2 // indirect github.com/modern-go/reflect2 v1.0.2 // indirect
github.com/pelletier/go-toml/v2 v2.2.1 // indirect github.com/pelletier/go-toml/v2 v2.2.1 // indirect
github.com/remyoudompheng/bigfft v0.0.0-20230129092748-24d4a6f8daec // indirect github.com/remyoudompheng/bigfft v0.0.0-20230129092748-24d4a6f8daec // indirect
github.com/spf13/cast v1.3.1 // indirect
github.com/tidwall/match v1.1.1 // indirect github.com/tidwall/match v1.1.1 // indirect
github.com/tidwall/pretty v1.2.0 // indirect github.com/tidwall/pretty v1.2.0 // indirect
github.com/tklauser/go-sysconf v0.3.12 // indirect github.com/tklauser/go-sysconf v0.3.12 // indirect
......
...@@ -4,6 +4,7 @@ import ( ...@@ -4,6 +4,7 @@ import (
"errors" "errors"
"fmt" "fmt"
"net/http" "net/http"
"one-api/common"
"one-api/dto" "one-api/dto"
relaycommon "one-api/relay/common" relaycommon "one-api/relay/common"
"one-api/relay/helper" "one-api/relay/helper"
...@@ -16,12 +17,17 @@ import ( ...@@ -16,12 +17,17 @@ import (
func AudioHelper(c *gin.Context, info *relaycommon.RelayInfo) (newAPIError *types.NewAPIError) { func AudioHelper(c *gin.Context, info *relaycommon.RelayInfo) (newAPIError *types.NewAPIError) {
info.InitChannelMeta(c) info.InitChannelMeta(c)
audioRequest, ok := info.Request.(*dto.AudioRequest) audioReq, ok := info.Request.(*dto.AudioRequest)
if !ok { if !ok {
return types.NewError(errors.New("invalid request type"), types.ErrorCodeInvalidRequest, types.ErrOptionWithSkipRetry()) return types.NewError(errors.New("invalid request type"), types.ErrorCodeInvalidRequest, types.ErrOptionWithSkipRetry())
} }
err := helper.ModelMappedHelper(c, info, audioRequest) request, err := common.DeepCopy(audioReq)
if err != nil {
return types.NewError(fmt.Errorf("failed to copy request to AudioRequest: %w", err), types.ErrorCodeInvalidRequest, types.ErrOptionWithSkipRetry())
}
err = helper.ModelMappedHelper(c, info, request)
if err != nil { if err != nil {
return types.NewError(err, types.ErrorCodeChannelModelMappedError, types.ErrOptionWithSkipRetry()) return types.NewError(err, types.ErrorCodeChannelModelMappedError, types.ErrOptionWithSkipRetry())
} }
...@@ -32,7 +38,7 @@ func AudioHelper(c *gin.Context, info *relaycommon.RelayInfo) (newAPIError *type ...@@ -32,7 +38,7 @@ func AudioHelper(c *gin.Context, info *relaycommon.RelayInfo) (newAPIError *type
} }
adaptor.Init(info) adaptor.Init(info)
ioReader, err := adaptor.ConvertAudioRequest(c, info, *audioRequest) ioReader, err := adaptor.ConvertAudioRequest(c, info, *request)
if err != nil { if err != nil {
return types.NewError(err, types.ErrorCodeConvertRequestFailed, types.ErrOptionWithSkipRetry()) return types.NewError(err, types.ErrorCodeConvertRequestFailed, types.ErrOptionWithSkipRetry())
} }
......
...@@ -120,15 +120,14 @@ func (a *Adaptor) DoResponse(c *gin.Context, resp *http.Response, info *relaycom ...@@ -120,15 +120,14 @@ func (a *Adaptor) DoResponse(c *gin.Context, resp *http.Response, info *relaycom
switch info.RelayFormat { switch info.RelayFormat {
case types.RelayFormatClaude: case types.RelayFormatClaude:
if info.IsStream { if info.IsStream {
err, usage = claude.ClaudeStreamHandler(c, resp, info, claude.RequestModeMessage) return claude.ClaudeStreamHandler(c, resp, info, claude.RequestModeMessage)
} else { } else {
err, usage = claude.ClaudeHandler(c, resp, info, claude.RequestModeMessage) return claude.ClaudeHandler(c, resp, info, claude.RequestModeMessage)
} }
default: default:
adaptor := openai.Adaptor{} adaptor := openai.Adaptor{}
return adaptor.DoResponse(c, resp, info) return adaptor.DoResponse(c, resp, info)
} }
return
} }
func (a *Adaptor) GetModelList() []string { func (a *Adaptor) GetModelList() []string {
......
...@@ -102,9 +102,9 @@ func (a *Adaptor) DoRequest(c *gin.Context, info *relaycommon.RelayInfo, request ...@@ -102,9 +102,9 @@ func (a *Adaptor) DoRequest(c *gin.Context, info *relaycommon.RelayInfo, request
func (a *Adaptor) DoResponse(c *gin.Context, resp *http.Response, info *relaycommon.RelayInfo) (usage any, err *types.NewAPIError) { func (a *Adaptor) DoResponse(c *gin.Context, resp *http.Response, info *relaycommon.RelayInfo) (usage any, err *types.NewAPIError) {
if info.IsStream { if info.IsStream {
err, usage = ClaudeStreamHandler(c, resp, info, a.RequestMode) return ClaudeStreamHandler(c, resp, info, a.RequestMode)
} else { } else {
err, usage = ClaudeHandler(c, resp, info, a.RequestMode) return ClaudeHandler(c, resp, info, a.RequestMode)
} }
return return
} }
......
...@@ -674,7 +674,7 @@ func HandleStreamFinalResponse(c *gin.Context, info *relaycommon.RelayInfo, clau ...@@ -674,7 +674,7 @@ func HandleStreamFinalResponse(c *gin.Context, info *relaycommon.RelayInfo, clau
} }
} }
func ClaudeStreamHandler(c *gin.Context, resp *http.Response, info *relaycommon.RelayInfo, requestMode int) (*types.NewAPIError, *dto.Usage) { func ClaudeStreamHandler(c *gin.Context, resp *http.Response, info *relaycommon.RelayInfo, requestMode int) (*dto.Usage, *types.NewAPIError) {
claudeInfo := &ClaudeResponseInfo{ claudeInfo := &ClaudeResponseInfo{
ResponseId: helper.GetResponseID(c), ResponseId: helper.GetResponseID(c),
Created: common.GetTimestamp(), Created: common.GetTimestamp(),
...@@ -691,11 +691,11 @@ func ClaudeStreamHandler(c *gin.Context, resp *http.Response, info *relaycommon. ...@@ -691,11 +691,11 @@ func ClaudeStreamHandler(c *gin.Context, resp *http.Response, info *relaycommon.
return true return true
}) })
if err != nil { if err != nil {
return err, nil return nil, err
} }
HandleStreamFinalResponse(c, info, claudeInfo, requestMode) HandleStreamFinalResponse(c, info, claudeInfo, requestMode)
return nil, claudeInfo.Usage return claudeInfo.Usage, nil
} }
func HandleClaudeResponseData(c *gin.Context, info *relaycommon.RelayInfo, claudeInfo *ClaudeResponseInfo, data []byte, requestMode int) *types.NewAPIError { func HandleClaudeResponseData(c *gin.Context, info *relaycommon.RelayInfo, claudeInfo *ClaudeResponseInfo, data []byte, requestMode int) *types.NewAPIError {
...@@ -740,7 +740,7 @@ func HandleClaudeResponseData(c *gin.Context, info *relaycommon.RelayInfo, claud ...@@ -740,7 +740,7 @@ func HandleClaudeResponseData(c *gin.Context, info *relaycommon.RelayInfo, claud
return nil return nil
} }
func ClaudeHandler(c *gin.Context, resp *http.Response, info *relaycommon.RelayInfo, requestMode int) (*types.NewAPIError, *dto.Usage) { func ClaudeHandler(c *gin.Context, resp *http.Response, info *relaycommon.RelayInfo, requestMode int) (*dto.Usage, *types.NewAPIError) {
defer service.CloseResponseBodyGracefully(resp) defer service.CloseResponseBodyGracefully(resp)
claudeInfo := &ClaudeResponseInfo{ claudeInfo := &ClaudeResponseInfo{
...@@ -752,16 +752,16 @@ func ClaudeHandler(c *gin.Context, resp *http.Response, info *relaycommon.RelayI ...@@ -752,16 +752,16 @@ func ClaudeHandler(c *gin.Context, resp *http.Response, info *relaycommon.RelayI
} }
responseBody, err := io.ReadAll(resp.Body) responseBody, err := io.ReadAll(resp.Body)
if err != nil { if err != nil {
return types.NewError(err, types.ErrorCodeBadResponseBody), nil return nil, types.NewError(err, types.ErrorCodeBadResponseBody)
} }
if common.DebugEnabled { if common.DebugEnabled {
println("responseBody: ", string(responseBody)) println("responseBody: ", string(responseBody))
} }
handleErr := HandleClaudeResponseData(c, info, claudeInfo, responseBody, requestMode) handleErr := HandleClaudeResponseData(c, info, claudeInfo, responseBody, requestMode)
if handleErr != nil { if handleErr != nil {
return handleErr, nil return nil, handleErr
} }
return nil, claudeInfo.Usage return claudeInfo.Usage, nil
} }
func mapToolChoice(toolChoice any, parallelToolCalls *bool) *dto.ClaudeToolChoice { func mapToolChoice(toolChoice any, parallelToolCalls *bool) *dto.ClaudeToolChoice {
......
...@@ -89,17 +89,16 @@ func (a *Adaptor) ConvertEmbeddingRequest(c *gin.Context, info *relaycommon.Rela ...@@ -89,17 +89,16 @@ func (a *Adaptor) ConvertEmbeddingRequest(c *gin.Context, info *relaycommon.Rela
func (a *Adaptor) DoResponse(c *gin.Context, resp *http.Response, info *relaycommon.RelayInfo) (usage any, err *types.NewAPIError) { func (a *Adaptor) DoResponse(c *gin.Context, resp *http.Response, info *relaycommon.RelayInfo) (usage any, err *types.NewAPIError) {
switch info.RelayFormat { switch info.RelayFormat {
case types.RelayFormatOpenAI:
adaptor := openai.Adaptor{}
return adaptor.DoResponse(c, resp, info)
case types.RelayFormatClaude: case types.RelayFormatClaude:
if info.IsStream { if info.IsStream {
err, usage = claude.ClaudeStreamHandler(c, resp, info, claude.RequestModeMessage) return claude.ClaudeStreamHandler(c, resp, info, claude.RequestModeMessage)
} else { } else {
err, usage = claude.ClaudeHandler(c, resp, info, claude.RequestModeMessage) return claude.ClaudeHandler(c, resp, info, claude.RequestModeMessage)
} }
default:
adaptor := openai.Adaptor{}
return adaptor.DoResponse(c, resp, info)
} }
return
} }
func (a *Adaptor) GetModelList() []string { func (a *Adaptor) GetModelList() []string {
......
...@@ -279,31 +279,31 @@ func (a *Adaptor) DoResponse(c *gin.Context, resp *http.Response, info *relaycom ...@@ -279,31 +279,31 @@ func (a *Adaptor) DoResponse(c *gin.Context, resp *http.Response, info *relaycom
if info.IsStream { if info.IsStream {
switch a.RequestMode { switch a.RequestMode {
case RequestModeClaude: case RequestModeClaude:
err, usage = claude.ClaudeStreamHandler(c, resp, info, claude.RequestModeMessage) return claude.ClaudeStreamHandler(c, resp, info, claude.RequestModeMessage)
case RequestModeGemini: case RequestModeGemini:
if info.RelayMode == constant.RelayModeGemini { if info.RelayMode == constant.RelayModeGemini {
usage, err = gemini.GeminiTextGenerationStreamHandler(c, info, resp) return gemini.GeminiTextGenerationStreamHandler(c, info, resp)
} else { } else {
usage, err = gemini.GeminiChatStreamHandler(c, info, resp) return gemini.GeminiChatStreamHandler(c, info, resp)
} }
case RequestModeLlama: case RequestModeLlama:
usage, err = openai.OaiStreamHandler(c, info, resp) return openai.OaiStreamHandler(c, info, resp)
} }
} else { } else {
switch a.RequestMode { switch a.RequestMode {
case RequestModeClaude: case RequestModeClaude:
err, usage = claude.ClaudeHandler(c, resp, info, claude.RequestModeMessage) return claude.ClaudeHandler(c, resp, info, claude.RequestModeMessage)
case RequestModeGemini: case RequestModeGemini:
if info.RelayMode == constant.RelayModeGemini { if info.RelayMode == constant.RelayModeGemini {
usage, err = gemini.GeminiTextGenerationHandler(c, info, resp) return gemini.GeminiTextGenerationHandler(c, info, resp)
} else { } else {
if strings.HasPrefix(info.UpstreamModelName, "imagen") { if strings.HasPrefix(info.UpstreamModelName, "imagen") {
return gemini.GeminiImageHandler(c, info, resp) return gemini.GeminiImageHandler(c, info, resp)
} }
usage, err = gemini.GeminiChatHandler(c, info, resp) return gemini.GeminiChatHandler(c, info, resp)
} }
case RequestModeLlama: case RequestModeLlama:
usage, err = openai.OpenaiHandler(c, info, resp) return openai.OpenaiHandler(c, info, resp)
} }
} }
return return
......
...@@ -7,6 +7,7 @@ import ( ...@@ -7,6 +7,7 @@ import (
"net/http" "net/http"
"one-api/dto" "one-api/dto"
"one-api/relay/channel" "one-api/relay/channel"
"one-api/relay/channel/claude"
"one-api/relay/channel/openai" "one-api/relay/channel/openai"
relaycommon "one-api/relay/common" relaycommon "one-api/relay/common"
relayconstant "one-api/relay/constant" relayconstant "one-api/relay/constant"
...@@ -23,10 +24,8 @@ func (a *Adaptor) ConvertGeminiRequest(*gin.Context, *relaycommon.RelayInfo, *dt ...@@ -23,10 +24,8 @@ func (a *Adaptor) ConvertGeminiRequest(*gin.Context, *relaycommon.RelayInfo, *dt
return nil, errors.New("not implemented") return nil, errors.New("not implemented")
} }
func (a *Adaptor) ConvertClaudeRequest(*gin.Context, *relaycommon.RelayInfo, *dto.ClaudeRequest) (any, error) { func (a *Adaptor) ConvertClaudeRequest(c *gin.Context, info *relaycommon.RelayInfo, req *dto.ClaudeRequest) (any, error) {
//TODO implement me return req, nil
panic("implement me")
return nil, nil
} }
func (a *Adaptor) ConvertAudioRequest(c *gin.Context, info *relaycommon.RelayInfo, request dto.AudioRequest) (io.Reader, error) { func (a *Adaptor) ConvertAudioRequest(c *gin.Context, info *relaycommon.RelayInfo, request dto.AudioRequest) (io.Reader, error) {
...@@ -43,12 +42,16 @@ func (a *Adaptor) Init(info *relaycommon.RelayInfo) { ...@@ -43,12 +42,16 @@ func (a *Adaptor) Init(info *relaycommon.RelayInfo) {
} }
func (a *Adaptor) GetRequestURL(info *relaycommon.RelayInfo) (string, error) { func (a *Adaptor) GetRequestURL(info *relaycommon.RelayInfo) (string, error) {
baseUrl := fmt.Sprintf("%s/api/paas/v4", info.ChannelBaseUrl) switch info.RelayFormat {
switch info.RelayMode { case types.RelayFormatClaude:
case relayconstant.RelayModeEmbeddings: return fmt.Sprintf("%s/api/anthropic/v1/messages", info.ChannelBaseUrl), nil
return fmt.Sprintf("%s/embeddings", baseUrl), nil
default: default:
return fmt.Sprintf("%s/chat/completions", baseUrl), nil switch info.RelayMode {
case relayconstant.RelayModeEmbeddings:
return fmt.Sprintf("%s/api/paas/v4/embeddings", info.ChannelBaseUrl), nil
default:
return fmt.Sprintf("%s/api/paas/v4/chat/completions", info.ChannelBaseUrl), nil
}
} }
} }
...@@ -86,12 +89,17 @@ func (a *Adaptor) DoRequest(c *gin.Context, info *relaycommon.RelayInfo, request ...@@ -86,12 +89,17 @@ func (a *Adaptor) DoRequest(c *gin.Context, info *relaycommon.RelayInfo, request
} }
func (a *Adaptor) DoResponse(c *gin.Context, resp *http.Response, info *relaycommon.RelayInfo) (usage any, err *types.NewAPIError) { func (a *Adaptor) DoResponse(c *gin.Context, resp *http.Response, info *relaycommon.RelayInfo) (usage any, err *types.NewAPIError) {
if info.IsStream { switch info.RelayFormat {
usage, err = openai.OaiStreamHandler(c, info, resp) case types.RelayFormatClaude:
} else { if info.IsStream {
usage, err = openai.OpenaiHandler(c, info, resp) return claude.ClaudeStreamHandler(c, resp, info, claude.RequestModeMessage)
} else {
return claude.ClaudeHandler(c, resp, info, claude.RequestModeMessage)
}
default:
adaptor := openai.Adaptor{}
return adaptor.DoResponse(c, resp, info)
} }
return
} }
func (a *Adaptor) GetModelList() []string { func (a *Adaptor) GetModelList() []string {
......
...@@ -21,13 +21,18 @@ func ClaudeHelper(c *gin.Context, info *relaycommon.RelayInfo) (newAPIError *typ ...@@ -21,13 +21,18 @@ func ClaudeHelper(c *gin.Context, info *relaycommon.RelayInfo) (newAPIError *typ
info.InitChannelMeta(c) info.InitChannelMeta(c)
textRequest, ok := info.Request.(*dto.ClaudeRequest) claudeReq, ok := info.Request.(*dto.ClaudeRequest)
if !ok { if !ok {
common.FatalLog(fmt.Sprintf("invalid request type, expected *dto.ClaudeRequest, got %T", info.Request)) return types.NewErrorWithStatusCode(fmt.Errorf("invalid request type, expected *dto.ClaudeRequest, got %T", info.Request), types.ErrorCodeInvalidRequest, http.StatusBadRequest, types.ErrOptionWithSkipRetry())
} }
err := helper.ModelMappedHelper(c, info, textRequest) request, err := common.DeepCopy(claudeReq)
if err != nil {
return types.NewError(fmt.Errorf("failed to copy request to ClaudeRequest: %w", err), types.ErrorCodeInvalidRequest, types.ErrOptionWithSkipRetry())
}
err = helper.ModelMappedHelper(c, info, request)
if err != nil { if err != nil {
return types.NewError(err, types.ErrorCodeChannelModelMappedError, types.ErrOptionWithSkipRetry()) return types.NewError(err, types.ErrorCodeChannelModelMappedError, types.ErrOptionWithSkipRetry())
} }
...@@ -38,30 +43,30 @@ func ClaudeHelper(c *gin.Context, info *relaycommon.RelayInfo) (newAPIError *typ ...@@ -38,30 +43,30 @@ func ClaudeHelper(c *gin.Context, info *relaycommon.RelayInfo) (newAPIError *typ
} }
adaptor.Init(info) adaptor.Init(info)
if textRequest.MaxTokens == 0 { if request.MaxTokens == 0 {
textRequest.MaxTokens = uint(model_setting.GetClaudeSettings().GetDefaultMaxTokens(textRequest.Model)) request.MaxTokens = uint(model_setting.GetClaudeSettings().GetDefaultMaxTokens(request.Model))
} }
if model_setting.GetClaudeSettings().ThinkingAdapterEnabled && if model_setting.GetClaudeSettings().ThinkingAdapterEnabled &&
strings.HasSuffix(textRequest.Model, "-thinking") { strings.HasSuffix(request.Model, "-thinking") {
if textRequest.Thinking == nil { if request.Thinking == nil {
// 因为BudgetTokens 必须大于1024 // 因为BudgetTokens 必须大于1024
if textRequest.MaxTokens < 1280 { if request.MaxTokens < 1280 {
textRequest.MaxTokens = 1280 request.MaxTokens = 1280
} }
// BudgetTokens 为 max_tokens 的 80% // BudgetTokens 为 max_tokens 的 80%
textRequest.Thinking = &dto.Thinking{ request.Thinking = &dto.Thinking{
Type: "enabled", Type: "enabled",
BudgetTokens: common.GetPointer[int](int(float64(textRequest.MaxTokens) * model_setting.GetClaudeSettings().ThinkingAdapterBudgetTokensPercentage)), BudgetTokens: common.GetPointer[int](int(float64(request.MaxTokens) * model_setting.GetClaudeSettings().ThinkingAdapterBudgetTokensPercentage)),
} }
// TODO: 临时处理 // TODO: 临时处理
// https://docs.anthropic.com/en/docs/build-with-claude/extended-thinking#important-considerations-when-using-extended-thinking // https://docs.anthropic.com/en/docs/build-with-claude/extended-thinking#important-considerations-when-using-extended-thinking
textRequest.TopP = 0 request.TopP = 0
textRequest.Temperature = common.GetPointer[float64](1.0) request.Temperature = common.GetPointer[float64](1.0)
} }
textRequest.Model = strings.TrimSuffix(textRequest.Model, "-thinking") request.Model = strings.TrimSuffix(request.Model, "-thinking")
info.UpstreamModelName = textRequest.Model info.UpstreamModelName = request.Model
} }
var requestBody io.Reader var requestBody io.Reader
...@@ -72,7 +77,7 @@ func ClaudeHelper(c *gin.Context, info *relaycommon.RelayInfo) (newAPIError *typ ...@@ -72,7 +77,7 @@ func ClaudeHelper(c *gin.Context, info *relaycommon.RelayInfo) (newAPIError *typ
} }
requestBody = bytes.NewBuffer(body) requestBody = bytes.NewBuffer(body)
} else { } else {
convertedRequest, err := adaptor.ConvertClaudeRequest(c, info, textRequest) convertedRequest, err := adaptor.ConvertClaudeRequest(c, info, request)
if err != nil { if err != nil {
return types.NewError(err, types.ErrorCodeConvertRequestFailed, types.ErrOptionWithSkipRetry()) return types.NewError(err, types.ErrorCodeConvertRequestFailed, types.ErrOptionWithSkipRetry())
} }
......
...@@ -158,7 +158,14 @@ func (info *RelayInfo) InitChannelMeta(c *gin.Context) { ...@@ -158,7 +158,14 @@ func (info *RelayInfo) InitChannelMeta(c *gin.Context) {
if streamSupportedChannels[channelMeta.ChannelType] { if streamSupportedChannels[channelMeta.ChannelType] {
channelMeta.SupportStreamOptions = true channelMeta.SupportStreamOptions = true
} }
info.ChannelMeta = channelMeta info.ChannelMeta = channelMeta
// reset some fields based on channel meta
// 重置某些字段,例如模型名称等
if info.Request != nil {
info.Request.SetModelName(info.OriginModelName)
}
} }
func (info *RelayInfo) ToString() string { func (info *RelayInfo) ToString() string {
...@@ -470,6 +477,7 @@ func GenTaskRelayInfo(c *gin.Context) (*TaskRelayInfo, error) { ...@@ -470,6 +477,7 @@ func GenTaskRelayInfo(c *gin.Context) (*TaskRelayInfo, error) {
info := &TaskRelayInfo{ info := &TaskRelayInfo{
RelayInfo: relayInfo, RelayInfo: relayInfo,
} }
info.InitChannelMeta(c)
return info, nil return info, nil
} }
......
...@@ -25,38 +25,40 @@ import ( ...@@ -25,38 +25,40 @@ import (
) )
func TextHelper(c *gin.Context, info *relaycommon.RelayInfo) (newAPIError *types.NewAPIError) { func TextHelper(c *gin.Context, info *relaycommon.RelayInfo) (newAPIError *types.NewAPIError) {
info.InitChannelMeta(c) info.InitChannelMeta(c)
textRequest, ok := info.Request.(*dto.GeneralOpenAIRequest) textReq, ok := info.Request.(*dto.GeneralOpenAIRequest)
if !ok { if !ok {
//return types.NewErrorWithStatusCode(errors.New("invalid request type"), types.ErrorCodeInvalidRequest, http.StatusBadRequest, types.ErrOptionWithSkipRetry()) return types.NewErrorWithStatusCode(fmt.Errorf("invalid request type, expected dto.GeneralOpenAIRequest, got %T", info.Request), types.ErrorCodeInvalidRequest, http.StatusBadRequest, types.ErrOptionWithSkipRetry())
common.FatalLog("invalid request type, expected dto.GeneralOpenAIRequest, got %T", info.Request) }
request, err := common.DeepCopy(textReq)
if err != nil {
return types.NewError(fmt.Errorf("failed to copy request to GeneralOpenAIRequest: %w", err), types.ErrorCodeInvalidRequest, types.ErrOptionWithSkipRetry())
} }
if textRequest.WebSearchOptions != nil { if request.WebSearchOptions != nil {
c.Set("chat_completion_web_search_context_size", textRequest.WebSearchOptions.SearchContextSize) c.Set("chat_completion_web_search_context_size", request.WebSearchOptions.SearchContextSize)
} }
err := helper.ModelMappedHelper(c, info, textRequest) err = helper.ModelMappedHelper(c, info, request)
if err != nil { if err != nil {
return types.NewError(err, types.ErrorCodeChannelModelMappedError, types.ErrOptionWithSkipRetry()) return types.NewError(err, types.ErrorCodeChannelModelMappedError, types.ErrOptionWithSkipRetry())
} }
includeUsage := true includeUsage := true
// 判断用户是否需要返回使用情况 // 判断用户是否需要返回使用情况
if textRequest.StreamOptions != nil { if request.StreamOptions != nil {
includeUsage = textRequest.StreamOptions.IncludeUsage includeUsage = request.StreamOptions.IncludeUsage
} }
// 如果不支持StreamOptions,将StreamOptions设置为nil // 如果不支持StreamOptions,将StreamOptions设置为nil
if !info.SupportStreamOptions || !textRequest.Stream { if !info.SupportStreamOptions || !request.Stream {
textRequest.StreamOptions = nil request.StreamOptions = nil
} else { } else {
// 如果支持StreamOptions,且请求中没有设置StreamOptions,根据配置文件设置StreamOptions // 如果支持StreamOptions,且请求中没有设置StreamOptions,根据配置文件设置StreamOptions
if constant.ForceStreamOption { if constant.ForceStreamOption {
textRequest.StreamOptions = &dto.StreamOptions{ request.StreamOptions = &dto.StreamOptions{
IncludeUsage: true, IncludeUsage: true,
} }
} }
...@@ -81,7 +83,7 @@ func TextHelper(c *gin.Context, info *relaycommon.RelayInfo) (newAPIError *types ...@@ -81,7 +83,7 @@ func TextHelper(c *gin.Context, info *relaycommon.RelayInfo) (newAPIError *types
} }
requestBody = bytes.NewBuffer(body) requestBody = bytes.NewBuffer(body)
} else { } else {
convertedRequest, err := adaptor.ConvertOpenAIRequest(c, info, textRequest) convertedRequest, err := adaptor.ConvertOpenAIRequest(c, info, request)
if err != nil { if err != nil {
return types.NewError(err, types.ErrorCodeConvertRequestFailed, types.ErrOptionWithSkipRetry()) return types.NewError(err, types.ErrorCodeConvertRequestFailed, types.ErrOptionWithSkipRetry())
} }
......
...@@ -16,15 +16,19 @@ import ( ...@@ -16,15 +16,19 @@ import (
) )
func EmbeddingHelper(c *gin.Context, info *relaycommon.RelayInfo) (newAPIError *types.NewAPIError) { func EmbeddingHelper(c *gin.Context, info *relaycommon.RelayInfo) (newAPIError *types.NewAPIError) {
info.InitChannelMeta(c) info.InitChannelMeta(c)
embeddingRequest, ok := info.Request.(*dto.EmbeddingRequest) embeddingReq, ok := info.Request.(*dto.EmbeddingRequest)
if !ok { if !ok {
common.FatalLog(fmt.Sprintf("invalid request type, expected *dto.EmbeddingRequest, got %T", info.Request)) return types.NewErrorWithStatusCode(fmt.Errorf("invalid request type, expected *dto.EmbeddingRequest, got %T", info.Request), types.ErrorCodeInvalidRequest, http.StatusBadRequest, types.ErrOptionWithSkipRetry())
}
request, err := common.DeepCopy(embeddingReq)
if err != nil {
return types.NewError(fmt.Errorf("failed to copy request to EmbeddingRequest: %w", err), types.ErrorCodeInvalidRequest, types.ErrOptionWithSkipRetry())
} }
err := helper.ModelMappedHelper(c, info, embeddingRequest) err = helper.ModelMappedHelper(c, info, request)
if err != nil { if err != nil {
return types.NewError(err, types.ErrorCodeChannelModelMappedError, types.ErrOptionWithSkipRetry()) return types.NewError(err, types.ErrorCodeChannelModelMappedError, types.ErrOptionWithSkipRetry())
} }
...@@ -35,7 +39,7 @@ func EmbeddingHelper(c *gin.Context, info *relaycommon.RelayInfo) (newAPIError * ...@@ -35,7 +39,7 @@ func EmbeddingHelper(c *gin.Context, info *relaycommon.RelayInfo) (newAPIError *
} }
adaptor.Init(info) adaptor.Init(info)
convertedRequest, err := adaptor.ConvertEmbeddingRequest(c, info, *embeddingRequest) convertedRequest, err := adaptor.ConvertEmbeddingRequest(c, info, *request)
if err != nil { if err != nil {
return types.NewError(err, types.ErrorCodeConvertRequestFailed, types.ErrOptionWithSkipRetry()) return types.NewError(err, types.ErrorCodeConvertRequestFailed, types.ErrOptionWithSkipRetry())
} }
......
...@@ -53,13 +53,18 @@ func trimModelThinking(modelName string) string { ...@@ -53,13 +53,18 @@ func trimModelThinking(modelName string) string {
func GeminiHelper(c *gin.Context, info *relaycommon.RelayInfo) (newAPIError *types.NewAPIError) { func GeminiHelper(c *gin.Context, info *relaycommon.RelayInfo) (newAPIError *types.NewAPIError) {
info.InitChannelMeta(c) info.InitChannelMeta(c)
request, ok := info.Request.(*dto.GeminiChatRequest) geminiReq, ok := info.Request.(*dto.GeminiChatRequest)
if !ok { if !ok {
common.FatalLog(fmt.Sprintf("invalid request type, expected *dto.GeminiChatRequest, got %T", info.Request)) return types.NewErrorWithStatusCode(fmt.Errorf("invalid request type, expected *dto.GeminiChatRequest, got %T", info.Request), types.ErrorCodeInvalidRequest, http.StatusBadRequest, types.ErrOptionWithSkipRetry())
}
request, err := common.DeepCopy(geminiReq)
if err != nil {
return types.NewError(fmt.Errorf("failed to copy request to GeminiChatRequest: %w", err), types.ErrorCodeInvalidRequest, types.ErrOptionWithSkipRetry())
} }
// model mapped 模型映射 // model mapped 模型映射
err := helper.ModelMappedHelper(c, info, request) err = helper.ModelMappedHelper(c, info, request)
if err != nil { if err != nil {
return types.NewError(err, types.ErrorCodeChannelModelMappedError, types.ErrOptionWithSkipRetry()) return types.NewError(err, types.ErrorCodeChannelModelMappedError, types.ErrOptionWithSkipRetry())
} }
...@@ -170,7 +175,7 @@ func GeminiEmbeddingHandler(c *gin.Context, info *relaycommon.RelayInfo) (newAPI ...@@ -170,7 +175,7 @@ func GeminiEmbeddingHandler(c *gin.Context, info *relaycommon.RelayInfo) (newAPI
isBatch := strings.HasSuffix(c.Request.URL.Path, "batchEmbedContents") isBatch := strings.HasSuffix(c.Request.URL.Path, "batchEmbedContents")
info.IsGeminiBatchEmbedding = isBatch info.IsGeminiBatchEmbedding = isBatch
var req any var req dto.Request
var err error var err error
var inputTexts []string var inputTexts []string
......
...@@ -4,15 +4,12 @@ import ( ...@@ -4,15 +4,12 @@ import (
"encoding/json" "encoding/json"
"errors" "errors"
"fmt" "fmt"
"github.com/gin-gonic/gin"
"one-api/dto" "one-api/dto"
common2 "one-api/logger"
"one-api/relay/common" "one-api/relay/common"
"one-api/types"
"github.com/gin-gonic/gin"
) )
func ModelMappedHelper(c *gin.Context, info *common.RelayInfo, request any) error { func ModelMappedHelper(c *gin.Context, info *common.RelayInfo, request dto.Request) error {
// map model name // map model name
modelMapping := c.GetString("model_mapping") modelMapping := c.GetString("model_mapping")
if modelMapping != "" && modelMapping != "{}" { if modelMapping != "" && modelMapping != "{}" {
...@@ -54,40 +51,7 @@ func ModelMappedHelper(c *gin.Context, info *common.RelayInfo, request any) erro ...@@ -54,40 +51,7 @@ func ModelMappedHelper(c *gin.Context, info *common.RelayInfo, request any) erro
} }
} }
if request != nil { if request != nil {
switch info.RelayFormat { request.SetModelName(info.UpstreamModelName)
case types.RelayFormatGemini:
// Gemini 模型映射
case types.RelayFormatClaude:
if claudeRequest, ok := request.(*dto.ClaudeRequest); ok {
claudeRequest.Model = info.UpstreamModelName
}
case types.RelayFormatOpenAIResponses:
if openAIResponsesRequest, ok := request.(*dto.OpenAIResponsesRequest); ok {
openAIResponsesRequest.Model = info.UpstreamModelName
}
case types.RelayFormatOpenAIAudio:
if openAIAudioRequest, ok := request.(*dto.AudioRequest); ok {
openAIAudioRequest.Model = info.UpstreamModelName
}
case types.RelayFormatOpenAIImage:
if imageRequest, ok := request.(*dto.ImageRequest); ok {
imageRequest.Model = info.UpstreamModelName
}
case types.RelayFormatRerank:
if rerankRequest, ok := request.(*dto.RerankRequest); ok {
rerankRequest.Model = info.UpstreamModelName
}
case types.RelayFormatEmbedding:
if embeddingRequest, ok := request.(*dto.EmbeddingRequest); ok {
embeddingRequest.Model = info.UpstreamModelName
}
default:
if openAIRequest, ok := request.(*dto.GeneralOpenAIRequest); ok {
openAIRequest.Model = info.UpstreamModelName
} else {
common2.LogWarn(c, fmt.Sprintf("model mapped but request type %T not supported", request))
}
}
} }
return nil return nil
} }
...@@ -20,16 +20,19 @@ import ( ...@@ -20,16 +20,19 @@ import (
) )
func ImageHelper(c *gin.Context, info *relaycommon.RelayInfo) (newAPIError *types.NewAPIError) { func ImageHelper(c *gin.Context, info *relaycommon.RelayInfo) (newAPIError *types.NewAPIError) {
info.InitChannelMeta(c) info.InitChannelMeta(c)
imageRequest, ok := info.Request.(*dto.ImageRequest) imageReq, ok := info.Request.(*dto.ImageRequest)
if !ok { if !ok {
common.FatalLog(fmt.Sprintf("invalid request type, expected dto.ImageRequest, got %T", info.Request)) return types.NewErrorWithStatusCode(fmt.Errorf("invalid request type, expected dto.ImageRequest, got %T", info.Request), types.ErrorCodeInvalidRequest, http.StatusBadRequest, types.ErrOptionWithSkipRetry())
}
request, err := common.DeepCopy(imageReq)
if err != nil {
return types.NewError(fmt.Errorf("failed to copy request to ImageRequest: %w", err), types.ErrorCodeInvalidRequest, types.ErrOptionWithSkipRetry())
} }
err := helper.ModelMappedHelper(c, info, imageRequest) err = helper.ModelMappedHelper(c, info, request)
if err != nil { if err != nil {
return types.NewError(err, types.ErrorCodeChannelModelMappedError, types.ErrOptionWithSkipRetry()) return types.NewError(err, types.ErrorCodeChannelModelMappedError, types.ErrOptionWithSkipRetry())
} }
...@@ -49,7 +52,7 @@ func ImageHelper(c *gin.Context, info *relaycommon.RelayInfo) (newAPIError *type ...@@ -49,7 +52,7 @@ func ImageHelper(c *gin.Context, info *relaycommon.RelayInfo) (newAPIError *type
} }
requestBody = bytes.NewBuffer(body) requestBody = bytes.NewBuffer(body)
} else { } else {
convertedRequest, err := adaptor.ConvertImageRequest(c, info, *imageRequest) convertedRequest, err := adaptor.ConvertImageRequest(c, info, *request)
if err != nil { if err != nil {
return types.NewError(err, types.ErrorCodeConvertRequestFailed) return types.NewError(err, types.ErrorCodeConvertRequestFailed)
} }
...@@ -102,21 +105,21 @@ func ImageHelper(c *gin.Context, info *relaycommon.RelayInfo) (newAPIError *type ...@@ -102,21 +105,21 @@ func ImageHelper(c *gin.Context, info *relaycommon.RelayInfo) (newAPIError *type
} }
if usage.(*dto.Usage).TotalTokens == 0 { if usage.(*dto.Usage).TotalTokens == 0 {
usage.(*dto.Usage).TotalTokens = int(imageRequest.N) usage.(*dto.Usage).TotalTokens = int(request.N)
} }
if usage.(*dto.Usage).PromptTokens == 0 { if usage.(*dto.Usage).PromptTokens == 0 {
usage.(*dto.Usage).PromptTokens = int(imageRequest.N) usage.(*dto.Usage).PromptTokens = int(request.N)
} }
quality := "standard" quality := "standard"
if imageRequest.Quality == "hd" { if request.Quality == "hd" {
quality = "hd" quality = "hd"
} }
var logContent string var logContent string
if len(imageRequest.Size) > 0 { if len(request.Size) > 0 {
logContent = fmt.Sprintf("大小 %s, 品质 %s", imageRequest.Size, quality) logContent = fmt.Sprintf("大小 %s, 品质 %s", request.Size, quality)
} }
postConsumeQuota(c, info, usage.(*dto.Usage), logContent) postConsumeQuota(c, info, usage.(*dto.Usage), logContent)
......
...@@ -16,23 +16,20 @@ import ( ...@@ -16,23 +16,20 @@ import (
"github.com/gin-gonic/gin" "github.com/gin-gonic/gin"
) )
func getRerankPromptToken(rerankRequest dto.RerankRequest) int {
token := service.CountTokenInput(rerankRequest.Query, rerankRequest.Model)
for _, document := range rerankRequest.Documents {
tkm := service.CountTokenInput(document, rerankRequest.Model)
token += tkm
}
return token
}
func RerankHelper(c *gin.Context, info *relaycommon.RelayInfo) (newAPIError *types.NewAPIError) { func RerankHelper(c *gin.Context, info *relaycommon.RelayInfo) (newAPIError *types.NewAPIError) {
info.InitChannelMeta(c)
rerankRequest, ok := info.Request.(*dto.RerankRequest) rerankReq, ok := info.Request.(*dto.RerankRequest)
if !ok { if !ok {
common.FatalLog(fmt.Sprintf("invalid request type, expected dto.RerankRequest, got %T", info.Request)) return types.NewErrorWithStatusCode(fmt.Errorf("invalid request type, expected dto.RerankRequest, got %T", info.Request), types.ErrorCodeInvalidRequest, http.StatusBadRequest, types.ErrOptionWithSkipRetry())
}
request, err := common.DeepCopy(rerankReq)
if err != nil {
return types.NewError(fmt.Errorf("failed to copy request to ImageRequest: %w", err), types.ErrorCodeInvalidRequest, types.ErrOptionWithSkipRetry())
} }
err := helper.ModelMappedHelper(c, info, rerankRequest) err = helper.ModelMappedHelper(c, info, request)
if err != nil { if err != nil {
return types.NewError(err, types.ErrorCodeChannelModelMappedError, types.ErrOptionWithSkipRetry()) return types.NewError(err, types.ErrorCodeChannelModelMappedError, types.ErrOptionWithSkipRetry())
} }
...@@ -51,7 +48,7 @@ func RerankHelper(c *gin.Context, info *relaycommon.RelayInfo) (newAPIError *typ ...@@ -51,7 +48,7 @@ func RerankHelper(c *gin.Context, info *relaycommon.RelayInfo) (newAPIError *typ
} }
requestBody = bytes.NewBuffer(body) requestBody = bytes.NewBuffer(body)
} else { } else {
convertedRequest, err := adaptor.ConvertRerankRequest(c, info.RelayMode, *rerankRequest) convertedRequest, err := adaptor.ConvertRerankRequest(c, info.RelayMode, *request)
if err != nil { if err != nil {
return types.NewError(err, types.ErrorCodeConvertRequestFailed, types.ErrOptionWithSkipRetry()) return types.NewError(err, types.ErrorCodeConvertRequestFailed, types.ErrOptionWithSkipRetry())
} }
......
...@@ -20,12 +20,17 @@ import ( ...@@ -20,12 +20,17 @@ import (
func ResponsesHelper(c *gin.Context, info *relaycommon.RelayInfo) (newAPIError *types.NewAPIError) { func ResponsesHelper(c *gin.Context, info *relaycommon.RelayInfo) (newAPIError *types.NewAPIError) {
info.InitChannelMeta(c) info.InitChannelMeta(c)
request, ok := info.Request.(*dto.OpenAIResponsesRequest) responsesReq, ok := info.Request.(*dto.OpenAIResponsesRequest)
if !ok { if !ok {
common.FatalLog(fmt.Sprintf("invalid request type, expected dto.OpenAIResponsesRequest, got %T", info.Request)) return types.NewErrorWithStatusCode(fmt.Errorf("invalid request type, expected dto.OpenAIResponsesRequest, got %T", info.Request), types.ErrorCodeInvalidRequest, http.StatusBadRequest, types.ErrOptionWithSkipRetry())
} }
err := helper.ModelMappedHelper(c, info, request) request, err := common.DeepCopy(responsesReq)
if err != nil {
return types.NewError(fmt.Errorf("failed to copy request to GeneralOpenAIRequest: %w", err), types.ErrorCodeInvalidRequest, types.ErrOptionWithSkipRetry())
}
err = helper.ModelMappedHelper(c, info, request)
if err != nil { if err != nil {
return types.NewError(err, types.ErrorCodeChannelModelMappedError, types.ErrOptionWithSkipRetry()) return types.NewError(err, types.ErrorCodeChannelModelMappedError, types.ErrOptionWithSkipRetry())
} }
......
...@@ -145,6 +145,17 @@ func SetApiRouter(router *gin.Engine) { ...@@ -145,6 +145,17 @@ func SetApiRouter(router *gin.Engine) {
tokenRoute.DELETE("/:id", controller.DeleteToken) tokenRoute.DELETE("/:id", controller.DeleteToken)
tokenRoute.POST("/batch", controller.DeleteTokenBatch) tokenRoute.POST("/batch", controller.DeleteTokenBatch)
} }
usageRoute := apiRouter.Group("/usage")
usageRoute.Use(middleware.CriticalRateLimit())
{
tokenUsageRoute := usageRoute.Group("/token")
tokenUsageRoute.Use(middleware.TokenAuth())
{
tokenUsageRoute.GET("/", controller.GetTokenUsage)
}
}
redemptionRoute := apiRouter.Group("/redemption") redemptionRoute := apiRouter.Group("/redemption")
redemptionRoute.Use(middleware.AdminAuth()) redemptionRoute.Use(middleware.AdminAuth())
{ {
...@@ -172,7 +183,6 @@ func SetApiRouter(router *gin.Engine) { ...@@ -172,7 +183,6 @@ func SetApiRouter(router *gin.Engine) {
logRoute.Use(middleware.CORS()) logRoute.Use(middleware.CORS())
{ {
logRoute.GET("/token", controller.GetLogByKey) logRoute.GET("/token", controller.GetLogByKey)
} }
groupRoute := apiRouter.Group("/group") groupRoute := apiRouter.Group("/group")
groupRoute.Use(middleware.AdminAuth()) groupRoute.Use(middleware.AdminAuth())
......
Markdown is supported
0% or
You are about to add 0 people to the discussion. Proceed with caution.
Finish editing this message first!
Please register or sign in to comment