Commit 252aff12 by wzxjohn Committed by GitHub

Merge branch 'alpha' into feature/simple_stripe

parents 35348892 31adc393
...@@ -2,6 +2,7 @@ FROM oven/bun:latest AS builder ...@@ -2,6 +2,7 @@ FROM oven/bun:latest AS builder
WORKDIR /build WORKDIR /build
COPY web/package.json . COPY web/package.json .
COPY web/bun.lock .
RUN bun install RUN bun install
COPY ./web . COPY ./web .
COPY ./VERSION . COPY ./VERSION .
......
...@@ -4,6 +4,7 @@ import ( ...@@ -4,6 +4,7 @@ import (
"bytes" "bytes"
"github.com/gin-gonic/gin" "github.com/gin-gonic/gin"
"io" "io"
"net/http"
"one-api/constant" "one-api/constant"
"strings" "strings"
"time" "time"
...@@ -32,7 +33,7 @@ func UnmarshalBodyReusable(c *gin.Context, v any) error { ...@@ -32,7 +33,7 @@ func UnmarshalBodyReusable(c *gin.Context, v any) error {
} }
contentType := c.Request.Header.Get("Content-Type") contentType := c.Request.Header.Get("Content-Type")
if strings.HasPrefix(contentType, "application/json") { if strings.HasPrefix(contentType, "application/json") {
err = UnmarshalJson(requestBody, &v) err = Unmarshal(requestBody, &v)
} else { } else {
// skip for now // skip for now
// TODO: someday non json request have variant model, we will need to implementation this // TODO: someday non json request have variant model, we will need to implementation this
...@@ -86,3 +87,25 @@ func GetContextKeyType[T any](c *gin.Context, key constant.ContextKey) (T, bool) ...@@ -86,3 +87,25 @@ func GetContextKeyType[T any](c *gin.Context, key constant.ContextKey) (T, bool)
var t T var t T
return t, false return t, false
} }
func ApiError(c *gin.Context, err error) {
c.JSON(http.StatusOK, gin.H{
"success": false,
"message": err.Error(),
})
}
func ApiErrorMsg(c *gin.Context, msg string) {
c.JSON(http.StatusOK, gin.H{
"success": false,
"message": msg,
})
}
func ApiSuccess(c *gin.Context, data any) {
c.JSON(http.StatusOK, gin.H{
"success": true,
"message": "",
"data": data,
})
}
...@@ -5,7 +5,7 @@ import ( ...@@ -5,7 +5,7 @@ import (
"encoding/json" "encoding/json"
) )
func UnmarshalJson(data []byte, v any) error { func Unmarshal(data []byte, v any) error {
return json.Unmarshal(data, v) return json.Unmarshal(data, v)
} }
...@@ -17,6 +17,6 @@ func DecodeJson(reader *bytes.Reader, v any) error { ...@@ -17,6 +17,6 @@ func DecodeJson(reader *bytes.Reader, v any) error {
return json.NewDecoder(reader).Decode(v) return json.NewDecoder(reader).Decode(v)
} }
func EncodeJson(v any) ([]byte, error) { func Marshal(v any) ([]byte, error) {
return json.Marshal(v) return json.Marshal(v)
} }
package common package common
import ( import (
"github.com/gin-gonic/gin"
"strconv" "strconv"
"github.com/gin-gonic/gin"
) )
type PageInfo struct { type PageInfo struct {
Page int `json:"page"` // page num 页码 Page int `json:"page"` // page num 页码
PageSize int `json:"page_size"` // page size 页大小 PageSize int `json:"page_size"` // page size 页大小
StartTimestamp int64 `json:"start_timestamp"` // 秒级
EndTimestamp int64 `json:"end_timestamp"` // 秒级
Total int `json:"total"` // 总条数,后设置 Total int `json:"total"` // 总条数,后设置
Items any `json:"items"` // 数据,后设置 Items any `json:"items"` // 数据,后设置
...@@ -39,11 +38,14 @@ func (p *PageInfo) SetItems(items any) { ...@@ -39,11 +38,14 @@ func (p *PageInfo) SetItems(items any) {
p.Items = items p.Items = items
} }
func GetPageQuery(c *gin.Context) (*PageInfo, error) { func GetPageQuery(c *gin.Context) *PageInfo {
pageInfo := &PageInfo{} pageInfo := &PageInfo{}
err := c.BindQuery(pageInfo) // 手动获取并处理每个参数
if err != nil { if page, err := strconv.Atoi(c.Query("page")); err == nil {
return nil, err pageInfo.Page = page
}
if pageSize, err := strconv.Atoi(c.Query("page_size")); err == nil {
pageInfo.PageSize = pageSize
} }
if pageInfo.Page < 1 { if pageInfo.Page < 1 {
// 兼容 // 兼容
...@@ -56,7 +58,25 @@ func GetPageQuery(c *gin.Context) (*PageInfo, error) { ...@@ -56,7 +58,25 @@ func GetPageQuery(c *gin.Context) (*PageInfo, error) {
} }
if pageInfo.PageSize == 0 { if pageInfo.PageSize == 0 {
pageInfo.PageSize = ItemsPerPage // 兼容
pageSize, _ := strconv.Atoi(c.Query("ps"))
if pageSize != 0 {
pageInfo.PageSize = pageSize
}
if pageInfo.PageSize == 0 {
pageSize, _ = strconv.Atoi(c.Query("size")) // token page
if pageSize != 0 {
pageInfo.PageSize = pageSize
}
}
if pageInfo.PageSize == 0 {
pageInfo.PageSize = ItemsPerPage
}
} }
return pageInfo, nil
if pageInfo.PageSize > 100 {
pageInfo.PageSize = 100
}
return pageInfo
} }
...@@ -32,16 +32,30 @@ func MapToJsonStr(m map[string]interface{}) string { ...@@ -32,16 +32,30 @@ func MapToJsonStr(m map[string]interface{}) string {
return string(bytes) return string(bytes)
} }
func StrToMap(str string) map[string]interface{} { func StrToMap(str string) (map[string]interface{}, error) {
m := make(map[string]interface{}) m := make(map[string]interface{})
err := json.Unmarshal([]byte(str), &m) err := Unmarshal([]byte(str), &m)
if err != nil { if err != nil {
return nil return nil, err
} }
return m return m, nil
} }
func IsJsonStr(str string) bool { func StrToJsonArray(str string) ([]interface{}, error) {
var js []interface{}
err := json.Unmarshal([]byte(str), &js)
if err != nil {
return nil, err
}
return js, nil
}
func IsJsonArray(str string) bool {
var js []interface{}
return json.Unmarshal([]byte(str), &js) == nil
}
func IsJsonObject(str string) bool {
var js map[string]interface{} var js map[string]interface{}
return json.Unmarshal([]byte(str), &js) == nil return json.Unmarshal([]byte(str), &js) == nil
} }
......
...@@ -17,11 +17,20 @@ const ( ...@@ -17,11 +17,20 @@ const (
ContextKeyTokenModelLimit ContextKey = "token_model_limit" ContextKeyTokenModelLimit ContextKey = "token_model_limit"
/* channel related keys */ /* channel related keys */
ContextKeyBaseUrl ContextKey = "base_url" ContextKeyChannelId ContextKey = "channel_id"
ContextKeyChannelType ContextKey = "channel_type" ContextKeyChannelName ContextKey = "channel_name"
ContextKeyChannelId ContextKey = "channel_id" ContextKeyChannelCreateTime ContextKey = "channel_create_time"
ContextKeyChannelSetting ContextKey = "channel_setting" ContextKeyChannelBaseUrl ContextKey = "base_url"
ContextKeyParamOverride ContextKey = "param_override" ContextKeyChannelType ContextKey = "channel_type"
ContextKeyChannelSetting ContextKey = "channel_setting"
ContextKeyChannelParamOverride ContextKey = "param_override"
ContextKeyChannelOrganization ContextKey = "channel_organization"
ContextKeyChannelAutoBan ContextKey = "auto_ban"
ContextKeyChannelModelMapping ContextKey = "model_mapping"
ContextKeyChannelStatusCodeMapping ContextKey = "status_code_mapping"
ContextKeyChannelIsMultiKey ContextKey = "channel_is_multi_key"
ContextKeyChannelMultiKeyIndex ContextKey = "channel_multi_key_index"
ContextKeyChannelKey ContextKey = "channel_key"
/* user related keys */ /* user related keys */
ContextKeyUserId ContextKey = "id" ContextKeyUserId ContextKey = "id"
......
package constant
type MultiKeyMode string
const (
MultiKeyModeRandom MultiKeyMode = "random" // 随机
MultiKeyModePolling MultiKeyMode = "polling" // 轮询
)
...@@ -4,7 +4,6 @@ import ( ...@@ -4,7 +4,6 @@ import (
"encoding/json" "encoding/json"
"errors" "errors"
"fmt" "fmt"
"github.com/shopspring/decimal"
"io" "io"
"net/http" "net/http"
"one-api/common" "one-api/common"
...@@ -12,9 +11,12 @@ import ( ...@@ -12,9 +11,12 @@ import (
"one-api/model" "one-api/model"
"one-api/service" "one-api/service"
"one-api/setting" "one-api/setting"
"one-api/types"
"strconv" "strconv"
"time" "time"
"github.com/shopspring/decimal"
"github.com/gin-gonic/gin" "github.com/gin-gonic/gin"
) )
...@@ -409,26 +411,24 @@ func updateChannelBalance(channel *model.Channel) (float64, error) { ...@@ -409,26 +411,24 @@ func updateChannelBalance(channel *model.Channel) (float64, error) {
func UpdateChannelBalance(c *gin.Context) { func UpdateChannelBalance(c *gin.Context) {
id, err := strconv.Atoi(c.Param("id")) id, err := strconv.Atoi(c.Param("id"))
if err != nil { if err != nil {
c.JSON(http.StatusOK, gin.H{ common.ApiError(c, err)
"success": false,
"message": err.Error(),
})
return return
} }
channel, err := model.GetChannelById(id, true) channel, err := model.CacheGetChannel(id)
if err != nil { if err != nil {
common.ApiError(c, err)
return
}
if channel.ChannelInfo.IsMultiKey {
c.JSON(http.StatusOK, gin.H{ c.JSON(http.StatusOK, gin.H{
"success": false, "success": false,
"message": err.Error(), "message": "多密钥渠道不支持余额查询",
}) })
return return
} }
balance, err := updateChannelBalance(channel) balance, err := updateChannelBalance(channel)
if err != nil { if err != nil {
c.JSON(http.StatusOK, gin.H{ common.ApiError(c, err)
"success": false,
"message": err.Error(),
})
return return
} }
c.JSON(http.StatusOK, gin.H{ c.JSON(http.StatusOK, gin.H{
...@@ -436,7 +436,6 @@ func UpdateChannelBalance(c *gin.Context) { ...@@ -436,7 +436,6 @@ func UpdateChannelBalance(c *gin.Context) {
"message": "", "message": "",
"balance": balance, "balance": balance,
}) })
return
} }
func updateAllChannelsBalance() error { func updateAllChannelsBalance() error {
...@@ -448,6 +447,9 @@ func updateAllChannelsBalance() error { ...@@ -448,6 +447,9 @@ func updateAllChannelsBalance() error {
if channel.Status != common.ChannelStatusEnabled { if channel.Status != common.ChannelStatusEnabled {
continue continue
} }
if channel.ChannelInfo.IsMultiKey {
continue // skip multi-key channels
}
// TODO: support Azure // TODO: support Azure
//if channel.Type != common.ChannelTypeOpenAI && channel.Type != common.ChannelTypeCustom { //if channel.Type != common.ChannelTypeOpenAI && channel.Type != common.ChannelTypeCustom {
// continue // continue
...@@ -458,7 +460,7 @@ func updateAllChannelsBalance() error { ...@@ -458,7 +460,7 @@ func updateAllChannelsBalance() error {
} else { } else {
// err is nil & balance <= 0 means quota is used up // err is nil & balance <= 0 means quota is used up
if balance <= 0 { if balance <= 0 {
service.DisableChannel(channel.Id, channel.Name, "余额不足") service.DisableChannel(*types.NewChannelError(channel.Id, channel.Type, channel.Name, channel.ChannelInfo.IsMultiKey, "", channel.GetAutoBan()), "余额不足")
} }
} }
time.Sleep(common.RequestInterval) time.Sleep(common.RequestInterval)
...@@ -470,10 +472,7 @@ func UpdateAllChannelsBalance(c *gin.Context) { ...@@ -470,10 +472,7 @@ func UpdateAllChannelsBalance(c *gin.Context) {
// TODO: make it async // TODO: make it async
err := updateAllChannelsBalance() err := updateAllChannelsBalance()
if err != nil { if err != nil {
c.JSON(http.StatusOK, gin.H{ common.ApiError(c, err)
"success": false,
"message": err.Error(),
})
return return
} }
c.JSON(http.StatusOK, gin.H{ c.JSON(http.StatusOK, gin.H{
......
...@@ -5,13 +5,14 @@ import ( ...@@ -5,13 +5,14 @@ import (
"encoding/json" "encoding/json"
"errors" "errors"
"fmt" "fmt"
"github.com/gin-contrib/sessions"
"github.com/gin-gonic/gin"
"net/http" "net/http"
"one-api/common" "one-api/common"
"one-api/model" "one-api/model"
"strconv" "strconv"
"time" "time"
"github.com/gin-contrib/sessions"
"github.com/gin-gonic/gin"
) )
type GitHubOAuthResponse struct { type GitHubOAuthResponse struct {
...@@ -103,10 +104,7 @@ func GitHubOAuth(c *gin.Context) { ...@@ -103,10 +104,7 @@ func GitHubOAuth(c *gin.Context) {
code := c.Query("code") code := c.Query("code")
githubUser, err := getGitHubUserInfoByCode(code) githubUser, err := getGitHubUserInfoByCode(code)
if err != nil { if err != nil {
c.JSON(http.StatusOK, gin.H{ common.ApiError(c, err)
"success": false,
"message": err.Error(),
})
return return
} }
user := model.User{ user := model.User{
...@@ -185,10 +183,7 @@ func GitHubBind(c *gin.Context) { ...@@ -185,10 +183,7 @@ func GitHubBind(c *gin.Context) {
code := c.Query("code") code := c.Query("code")
githubUser, err := getGitHubUserInfoByCode(code) githubUser, err := getGitHubUserInfoByCode(code)
if err != nil { if err != nil {
c.JSON(http.StatusOK, gin.H{ common.ApiError(c, err)
"success": false,
"message": err.Error(),
})
return return
} }
user := model.User{ user := model.User{
...@@ -207,19 +202,13 @@ func GitHubBind(c *gin.Context) { ...@@ -207,19 +202,13 @@ func GitHubBind(c *gin.Context) {
user.Id = id.(int) user.Id = id.(int)
err = user.FillUserById() err = user.FillUserById()
if err != nil { if err != nil {
c.JSON(http.StatusOK, gin.H{ common.ApiError(c, err)
"success": false,
"message": err.Error(),
})
return return
} }
user.GitHubId = githubUser.Login user.GitHubId = githubUser.Login
err = user.Update(false) err = user.Update(false)
if err != nil { if err != nil {
c.JSON(http.StatusOK, gin.H{ common.ApiError(c, err)
"success": false,
"message": err.Error(),
})
return return
} }
c.JSON(http.StatusOK, gin.H{ c.JSON(http.StatusOK, gin.H{
...@@ -239,10 +228,7 @@ func GenerateOAuthCode(c *gin.Context) { ...@@ -239,10 +228,7 @@ func GenerateOAuthCode(c *gin.Context) {
session.Set("oauth_state", state) session.Set("oauth_state", state)
err := session.Save() err := session.Save()
if err != nil { if err != nil {
c.JSON(http.StatusOK, gin.H{ common.ApiError(c, err)
"success": false,
"message": err.Error(),
})
return return
} }
c.JSON(http.StatusOK, gin.H{ c.JSON(http.StatusOK, gin.H{
......
...@@ -38,10 +38,7 @@ func LinuxDoBind(c *gin.Context) { ...@@ -38,10 +38,7 @@ func LinuxDoBind(c *gin.Context) {
code := c.Query("code") code := c.Query("code")
linuxdoUser, err := getLinuxdoUserInfoByCode(code, c) linuxdoUser, err := getLinuxdoUserInfoByCode(code, c)
if err != nil { if err != nil {
c.JSON(http.StatusOK, gin.H{ common.ApiError(c, err)
"success": false,
"message": err.Error(),
})
return return
} }
...@@ -63,20 +60,14 @@ func LinuxDoBind(c *gin.Context) { ...@@ -63,20 +60,14 @@ func LinuxDoBind(c *gin.Context) {
err = user.FillUserById() err = user.FillUserById()
if err != nil { if err != nil {
c.JSON(http.StatusOK, gin.H{ common.ApiError(c, err)
"success": false,
"message": err.Error(),
})
return return
} }
user.LinuxDOId = strconv.Itoa(linuxdoUser.Id) user.LinuxDOId = strconv.Itoa(linuxdoUser.Id)
err = user.Update(false) err = user.Update(false)
if err != nil { if err != nil {
c.JSON(http.StatusOK, gin.H{ common.ApiError(c, err)
"success": false,
"message": err.Error(),
})
return return
} }
...@@ -202,10 +193,7 @@ func LinuxdoOAuth(c *gin.Context) { ...@@ -202,10 +193,7 @@ func LinuxdoOAuth(c *gin.Context) {
code := c.Query("code") code := c.Query("code")
linuxdoUser, err := getLinuxdoUserInfoByCode(code, c) linuxdoUser, err := getLinuxdoUserInfoByCode(code, c)
if err != nil { if err != nil {
c.JSON(http.StatusOK, gin.H{ common.ApiError(c, err)
"success": false,
"message": err.Error(),
})
return return
} }
......
...@@ -10,14 +10,7 @@ import ( ...@@ -10,14 +10,7 @@ import (
) )
func GetAllLogs(c *gin.Context) { func GetAllLogs(c *gin.Context) {
p, _ := strconv.Atoi(c.Query("p")) pageInfo := common.GetPageQuery(c)
pageSize, _ := strconv.Atoi(c.Query("page_size"))
if p < 1 {
p = 1
}
if pageSize < 0 {
pageSize = common.ItemsPerPage
}
logType, _ := strconv.Atoi(c.Query("type")) logType, _ := strconv.Atoi(c.Query("type"))
startTimestamp, _ := strconv.ParseInt(c.Query("start_timestamp"), 10, 64) startTimestamp, _ := strconv.ParseInt(c.Query("start_timestamp"), 10, 64)
endTimestamp, _ := strconv.ParseInt(c.Query("end_timestamp"), 10, 64) endTimestamp, _ := strconv.ParseInt(c.Query("end_timestamp"), 10, 64)
...@@ -26,38 +19,19 @@ func GetAllLogs(c *gin.Context) { ...@@ -26,38 +19,19 @@ func GetAllLogs(c *gin.Context) {
modelName := c.Query("model_name") modelName := c.Query("model_name")
channel, _ := strconv.Atoi(c.Query("channel")) channel, _ := strconv.Atoi(c.Query("channel"))
group := c.Query("group") group := c.Query("group")
logs, total, err := model.GetAllLogs(logType, startTimestamp, endTimestamp, modelName, username, tokenName, (p-1)*pageSize, pageSize, channel, group) logs, total, err := model.GetAllLogs(logType, startTimestamp, endTimestamp, modelName, username, tokenName, pageInfo.GetStartIdx(), pageInfo.GetPageSize(), channel, group)
if err != nil { if err != nil {
c.JSON(http.StatusOK, gin.H{ common.ApiError(c, err)
"success": false,
"message": err.Error(),
})
return return
} }
c.JSON(http.StatusOK, gin.H{ pageInfo.SetTotal(int(total))
"success": true, pageInfo.SetItems(logs)
"message": "", common.ApiSuccess(c, pageInfo)
"data": map[string]any{ return
"items": logs,
"total": total,
"page": p,
"page_size": pageSize,
},
})
} }
func GetUserLogs(c *gin.Context) { func GetUserLogs(c *gin.Context) {
p, _ := strconv.Atoi(c.Query("p")) pageInfo := common.GetPageQuery(c)
pageSize, _ := strconv.Atoi(c.Query("page_size"))
if p < 1 {
p = 1
}
if pageSize < 0 {
pageSize = common.ItemsPerPage
}
if pageSize > 100 {
pageSize = 100
}
userId := c.GetInt("id") userId := c.GetInt("id")
logType, _ := strconv.Atoi(c.Query("type")) logType, _ := strconv.Atoi(c.Query("type"))
startTimestamp, _ := strconv.ParseInt(c.Query("start_timestamp"), 10, 64) startTimestamp, _ := strconv.ParseInt(c.Query("start_timestamp"), 10, 64)
...@@ -65,24 +39,14 @@ func GetUserLogs(c *gin.Context) { ...@@ -65,24 +39,14 @@ func GetUserLogs(c *gin.Context) {
tokenName := c.Query("token_name") tokenName := c.Query("token_name")
modelName := c.Query("model_name") modelName := c.Query("model_name")
group := c.Query("group") group := c.Query("group")
logs, total, err := model.GetUserLogs(userId, logType, startTimestamp, endTimestamp, modelName, tokenName, (p-1)*pageSize, pageSize, group) logs, total, err := model.GetUserLogs(userId, logType, startTimestamp, endTimestamp, modelName, tokenName, pageInfo.GetStartIdx(), pageInfo.GetPageSize(), group)
if err != nil { if err != nil {
c.JSON(http.StatusOK, gin.H{ common.ApiError(c, err)
"success": false,
"message": err.Error(),
})
return return
} }
c.JSON(http.StatusOK, gin.H{ pageInfo.SetTotal(int(total))
"success": true, pageInfo.SetItems(logs)
"message": "", common.ApiSuccess(c, pageInfo)
"data": map[string]any{
"items": logs,
"total": total,
"page": p,
"page_size": pageSize,
},
})
return return
} }
...@@ -90,10 +54,7 @@ func SearchAllLogs(c *gin.Context) { ...@@ -90,10 +54,7 @@ func SearchAllLogs(c *gin.Context) {
keyword := c.Query("keyword") keyword := c.Query("keyword")
logs, err := model.SearchAllLogs(keyword) logs, err := model.SearchAllLogs(keyword)
if err != nil { if err != nil {
c.JSON(http.StatusOK, gin.H{ common.ApiError(c, err)
"success": false,
"message": err.Error(),
})
return return
} }
c.JSON(http.StatusOK, gin.H{ c.JSON(http.StatusOK, gin.H{
...@@ -109,10 +70,7 @@ func SearchUserLogs(c *gin.Context) { ...@@ -109,10 +70,7 @@ func SearchUserLogs(c *gin.Context) {
userId := c.GetInt("id") userId := c.GetInt("id")
logs, err := model.SearchUserLogs(userId, keyword) logs, err := model.SearchUserLogs(userId, keyword)
if err != nil { if err != nil {
c.JSON(http.StatusOK, gin.H{ common.ApiError(c, err)
"success": false,
"message": err.Error(),
})
return return
} }
c.JSON(http.StatusOK, gin.H{ c.JSON(http.StatusOK, gin.H{
...@@ -198,10 +156,7 @@ func DeleteHistoryLogs(c *gin.Context) { ...@@ -198,10 +156,7 @@ func DeleteHistoryLogs(c *gin.Context) {
} }
count, err := model.DeleteOldLog(c.Request.Context(), targetTimestamp, 100) count, err := model.DeleteOldLog(c.Request.Context(), targetTimestamp, 100)
if err != nil { if err != nil {
c.JSON(http.StatusOK, gin.H{ common.ApiError(c, err)
"success": false,
"message": err.Error(),
})
return return
} }
c.JSON(http.StatusOK, gin.H{ c.JSON(http.StatusOK, gin.H{
......
...@@ -5,7 +5,6 @@ import ( ...@@ -5,7 +5,6 @@ import (
"context" "context"
"encoding/json" "encoding/json"
"fmt" "fmt"
"github.com/gin-gonic/gin"
"io" "io"
"net/http" "net/http"
"one-api/common" "one-api/common"
...@@ -13,8 +12,9 @@ import ( ...@@ -13,8 +12,9 @@ import (
"one-api/model" "one-api/model"
"one-api/service" "one-api/service"
"one-api/setting" "one-api/setting"
"strconv"
"time" "time"
"github.com/gin-gonic/gin"
) )
func UpdateMidjourneyTaskBulk() { func UpdateMidjourneyTaskBulk() {
...@@ -213,14 +213,7 @@ func checkMjTaskNeedUpdate(oldTask *model.Midjourney, newTask dto.MidjourneyDto) ...@@ -213,14 +213,7 @@ func checkMjTaskNeedUpdate(oldTask *model.Midjourney, newTask dto.MidjourneyDto)
} }
func GetAllMidjourney(c *gin.Context) { func GetAllMidjourney(c *gin.Context) {
p, _ := strconv.Atoi(c.Query("p")) pageInfo := common.GetPageQuery(c)
if p < 1 {
p = 1
}
pageSize, _ := strconv.Atoi(c.Query("page_size"))
if pageSize <= 0 {
pageSize = common.ItemsPerPage
}
// 解析其他查询参数 // 解析其他查询参数
queryParams := model.TaskQueryParams{ queryParams := model.TaskQueryParams{
...@@ -230,7 +223,7 @@ func GetAllMidjourney(c *gin.Context) { ...@@ -230,7 +223,7 @@ func GetAllMidjourney(c *gin.Context) {
EndTimestamp: c.Query("end_timestamp"), EndTimestamp: c.Query("end_timestamp"),
} }
items := model.GetAllTasks((p-1)*pageSize, pageSize, queryParams) items := model.GetAllTasks(pageInfo.GetStartIdx(), pageInfo.GetPageSize(), queryParams)
total := model.CountAllTasks(queryParams) total := model.CountAllTasks(queryParams)
if setting.MjForwardUrlEnabled { if setting.MjForwardUrlEnabled {
...@@ -239,27 +232,13 @@ func GetAllMidjourney(c *gin.Context) { ...@@ -239,27 +232,13 @@ func GetAllMidjourney(c *gin.Context) {
items[i] = midjourney items[i] = midjourney
} }
} }
c.JSON(200, gin.H{ pageInfo.SetTotal(int(total))
"success": true, pageInfo.SetItems(items)
"message": "", common.ApiSuccess(c, pageInfo)
"data": gin.H{
"items": items,
"total": total,
"page": p,
"page_size": pageSize,
},
})
} }
func GetUserMidjourney(c *gin.Context) { func GetUserMidjourney(c *gin.Context) {
p, _ := strconv.Atoi(c.Query("p")) pageInfo := common.GetPageQuery(c)
if p < 1 {
p = 1
}
pageSize, _ := strconv.Atoi(c.Query("page_size"))
if pageSize <= 0 {
pageSize = common.ItemsPerPage
}
userId := c.GetInt("id") userId := c.GetInt("id")
...@@ -269,7 +248,7 @@ func GetUserMidjourney(c *gin.Context) { ...@@ -269,7 +248,7 @@ func GetUserMidjourney(c *gin.Context) {
EndTimestamp: c.Query("end_timestamp"), EndTimestamp: c.Query("end_timestamp"),
} }
items := model.GetAllUserTask(userId, (p-1)*pageSize, pageSize, queryParams) items := model.GetAllUserTask(userId, pageInfo.GetStartIdx(), pageInfo.GetPageSize(), queryParams)
total := model.CountAllUserTask(userId, queryParams) total := model.CountAllUserTask(userId, queryParams)
if setting.MjForwardUrlEnabled { if setting.MjForwardUrlEnabled {
...@@ -278,14 +257,7 @@ func GetUserMidjourney(c *gin.Context) { ...@@ -278,14 +257,7 @@ func GetUserMidjourney(c *gin.Context) {
items[i] = midjourney items[i] = midjourney
} }
} }
c.JSON(200, gin.H{ pageInfo.SetTotal(int(total))
"success": true, pageInfo.SetItems(items)
"message": "", common.ApiSuccess(c, pageInfo)
"data": gin.H{
"items": items,
"total": total,
"page": p,
"page_size": pageSize,
},
})
} }
...@@ -217,10 +217,7 @@ func SendEmailVerification(c *gin.Context) { ...@@ -217,10 +217,7 @@ func SendEmailVerification(c *gin.Context) {
"<p>验证码 %d 分钟内有效,如果不是本人操作,请忽略。</p>", common.SystemName, code, common.VerificationValidMinutes) "<p>验证码 %d 分钟内有效,如果不是本人操作,请忽略。</p>", common.SystemName, code, common.VerificationValidMinutes)
err := common.SendEmail(subject, email, content) err := common.SendEmail(subject, email, content)
if err != nil { if err != nil {
c.JSON(http.StatusOK, gin.H{ common.ApiError(c, err)
"success": false,
"message": err.Error(),
})
return return
} }
c.JSON(http.StatusOK, gin.H{ c.JSON(http.StatusOK, gin.H{
...@@ -256,10 +253,7 @@ func SendPasswordResetEmail(c *gin.Context) { ...@@ -256,10 +253,7 @@ func SendPasswordResetEmail(c *gin.Context) {
"<p>重置链接 %d 分钟内有效,如果不是本人操作,请忽略。</p>", common.SystemName, link, link, common.VerificationValidMinutes) "<p>重置链接 %d 分钟内有效,如果不是本人操作,请忽略。</p>", common.SystemName, link, link, common.VerificationValidMinutes)
err := common.SendEmail(subject, email, content) err := common.SendEmail(subject, email, content)
if err != nil { if err != nil {
c.JSON(http.StatusOK, gin.H{ common.ApiError(c, err)
"success": false,
"message": err.Error(),
})
return return
} }
c.JSON(http.StatusOK, gin.H{ c.JSON(http.StatusOK, gin.H{
...@@ -294,10 +288,7 @@ func ResetPassword(c *gin.Context) { ...@@ -294,10 +288,7 @@ func ResetPassword(c *gin.Context) {
password := common.GenerateVerificationCode(12) password := common.GenerateVerificationCode(12)
err = model.ResetUserPasswordByEmail(req.Email, password) err = model.ResetUserPasswordByEmail(req.Email, password)
if err != nil { if err != nil {
c.JSON(http.StatusOK, gin.H{ common.ApiError(c, err)
"success": false,
"message": err.Error(),
})
return return
} }
common.DeleteKey(req.Email, common.PasswordResetPurpose) common.DeleteKey(req.Email, common.PasswordResetPurpose)
......
...@@ -126,10 +126,7 @@ func OidcAuth(c *gin.Context) { ...@@ -126,10 +126,7 @@ func OidcAuth(c *gin.Context) {
code := c.Query("code") code := c.Query("code")
oidcUser, err := getOidcUserInfoByCode(code) oidcUser, err := getOidcUserInfoByCode(code)
if err != nil { if err != nil {
c.JSON(http.StatusOK, gin.H{ common.ApiError(c, err)
"success": false,
"message": err.Error(),
})
return return
} }
user := model.User{ user := model.User{
...@@ -195,10 +192,7 @@ func OidcBind(c *gin.Context) { ...@@ -195,10 +192,7 @@ func OidcBind(c *gin.Context) {
code := c.Query("code") code := c.Query("code")
oidcUser, err := getOidcUserInfoByCode(code) oidcUser, err := getOidcUserInfoByCode(code)
if err != nil { if err != nil {
c.JSON(http.StatusOK, gin.H{ common.ApiError(c, err)
"success": false,
"message": err.Error(),
})
return return
} }
user := model.User{ user := model.User{
...@@ -217,19 +211,13 @@ func OidcBind(c *gin.Context) { ...@@ -217,19 +211,13 @@ func OidcBind(c *gin.Context) {
user.Id = id.(int) user.Id = id.(int)
err = user.FillUserById() err = user.FillUserById()
if err != nil { if err != nil {
c.JSON(http.StatusOK, gin.H{ common.ApiError(c, err)
"success": false,
"message": err.Error(),
})
return return
} }
user.OidcId = oidcUser.OpenID user.OidcId = oidcUser.OpenID
err = user.Update(false) err = user.Update(false)
if err != nil { if err != nil {
c.JSON(http.StatusOK, gin.H{ common.ApiError(c, err)
"success": false,
"message": err.Error(),
})
return return
} }
c.JSON(http.StatusOK, gin.H{ c.JSON(http.StatusOK, gin.H{
......
...@@ -160,10 +160,7 @@ func UpdateOption(c *gin.Context) { ...@@ -160,10 +160,7 @@ func UpdateOption(c *gin.Context) {
} }
err = model.UpdateOption(option.Key, option.Value) err = model.UpdateOption(option.Key, option.Value)
if err != nil { if err != nil {
c.JSON(http.StatusOK, gin.H{ common.ApiError(c, err)
"success": false,
"message": err.Error(),
})
return return
} }
c.JSON(http.StatusOK, gin.H{ c.JSON(http.StatusOK, gin.H{
......
...@@ -3,45 +3,44 @@ package controller ...@@ -3,45 +3,44 @@ package controller
import ( import (
"errors" "errors"
"fmt" "fmt"
"net/http"
"one-api/common" "one-api/common"
"one-api/constant" "one-api/constant"
"one-api/dto" "one-api/dto"
"one-api/middleware" "one-api/middleware"
"one-api/model" "one-api/model"
"one-api/service"
"one-api/setting" "one-api/setting"
"one-api/types"
"time" "time"
"github.com/gin-gonic/gin" "github.com/gin-gonic/gin"
) )
func Playground(c *gin.Context) { func Playground(c *gin.Context) {
var openaiErr *dto.OpenAIErrorWithStatusCode var newAPIError *types.NewAPIError
defer func() { defer func() {
if openaiErr != nil { if newAPIError != nil {
c.JSON(openaiErr.StatusCode, gin.H{ c.JSON(newAPIError.StatusCode, gin.H{
"error": openaiErr.Error, "error": newAPIError.ToOpenAIError(),
}) })
} }
}() }()
useAccessToken := c.GetBool("use_access_token") useAccessToken := c.GetBool("use_access_token")
if useAccessToken { if useAccessToken {
openaiErr = service.OpenAIErrorWrapperLocal(errors.New("暂不支持使用 access token"), "access_token_not_supported", http.StatusBadRequest) newAPIError = types.NewError(errors.New("暂不支持使用 access token"), types.ErrorCodeAccessDenied)
return return
} }
playgroundRequest := &dto.PlayGroundRequest{} playgroundRequest := &dto.PlayGroundRequest{}
err := common.UnmarshalBodyReusable(c, playgroundRequest) err := common.UnmarshalBodyReusable(c, playgroundRequest)
if err != nil { if err != nil {
openaiErr = service.OpenAIErrorWrapperLocal(err, "unmarshal_request_failed", http.StatusBadRequest) newAPIError = types.NewError(err, types.ErrorCodeInvalidRequest)
return return
} }
if playgroundRequest.Model == "" { if playgroundRequest.Model == "" {
openaiErr = service.OpenAIErrorWrapperLocal(errors.New("请选择模型"), "model_required", http.StatusBadRequest) newAPIError = types.NewError(errors.New("请选择模型"), types.ErrorCodeInvalidRequest)
return return
} }
c.Set("original_model", playgroundRequest.Model) c.Set("original_model", playgroundRequest.Model)
...@@ -52,26 +51,32 @@ func Playground(c *gin.Context) { ...@@ -52,26 +51,32 @@ func Playground(c *gin.Context) {
group = userGroup group = userGroup
} else { } else {
if !setting.GroupInUserUsableGroups(group) && group != userGroup { if !setting.GroupInUserUsableGroups(group) && group != userGroup {
openaiErr = service.OpenAIErrorWrapperLocal(errors.New("无权访问该分组"), "group_not_allowed", http.StatusForbidden) newAPIError = types.NewError(errors.New("无权访问该分组"), types.ErrorCodeAccessDenied)
return return
} }
c.Set("group", group) c.Set("group", group)
} }
c.Set("token_name", "playground-"+group)
channel, finalGroup, err := model.CacheGetRandomSatisfiedChannel(c, group, playgroundRequest.Model, 0) userId := c.GetInt("id")
//c.Set("token_name", "playground-"+group)
tempToken := &model.Token{
UserId: userId,
Name: fmt.Sprintf("playground-%s", group),
Group: group,
}
_ = middleware.SetupContextForToken(c, tempToken)
_, err = getChannel(c, group, playgroundRequest.Model, 0)
if err != nil { if err != nil {
message := fmt.Sprintf("当前分组 %s 下对于模型 %s 无可用渠道", finalGroup, playgroundRequest.Model) newAPIError = types.NewError(err, types.ErrorCodeGetChannelFailed)
openaiErr = service.OpenAIErrorWrapperLocal(errors.New(message), "get_playground_channel_failed", http.StatusInternalServerError)
return return
} }
middleware.SetupContextForSelectedChannel(c, channel, playgroundRequest.Model) //middleware.SetupContextForSelectedChannel(c, channel, playgroundRequest.Model)
common.SetContextKey(c, constant.ContextKeyRequestStartTime, time.Now()) common.SetContextKey(c, constant.ContextKeyRequestStartTime, time.Now())
// Write user context to ensure acceptUnsetRatio is available // Write user context to ensure acceptUnsetRatio is available
userId := c.GetInt("id")
userCache, err := model.GetUserCache(userId) userCache, err := model.GetUserCache(userId)
if err != nil { if err != nil {
openaiErr = service.OpenAIErrorWrapperLocal(err, "get_user_cache_failed", http.StatusInternalServerError) newAPIError = types.NewError(err, types.ErrorCodeQueryDataError)
return return
} }
userCache.WriteContext(c) userCache.WriteContext(c)
......
package controller package controller
import ( import (
"errors"
"net/http" "net/http"
"one-api/common" "one-api/common"
"one-api/model" "one-api/model"
"strconv" "strconv"
"errors"
"github.com/gin-gonic/gin" "github.com/gin-gonic/gin"
) )
func GetAllRedemptions(c *gin.Context) { func GetAllRedemptions(c *gin.Context) {
p, _ := strconv.Atoi(c.Query("p")) pageInfo := common.GetPageQuery(c)
pageSize, _ := strconv.Atoi(c.Query("page_size")) redemptions, total, err := model.GetAllRedemptions(pageInfo.GetStartIdx(), pageInfo.GetPageSize())
if p < 0 {
p = 0
}
if pageSize < 1 {
pageSize = common.ItemsPerPage
}
redemptions, total, err := model.GetAllRedemptions((p-1)*pageSize, pageSize)
if err != nil { if err != nil {
c.JSON(http.StatusOK, gin.H{ common.ApiError(c, err)
"success": false,
"message": err.Error(),
})
return return
} }
c.JSON(http.StatusOK, gin.H{ pageInfo.SetTotal(int(total))
"success": true, pageInfo.SetItems(redemptions)
"message": "", common.ApiSuccess(c, pageInfo)
"data": gin.H{
"items": redemptions,
"total": total,
"page": p,
"page_size": pageSize,
},
})
return return
} }
func SearchRedemptions(c *gin.Context) { func SearchRedemptions(c *gin.Context) {
keyword := c.Query("keyword") keyword := c.Query("keyword")
p, _ := strconv.Atoi(c.Query("p")) pageInfo := common.GetPageQuery(c)
pageSize, _ := strconv.Atoi(c.Query("page_size")) redemptions, total, err := model.SearchRedemptions(keyword, pageInfo.GetStartIdx(), pageInfo.GetPageSize())
if p < 0 {
p = 0
}
if pageSize < 1 {
pageSize = common.ItemsPerPage
}
redemptions, total, err := model.SearchRedemptions(keyword, (p-1)*pageSize, pageSize)
if err != nil { if err != nil {
c.JSON(http.StatusOK, gin.H{ common.ApiError(c, err)
"success": false,
"message": err.Error(),
})
return return
} }
c.JSON(http.StatusOK, gin.H{ pageInfo.SetTotal(int(total))
"success": true, pageInfo.SetItems(redemptions)
"message": "", common.ApiSuccess(c, pageInfo)
"data": gin.H{
"items": redemptions,
"total": total,
"page": p,
"page_size": pageSize,
},
})
return return
} }
func GetRedemption(c *gin.Context) { func GetRedemption(c *gin.Context) {
id, err := strconv.Atoi(c.Param("id")) id, err := strconv.Atoi(c.Param("id"))
if err != nil { if err != nil {
c.JSON(http.StatusOK, gin.H{ common.ApiError(c, err)
"success": false,
"message": err.Error(),
})
return return
} }
redemption, err := model.GetRedemptionById(id) redemption, err := model.GetRedemptionById(id)
if err != nil { if err != nil {
c.JSON(http.StatusOK, gin.H{ common.ApiError(c, err)
"success": false,
"message": err.Error(),
})
return return
} }
c.JSON(http.StatusOK, gin.H{ c.JSON(http.StatusOK, gin.H{
...@@ -100,10 +60,7 @@ func AddRedemption(c *gin.Context) { ...@@ -100,10 +60,7 @@ func AddRedemption(c *gin.Context) {
redemption := model.Redemption{} redemption := model.Redemption{}
err := c.ShouldBindJSON(&redemption) err := c.ShouldBindJSON(&redemption)
if err != nil { if err != nil {
c.JSON(http.StatusOK, gin.H{ common.ApiError(c, err)
"success": false,
"message": err.Error(),
})
return return
} }
if len(redemption.Name) == 0 || len(redemption.Name) > 20 { if len(redemption.Name) == 0 || len(redemption.Name) > 20 {
...@@ -165,10 +122,7 @@ func DeleteRedemption(c *gin.Context) { ...@@ -165,10 +122,7 @@ func DeleteRedemption(c *gin.Context) {
id, _ := strconv.Atoi(c.Param("id")) id, _ := strconv.Atoi(c.Param("id"))
err := model.DeleteRedemptionById(id) err := model.DeleteRedemptionById(id)
if err != nil { if err != nil {
c.JSON(http.StatusOK, gin.H{ common.ApiError(c, err)
"success": false,
"message": err.Error(),
})
return return
} }
c.JSON(http.StatusOK, gin.H{ c.JSON(http.StatusOK, gin.H{
...@@ -183,18 +137,12 @@ func UpdateRedemption(c *gin.Context) { ...@@ -183,18 +137,12 @@ func UpdateRedemption(c *gin.Context) {
redemption := model.Redemption{} redemption := model.Redemption{}
err := c.ShouldBindJSON(&redemption) err := c.ShouldBindJSON(&redemption)
if err != nil { if err != nil {
c.JSON(http.StatusOK, gin.H{ common.ApiError(c, err)
"success": false,
"message": err.Error(),
})
return return
} }
cleanRedemption, err := model.GetRedemptionById(redemption.Id) cleanRedemption, err := model.GetRedemptionById(redemption.Id)
if err != nil { if err != nil {
c.JSON(http.StatusOK, gin.H{ common.ApiError(c, err)
"success": false,
"message": err.Error(),
})
return return
} }
if statusOnly == "" { if statusOnly == "" {
...@@ -212,10 +160,7 @@ func UpdateRedemption(c *gin.Context) { ...@@ -212,10 +160,7 @@ func UpdateRedemption(c *gin.Context) {
} }
err = cleanRedemption.Update() err = cleanRedemption.Update()
if err != nil { if err != nil {
c.JSON(http.StatusOK, gin.H{ common.ApiError(c, err)
"success": false,
"message": err.Error(),
})
return return
} }
c.JSON(http.StatusOK, gin.H{ c.JSON(http.StatusOK, gin.H{
...@@ -229,16 +174,13 @@ func UpdateRedemption(c *gin.Context) { ...@@ -229,16 +174,13 @@ func UpdateRedemption(c *gin.Context) {
func DeleteInvalidRedemption(c *gin.Context) { func DeleteInvalidRedemption(c *gin.Context) {
rows, err := model.DeleteInvalidRedemptions() rows, err := model.DeleteInvalidRedemptions()
if err != nil { if err != nil {
c.JSON(http.StatusOK, gin.H{ common.ApiError(c, err)
"success": false,
"message": err.Error(),
})
return return
} }
c.JSON(http.StatusOK, gin.H{ c.JSON(http.StatusOK, gin.H{
"success": true, "success": true,
"message": "", "message": "",
"data": rows, "data": rows,
}) })
return return
} }
......
...@@ -5,8 +5,6 @@ import ( ...@@ -5,8 +5,6 @@ import (
"encoding/json" "encoding/json"
"errors" "errors"
"fmt" "fmt"
"github.com/gin-gonic/gin"
"github.com/samber/lo"
"io" "io"
"net/http" "net/http"
"one-api/common" "one-api/common"
...@@ -17,6 +15,9 @@ import ( ...@@ -17,6 +15,9 @@ import (
"sort" "sort"
"strconv" "strconv"
"time" "time"
"github.com/gin-gonic/gin"
"github.com/samber/lo"
) )
func UpdateTaskBulk() { func UpdateTaskBulk() {
...@@ -225,14 +226,7 @@ func checkTaskNeedUpdate(oldTask *model.Task, newTask dto.SunoDataResponse) bool ...@@ -225,14 +226,7 @@ func checkTaskNeedUpdate(oldTask *model.Task, newTask dto.SunoDataResponse) bool
} }
func GetAllTask(c *gin.Context) { func GetAllTask(c *gin.Context) {
p, _ := strconv.Atoi(c.Query("p")) pageInfo := common.GetPageQuery(c)
if p < 1 {
p = 1
}
pageSize, _ := strconv.Atoi(c.Query("page_size"))
if pageSize <= 0 {
pageSize = common.ItemsPerPage
}
startTimestamp, _ := strconv.ParseInt(c.Query("start_timestamp"), 10, 64) startTimestamp, _ := strconv.ParseInt(c.Query("start_timestamp"), 10, 64)
endTimestamp, _ := strconv.ParseInt(c.Query("end_timestamp"), 10, 64) endTimestamp, _ := strconv.ParseInt(c.Query("end_timestamp"), 10, 64)
...@@ -247,30 +241,15 @@ func GetAllTask(c *gin.Context) { ...@@ -247,30 +241,15 @@ func GetAllTask(c *gin.Context) {
ChannelID: c.Query("channel_id"), ChannelID: c.Query("channel_id"),
} }
items := model.TaskGetAllTasks((p-1)*pageSize, pageSize, queryParams) items := model.TaskGetAllTasks(pageInfo.GetStartIdx(), pageInfo.GetPageSize(), queryParams)
total := model.TaskCountAllTasks(queryParams) total := model.TaskCountAllTasks(queryParams)
pageInfo.SetTotal(int(total))
c.JSON(200, gin.H{ pageInfo.SetItems(items)
"success": true, common.ApiSuccess(c, pageInfo)
"message": "",
"data": gin.H{
"items": items,
"total": total,
"page": p,
"page_size": pageSize,
},
})
} }
func GetUserTask(c *gin.Context) { func GetUserTask(c *gin.Context) {
p, _ := strconv.Atoi(c.Query("p")) pageInfo := common.GetPageQuery(c)
if p < 1 {
p = 1
}
pageSize, _ := strconv.Atoi(c.Query("page_size"))
if pageSize <= 0 {
pageSize = common.ItemsPerPage
}
userId := c.GetInt("id") userId := c.GetInt("id")
...@@ -286,17 +265,9 @@ func GetUserTask(c *gin.Context) { ...@@ -286,17 +265,9 @@ func GetUserTask(c *gin.Context) {
EndTimestamp: endTimestamp, EndTimestamp: endTimestamp,
} }
items := model.TaskGetAllUserTask(userId, (p-1)*pageSize, pageSize, queryParams) items := model.TaskGetAllUserTask(userId, pageInfo.GetStartIdx(), pageInfo.GetPageSize(), queryParams)
total := model.TaskCountAllUserTask(userId, queryParams) total := model.TaskCountAllUserTask(userId, queryParams)
pageInfo.SetTotal(int(total))
c.JSON(200, gin.H{ pageInfo.SetItems(items)
"success": true, common.ApiSuccess(c, pageInfo)
"message": "",
"data": gin.H{
"items": items,
"total": total,
"page": p,
"page_size": pageSize,
},
})
} }
package controller package controller
import ( import (
"github.com/gin-gonic/gin"
"net/http" "net/http"
"one-api/common" "one-api/common"
"one-api/model" "one-api/model"
"strconv" "strconv"
"github.com/gin-gonic/gin"
) )
func GetAllTokens(c *gin.Context) { func GetAllTokens(c *gin.Context) {
userId := c.GetInt("id") userId := c.GetInt("id")
p, _ := strconv.Atoi(c.Query("p")) pageInfo := common.GetPageQuery(c)
size, _ := strconv.Atoi(c.Query("size")) tokens, err := model.GetAllUserTokens(userId, pageInfo.GetStartIdx(), pageInfo.GetPageSize())
if p < 1 {
p = 1
}
if size <= 0 {
size = common.ItemsPerPage
} else if size > 100 {
size = 100
}
tokens, err := model.GetAllUserTokens(userId, (p-1)*size, size)
if err != nil { if err != nil {
c.JSON(http.StatusOK, gin.H{ common.ApiError(c, err)
"success": false,
"message": err.Error(),
})
return return
} }
// Get total count for pagination
total, _ := model.CountUserTokens(userId) total, _ := model.CountUserTokens(userId)
pageInfo.SetTotal(int(total))
c.JSON(http.StatusOK, gin.H{ pageInfo.SetItems(tokens)
"success": true, common.ApiSuccess(c, pageInfo)
"message": "",
"data": gin.H{
"items": tokens,
"total": total,
"page": p,
"page_size": size,
},
})
return return
} }
...@@ -50,10 +30,7 @@ func SearchTokens(c *gin.Context) { ...@@ -50,10 +30,7 @@ func SearchTokens(c *gin.Context) {
token := c.Query("token") token := c.Query("token")
tokens, err := model.SearchUserTokens(userId, keyword, token) tokens, err := model.SearchUserTokens(userId, keyword, token)
if err != nil { if err != nil {
c.JSON(http.StatusOK, gin.H{ common.ApiError(c, err)
"success": false,
"message": err.Error(),
})
return return
} }
c.JSON(http.StatusOK, gin.H{ c.JSON(http.StatusOK, gin.H{
...@@ -68,18 +45,12 @@ func GetToken(c *gin.Context) { ...@@ -68,18 +45,12 @@ func GetToken(c *gin.Context) {
id, err := strconv.Atoi(c.Param("id")) id, err := strconv.Atoi(c.Param("id"))
userId := c.GetInt("id") userId := c.GetInt("id")
if err != nil { if err != nil {
c.JSON(http.StatusOK, gin.H{ common.ApiError(c, err)
"success": false,
"message": err.Error(),
})
return return
} }
token, err := model.GetTokenByIds(id, userId) token, err := model.GetTokenByIds(id, userId)
if err != nil { if err != nil {
c.JSON(http.StatusOK, gin.H{ common.ApiError(c, err)
"success": false,
"message": err.Error(),
})
return return
} }
c.JSON(http.StatusOK, gin.H{ c.JSON(http.StatusOK, gin.H{
...@@ -95,10 +66,7 @@ func GetTokenStatus(c *gin.Context) { ...@@ -95,10 +66,7 @@ func GetTokenStatus(c *gin.Context) {
userId := c.GetInt("id") userId := c.GetInt("id")
token, err := model.GetTokenByIds(tokenId, userId) token, err := model.GetTokenByIds(tokenId, userId)
if err != nil { if err != nil {
c.JSON(http.StatusOK, gin.H{ common.ApiError(c, err)
"success": false,
"message": err.Error(),
})
return return
} }
expiredAt := token.ExpiredTime expiredAt := token.ExpiredTime
...@@ -118,10 +86,7 @@ func AddToken(c *gin.Context) { ...@@ -118,10 +86,7 @@ func AddToken(c *gin.Context) {
token := model.Token{} token := model.Token{}
err := c.ShouldBindJSON(&token) err := c.ShouldBindJSON(&token)
if err != nil { if err != nil {
c.JSON(http.StatusOK, gin.H{ common.ApiError(c, err)
"success": false,
"message": err.Error(),
})
return return
} }
if len(token.Name) > 30 { if len(token.Name) > 30 {
...@@ -156,10 +121,7 @@ func AddToken(c *gin.Context) { ...@@ -156,10 +121,7 @@ func AddToken(c *gin.Context) {
} }
err = cleanToken.Insert() err = cleanToken.Insert()
if err != nil { if err != nil {
c.JSON(http.StatusOK, gin.H{ common.ApiError(c, err)
"success": false,
"message": err.Error(),
})
return return
} }
c.JSON(http.StatusOK, gin.H{ c.JSON(http.StatusOK, gin.H{
...@@ -174,10 +136,7 @@ func DeleteToken(c *gin.Context) { ...@@ -174,10 +136,7 @@ func DeleteToken(c *gin.Context) {
userId := c.GetInt("id") userId := c.GetInt("id")
err := model.DeleteTokenById(id, userId) err := model.DeleteTokenById(id, userId)
if err != nil { if err != nil {
c.JSON(http.StatusOK, gin.H{ common.ApiError(c, err)
"success": false,
"message": err.Error(),
})
return return
} }
c.JSON(http.StatusOK, gin.H{ c.JSON(http.StatusOK, gin.H{
...@@ -193,10 +152,7 @@ func UpdateToken(c *gin.Context) { ...@@ -193,10 +152,7 @@ func UpdateToken(c *gin.Context) {
token := model.Token{} token := model.Token{}
err := c.ShouldBindJSON(&token) err := c.ShouldBindJSON(&token)
if err != nil { if err != nil {
c.JSON(http.StatusOK, gin.H{ common.ApiError(c, err)
"success": false,
"message": err.Error(),
})
return return
} }
if len(token.Name) > 30 { if len(token.Name) > 30 {
...@@ -208,10 +164,7 @@ func UpdateToken(c *gin.Context) { ...@@ -208,10 +164,7 @@ func UpdateToken(c *gin.Context) {
} }
cleanToken, err := model.GetTokenByIds(token.Id, userId) cleanToken, err := model.GetTokenByIds(token.Id, userId)
if err != nil { if err != nil {
c.JSON(http.StatusOK, gin.H{ common.ApiError(c, err)
"success": false,
"message": err.Error(),
})
return return
} }
if token.Status == common.TokenStatusEnabled { if token.Status == common.TokenStatusEnabled {
...@@ -245,10 +198,7 @@ func UpdateToken(c *gin.Context) { ...@@ -245,10 +198,7 @@ func UpdateToken(c *gin.Context) {
} }
err = cleanToken.Update() err = cleanToken.Update()
if err != nil { if err != nil {
c.JSON(http.StatusOK, gin.H{ common.ApiError(c, err)
"success": false,
"message": err.Error(),
})
return return
} }
c.JSON(http.StatusOK, gin.H{ c.JSON(http.StatusOK, gin.H{
...@@ -275,10 +225,7 @@ func DeleteTokenBatch(c *gin.Context) { ...@@ -275,10 +225,7 @@ func DeleteTokenBatch(c *gin.Context) {
userId := c.GetInt("id") userId := c.GetInt("id")
count, err := model.BatchDeleteTokens(tokenBatch.Ids, userId) count, err := model.BatchDeleteTokens(tokenBatch.Ids, userId)
if err != nil { if err != nil {
c.JSON(http.StatusOK, gin.H{ common.ApiError(c, err)
"success": false,
"message": err.Error(),
})
return return
} }
c.JSON(http.StatusOK, gin.H{ c.JSON(http.StatusOK, gin.H{
......
package controller package controller
import ( import (
"github.com/gin-gonic/gin"
"net/http" "net/http"
"one-api/common"
"one-api/model" "one-api/model"
"strconv" "strconv"
"github.com/gin-gonic/gin"
) )
func GetAllQuotaDates(c *gin.Context) { func GetAllQuotaDates(c *gin.Context) {
...@@ -13,10 +15,7 @@ func GetAllQuotaDates(c *gin.Context) { ...@@ -13,10 +15,7 @@ func GetAllQuotaDates(c *gin.Context) {
username := c.Query("username") username := c.Query("username")
dates, err := model.GetAllQuotaDates(startTimestamp, endTimestamp, username) dates, err := model.GetAllQuotaDates(startTimestamp, endTimestamp, username)
if err != nil { if err != nil {
c.JSON(http.StatusOK, gin.H{ common.ApiError(c, err)
"success": false,
"message": err.Error(),
})
return return
} }
c.JSON(http.StatusOK, gin.H{ c.JSON(http.StatusOK, gin.H{
...@@ -41,10 +40,7 @@ func GetUserQuotaDates(c *gin.Context) { ...@@ -41,10 +40,7 @@ func GetUserQuotaDates(c *gin.Context) {
} }
dates, err := model.GetQuotaDataByUserId(userId, startTimestamp, endTimestamp) dates, err := model.GetQuotaDataByUserId(userId, startTimestamp, endTimestamp)
if err != nil { if err != nil {
c.JSON(http.StatusOK, gin.H{ common.ApiError(c, err)
"success": false,
"message": err.Error(),
})
return return
} }
c.JSON(http.StatusOK, gin.H{ c.JSON(http.StatusOK, gin.H{
......
...@@ -188,10 +188,7 @@ func Register(c *gin.Context) { ...@@ -188,10 +188,7 @@ func Register(c *gin.Context) {
cleanUser.Email = user.Email cleanUser.Email = user.Email
} }
if err := cleanUser.Insert(inviterId); err != nil { if err := cleanUser.Insert(inviterId); err != nil {
c.JSON(http.StatusOK, gin.H{ common.ApiError(c, err)
"success": false,
"message": err.Error(),
})
return return
} }
...@@ -247,81 +244,45 @@ func Register(c *gin.Context) { ...@@ -247,81 +244,45 @@ func Register(c *gin.Context) {
} }
func GetAllUsers(c *gin.Context) { func GetAllUsers(c *gin.Context) {
pageInfo, err := common.GetPageQuery(c) pageInfo := common.GetPageQuery(c)
if err != nil {
c.JSON(http.StatusOK, gin.H{
"success": false,
"message": "parse page query failed",
})
return
}
users, total, err := model.GetAllUsers(pageInfo) users, total, err := model.GetAllUsers(pageInfo)
if err != nil { if err != nil {
c.JSON(http.StatusOK, gin.H{ common.ApiError(c, err)
"success": false,
"message": err.Error(),
})
return return
} }
pageInfo.SetTotal(int(total)) pageInfo.SetTotal(int(total))
pageInfo.SetItems(users) pageInfo.SetItems(users)
c.JSON(http.StatusOK, gin.H{
"success": true, common.ApiSuccess(c, pageInfo)
"message": "",
"data": pageInfo,
})
return return
} }
func SearchUsers(c *gin.Context) { func SearchUsers(c *gin.Context) {
keyword := c.Query("keyword") keyword := c.Query("keyword")
group := c.Query("group") group := c.Query("group")
p, _ := strconv.Atoi(c.Query("p")) pageInfo := common.GetPageQuery(c)
pageSize, _ := strconv.Atoi(c.Query("page_size")) users, total, err := model.SearchUsers(keyword, group, pageInfo.GetStartIdx(), pageInfo.GetPageSize())
if p < 1 {
p = 1
}
if pageSize < 0 {
pageSize = common.ItemsPerPage
}
startIdx := (p - 1) * pageSize
users, total, err := model.SearchUsers(keyword, group, startIdx, pageSize)
if err != nil { if err != nil {
c.JSON(http.StatusOK, gin.H{ common.ApiError(c, err)
"success": false,
"message": err.Error(),
})
return return
} }
c.JSON(http.StatusOK, gin.H{
"success": true, pageInfo.SetTotal(int(total))
"message": "", pageInfo.SetItems(users)
"data": gin.H{ common.ApiSuccess(c, pageInfo)
"items": users,
"total": total,
"page": p,
"page_size": pageSize,
},
})
return return
} }
func GetUser(c *gin.Context) { func GetUser(c *gin.Context) {
id, err := strconv.Atoi(c.Param("id")) id, err := strconv.Atoi(c.Param("id"))
if err != nil { if err != nil {
c.JSON(http.StatusOK, gin.H{ common.ApiError(c, err)
"success": false,
"message": err.Error(),
})
return return
} }
user, err := model.GetUserById(id, false) user, err := model.GetUserById(id, false)
if err != nil { if err != nil {
c.JSON(http.StatusOK, gin.H{ common.ApiError(c, err)
"success": false,
"message": err.Error(),
})
return return
} }
myRole := c.GetInt("role") myRole := c.GetInt("role")
...@@ -344,10 +305,7 @@ func GenerateAccessToken(c *gin.Context) { ...@@ -344,10 +305,7 @@ func GenerateAccessToken(c *gin.Context) {
id := c.GetInt("id") id := c.GetInt("id")
user, err := model.GetUserById(id, true) user, err := model.GetUserById(id, true)
if err != nil { if err != nil {
c.JSON(http.StatusOK, gin.H{ common.ApiError(c, err)
"success": false,
"message": err.Error(),
})
return return
} }
// get rand int 28-32 // get rand int 28-32
...@@ -372,10 +330,7 @@ func GenerateAccessToken(c *gin.Context) { ...@@ -372,10 +330,7 @@ func GenerateAccessToken(c *gin.Context) {
} }
if err := user.Update(false); err != nil { if err := user.Update(false); err != nil {
c.JSON(http.StatusOK, gin.H{ common.ApiError(c, err)
"success": false,
"message": err.Error(),
})
return return
} }
...@@ -395,18 +350,12 @@ func TransferAffQuota(c *gin.Context) { ...@@ -395,18 +350,12 @@ func TransferAffQuota(c *gin.Context) {
id := c.GetInt("id") id := c.GetInt("id")
user, err := model.GetUserById(id, true) user, err := model.GetUserById(id, true)
if err != nil { if err != nil {
c.JSON(http.StatusOK, gin.H{ common.ApiError(c, err)
"success": false,
"message": err.Error(),
})
return return
} }
tran := TransferAffQuotaRequest{} tran := TransferAffQuotaRequest{}
if err := c.ShouldBindJSON(&tran); err != nil { if err := c.ShouldBindJSON(&tran); err != nil {
c.JSON(http.StatusOK, gin.H{ common.ApiError(c, err)
"success": false,
"message": err.Error(),
})
return return
} }
err = user.TransferAffQuotaToQuota(tran.Quota) err = user.TransferAffQuotaToQuota(tran.Quota)
...@@ -427,10 +376,7 @@ func GetAffCode(c *gin.Context) { ...@@ -427,10 +376,7 @@ func GetAffCode(c *gin.Context) {
id := c.GetInt("id") id := c.GetInt("id")
user, err := model.GetUserById(id, true) user, err := model.GetUserById(id, true)
if err != nil { if err != nil {
c.JSON(http.StatusOK, gin.H{ common.ApiError(c, err)
"success": false,
"message": err.Error(),
})
return return
} }
if user.AffCode == "" { if user.AffCode == "" {
...@@ -455,10 +401,7 @@ func GetSelf(c *gin.Context) { ...@@ -455,10 +401,7 @@ func GetSelf(c *gin.Context) {
id := c.GetInt("id") id := c.GetInt("id")
user, err := model.GetUserById(id, false) user, err := model.GetUserById(id, false)
if err != nil { if err != nil {
c.JSON(http.StatusOK, gin.H{ common.ApiError(c, err)
"success": false,
"message": err.Error(),
})
return return
} }
// Hide admin remarks: set to empty to trigger omitempty tag, ensuring the remark field is not included in JSON returned to regular users // Hide admin remarks: set to empty to trigger omitempty tag, ensuring the remark field is not included in JSON returned to regular users
...@@ -479,10 +422,7 @@ func GetUserModels(c *gin.Context) { ...@@ -479,10 +422,7 @@ func GetUserModels(c *gin.Context) {
} }
user, err := model.GetUserCache(id) user, err := model.GetUserCache(id)
if err != nil { if err != nil {
c.JSON(http.StatusOK, gin.H{ common.ApiError(c, err)
"success": false,
"message": err.Error(),
})
return return
} }
groups := setting.GetUserUsableGroups(user.Group) groups := setting.GetUserUsableGroups(user.Group)
...@@ -524,10 +464,7 @@ func UpdateUser(c *gin.Context) { ...@@ -524,10 +464,7 @@ func UpdateUser(c *gin.Context) {
} }
originUser, err := model.GetUserById(updatedUser.Id, false) originUser, err := model.GetUserById(updatedUser.Id, false)
if err != nil { if err != nil {
c.JSON(http.StatusOK, gin.H{ common.ApiError(c, err)
"success": false,
"message": err.Error(),
})
return return
} }
myRole := c.GetInt("role") myRole := c.GetInt("role")
...@@ -550,10 +487,7 @@ func UpdateUser(c *gin.Context) { ...@@ -550,10 +487,7 @@ func UpdateUser(c *gin.Context) {
} }
updatePassword := updatedUser.Password != "" updatePassword := updatedUser.Password != ""
if err := updatedUser.Edit(updatePassword); err != nil { if err := updatedUser.Edit(updatePassword); err != nil {
c.JSON(http.StatusOK, gin.H{ common.ApiError(c, err)
"success": false,
"message": err.Error(),
})
return return
} }
if originUser.Quota != updatedUser.Quota { if originUser.Quota != updatedUser.Quota {
...@@ -599,17 +533,11 @@ func UpdateSelf(c *gin.Context) { ...@@ -599,17 +533,11 @@ func UpdateSelf(c *gin.Context) {
} }
updatePassword, err := checkUpdatePassword(user.OriginalPassword, user.Password, cleanUser.Id) updatePassword, err := checkUpdatePassword(user.OriginalPassword, user.Password, cleanUser.Id)
if err != nil { if err != nil {
c.JSON(http.StatusOK, gin.H{ common.ApiError(c, err)
"success": false,
"message": err.Error(),
})
return return
} }
if err := cleanUser.Update(updatePassword); err != nil { if err := cleanUser.Update(updatePassword); err != nil {
c.JSON(http.StatusOK, gin.H{ common.ApiError(c, err)
"success": false,
"message": err.Error(),
})
return return
} }
...@@ -640,18 +568,12 @@ func checkUpdatePassword(originalPassword string, newPassword string, userId int ...@@ -640,18 +568,12 @@ func checkUpdatePassword(originalPassword string, newPassword string, userId int
func DeleteUser(c *gin.Context) { func DeleteUser(c *gin.Context) {
id, err := strconv.Atoi(c.Param("id")) id, err := strconv.Atoi(c.Param("id"))
if err != nil { if err != nil {
c.JSON(http.StatusOK, gin.H{ common.ApiError(c, err)
"success": false,
"message": err.Error(),
})
return return
} }
originUser, err := model.GetUserById(id, false) originUser, err := model.GetUserById(id, false)
if err != nil { if err != nil {
c.JSON(http.StatusOK, gin.H{ common.ApiError(c, err)
"success": false,
"message": err.Error(),
})
return return
} }
myRole := c.GetInt("role") myRole := c.GetInt("role")
...@@ -686,10 +608,7 @@ func DeleteSelf(c *gin.Context) { ...@@ -686,10 +608,7 @@ func DeleteSelf(c *gin.Context) {
err := model.DeleteUserById(id) err := model.DeleteUserById(id)
if err != nil { if err != nil {
c.JSON(http.StatusOK, gin.H{ common.ApiError(c, err)
"success": false,
"message": err.Error(),
})
return return
} }
c.JSON(http.StatusOK, gin.H{ c.JSON(http.StatusOK, gin.H{
...@@ -735,10 +654,7 @@ func CreateUser(c *gin.Context) { ...@@ -735,10 +654,7 @@ func CreateUser(c *gin.Context) {
DisplayName: user.DisplayName, DisplayName: user.DisplayName,
} }
if err := cleanUser.Insert(0); err != nil { if err := cleanUser.Insert(0); err != nil {
c.JSON(http.StatusOK, gin.H{ common.ApiError(c, err)
"success": false,
"message": err.Error(),
})
return return
} }
...@@ -848,10 +764,7 @@ func ManageUser(c *gin.Context) { ...@@ -848,10 +764,7 @@ func ManageUser(c *gin.Context) {
} }
if err := user.Update(false); err != nil { if err := user.Update(false); err != nil {
c.JSON(http.StatusOK, gin.H{ common.ApiError(c, err)
"success": false,
"message": err.Error(),
})
return return
} }
clearUser := model.User{ clearUser := model.User{
...@@ -883,20 +796,14 @@ func EmailBind(c *gin.Context) { ...@@ -883,20 +796,14 @@ func EmailBind(c *gin.Context) {
} }
err := user.FillUserById() err := user.FillUserById()
if err != nil { if err != nil {
c.JSON(http.StatusOK, gin.H{ common.ApiError(c, err)
"success": false,
"message": err.Error(),
})
return return
} }
user.Email = email user.Email = email
// no need to check if this email already taken, because we have used verification code to check it // no need to check if this email already taken, because we have used verification code to check it
err = user.Update(false) err = user.Update(false)
if err != nil { if err != nil {
c.JSON(http.StatusOK, gin.H{ common.ApiError(c, err)
"success": false,
"message": err.Error(),
})
return return
} }
c.JSON(http.StatusOK, gin.H{ c.JSON(http.StatusOK, gin.H{
...@@ -918,19 +825,13 @@ func TopUp(c *gin.Context) { ...@@ -918,19 +825,13 @@ func TopUp(c *gin.Context) {
req := topUpRequest{} req := topUpRequest{}
err := c.ShouldBindJSON(&req) err := c.ShouldBindJSON(&req)
if err != nil { if err != nil {
c.JSON(http.StatusOK, gin.H{ common.ApiError(c, err)
"success": false,
"message": err.Error(),
})
return return
} }
id := c.GetInt("id") id := c.GetInt("id")
quota, err := model.Redeem(req.Key, id) quota, err := model.Redeem(req.Key, id)
if err != nil { if err != nil {
c.JSON(http.StatusOK, gin.H{ common.ApiError(c, err)
"success": false,
"message": err.Error(),
})
return return
} }
c.JSON(http.StatusOK, gin.H{ c.JSON(http.StatusOK, gin.H{
...@@ -1013,10 +914,7 @@ func UpdateUserSetting(c *gin.Context) { ...@@ -1013,10 +914,7 @@ func UpdateUserSetting(c *gin.Context) {
userId := c.GetInt("id") userId := c.GetInt("id")
user, err := model.GetUserById(userId, true) user, err := model.GetUserById(userId, true)
if err != nil { if err != nil {
c.JSON(http.StatusOK, gin.H{ common.ApiError(c, err)
"success": false,
"message": err.Error(),
})
return return
} }
......
...@@ -4,13 +4,14 @@ import ( ...@@ -4,13 +4,14 @@ import (
"encoding/json" "encoding/json"
"errors" "errors"
"fmt" "fmt"
"github.com/gin-contrib/sessions"
"github.com/gin-gonic/gin"
"net/http" "net/http"
"one-api/common" "one-api/common"
"one-api/model" "one-api/model"
"strconv" "strconv"
"time" "time"
"github.com/gin-contrib/sessions"
"github.com/gin-gonic/gin"
) )
type wechatLoginResponse struct { type wechatLoginResponse struct {
...@@ -150,19 +151,13 @@ func WeChatBind(c *gin.Context) { ...@@ -150,19 +151,13 @@ func WeChatBind(c *gin.Context) {
} }
err = user.FillUserById() err = user.FillUserById()
if err != nil { if err != nil {
c.JSON(http.StatusOK, gin.H{ common.ApiError(c, err)
"success": false,
"message": err.Error(),
})
return return
} }
user.WeChatId = wechatId user.WeChatId = wechatId
err = user.Update(false) err = user.Update(false)
if err != nil { if err != nil {
c.JSON(http.StatusOK, gin.H{ common.ApiError(c, err)
"success": false,
"message": err.Error(),
})
return return
} }
c.JSON(http.StatusOK, gin.H{ c.JSON(http.StatusOK, gin.H{
......
...@@ -3,6 +3,7 @@ package dto ...@@ -3,6 +3,7 @@ package dto
import ( import (
"encoding/json" "encoding/json"
"one-api/common" "one-api/common"
"one-api/types"
) )
type ClaudeMetadata struct { type ClaudeMetadata struct {
...@@ -228,7 +229,7 @@ type ClaudeResponse struct { ...@@ -228,7 +229,7 @@ type ClaudeResponse struct {
Completion string `json:"completion,omitempty"` Completion string `json:"completion,omitempty"`
StopReason string `json:"stop_reason,omitempty"` StopReason string `json:"stop_reason,omitempty"`
Model string `json:"model,omitempty"` Model string `json:"model,omitempty"`
Error *ClaudeError `json:"error,omitempty"` Error *types.ClaudeError `json:"error,omitempty"`
Usage *ClaudeUsage `json:"usage,omitempty"` Usage *ClaudeUsage `json:"usage,omitempty"`
Index *int `json:"index,omitempty"` Index *int `json:"index,omitempty"`
ContentBlock *ClaudeMediaMessage `json:"content_block,omitempty"` ContentBlock *ClaudeMediaMessage `json:"content_block,omitempty"`
......
package dto package dto
import "one-api/types"
type OpenAIError struct { type OpenAIError struct {
Message string `json:"message"` Message string `json:"message"`
Type string `json:"type"` Type string `json:"type"`
...@@ -14,11 +16,11 @@ type OpenAIErrorWithStatusCode struct { ...@@ -14,11 +16,11 @@ type OpenAIErrorWithStatusCode struct {
} }
type GeneralErrorResponse struct { type GeneralErrorResponse struct {
Error OpenAIError `json:"error"` Error types.OpenAIError `json:"error"`
Message string `json:"message"` Message string `json:"message"`
Msg string `json:"msg"` Msg string `json:"msg"`
Err string `json:"err"` Err string `json:"err"`
ErrorMsg string `json:"error_msg"` ErrorMsg string `json:"error_msg"`
Header struct { Header struct {
Message string `json:"message"` Message string `json:"message"`
} `json:"header"` } `json:"header"`
......
...@@ -55,6 +55,7 @@ type GeneralOpenAIRequest struct { ...@@ -55,6 +55,7 @@ type GeneralOpenAIRequest struct {
EnableThinking any `json:"enable_thinking,omitempty"` // ali EnableThinking any `json:"enable_thinking,omitempty"` // ali
THINKING json.RawMessage `json:"thinking,omitempty"` // doubao THINKING json.RawMessage `json:"thinking,omitempty"` // doubao
ExtraBody json.RawMessage `json:"extra_body,omitempty"` ExtraBody json.RawMessage `json:"extra_body,omitempty"`
SearchParameters any `json:"search_parameters,omitempty"` //xai
WebSearchOptions *WebSearchOptions `json:"web_search_options,omitempty"` WebSearchOptions *WebSearchOptions `json:"web_search_options,omitempty"`
// OpenRouter Params // OpenRouter Params
Usage json.RawMessage `json:"usage,omitempty"` Usage json.RawMessage `json:"usage,omitempty"`
...@@ -65,8 +66,8 @@ type GeneralOpenAIRequest struct { ...@@ -65,8 +66,8 @@ type GeneralOpenAIRequest struct {
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.EncodeJson(r) data, _ := common.Marshal(r)
_ = common.UnmarshalJson(data, &result) _ = common.Unmarshal(data, &result)
return result return result
} }
......
package dto package dto
import "encoding/json" import (
"encoding/json"
"one-api/types"
)
type SimpleResponse struct { type SimpleResponse struct {
Usage `json:"usage"` Usage `json:"usage"`
...@@ -28,7 +31,7 @@ type OpenAITextResponse struct { ...@@ -28,7 +31,7 @@ type OpenAITextResponse struct {
Object string `json:"object"` Object string `json:"object"`
Created any `json:"created"` Created any `json:"created"`
Choices []OpenAITextResponseChoice `json:"choices"` Choices []OpenAITextResponseChoice `json:"choices"`
Error *OpenAIError `json:"error,omitempty"` Error *types.OpenAIError `json:"error,omitempty"`
Usage `json:"usage"` Usage `json:"usage"`
} }
...@@ -201,7 +204,7 @@ type OpenAIResponsesResponse struct { ...@@ -201,7 +204,7 @@ type OpenAIResponsesResponse struct {
Object string `json:"object"` Object string `json:"object"`
CreatedAt int `json:"created_at"` CreatedAt int `json:"created_at"`
Status string `json:"status"` Status string `json:"status"`
Error *OpenAIError `json:"error,omitempty"` Error *types.OpenAIError `json:"error,omitempty"`
IncompleteDetails *IncompleteDetails `json:"incomplete_details,omitempty"` IncompleteDetails *IncompleteDetails `json:"incomplete_details,omitempty"`
Instructions string `json:"instructions"` Instructions string `json:"instructions"`
MaxOutputTokens int `json:"max_output_tokens"` MaxOutputTokens int `json:"max_output_tokens"`
......
package dto package dto
import "one-api/types"
const ( const (
RealtimeEventTypeError = "error" RealtimeEventTypeError = "error"
RealtimeEventTypeSessionUpdate = "session.update" RealtimeEventTypeSessionUpdate = "session.update"
...@@ -23,12 +25,12 @@ type RealtimeEvent struct { ...@@ -23,12 +25,12 @@ type RealtimeEvent struct {
EventId string `json:"event_id"` EventId string `json:"event_id"`
Type string `json:"type"` Type string `json:"type"`
//PreviousItemId string `json:"previous_item_id"` //PreviousItemId string `json:"previous_item_id"`
Session *RealtimeSession `json:"session,omitempty"` Session *RealtimeSession `json:"session,omitempty"`
Item *RealtimeItem `json:"item,omitempty"` Item *RealtimeItem `json:"item,omitempty"`
Error *OpenAIError `json:"error,omitempty"` Error *types.OpenAIError `json:"error,omitempty"`
Response *RealtimeResponse `json:"response,omitempty"` Response *RealtimeResponse `json:"response,omitempty"`
Delta string `json:"delta,omitempty"` Delta string `json:"delta,omitempty"`
Audio string `json:"audio,omitempty"` Audio string `json:"audio,omitempty"`
} }
type RealtimeResponse struct { type RealtimeResponse struct {
......
package middleware package middleware
import ( import (
"fmt"
"net/http" "net/http"
"one-api/common" "one-api/common"
"one-api/model" "one-api/model"
...@@ -233,30 +234,41 @@ func TokenAuth() func(c *gin.Context) { ...@@ -233,30 +234,41 @@ func TokenAuth() func(c *gin.Context) {
userCache.WriteContext(c) userCache.WriteContext(c)
c.Set("id", token.UserId) err = SetupContextForToken(c, token, parts...)
c.Set("token_id", token.Id) if err != nil {
c.Set("token_key", token.Key) return
c.Set("token_name", token.Name)
c.Set("token_unlimited_quota", token.UnlimitedQuota)
if !token.UnlimitedQuota {
c.Set("token_quota", token.RemainQuota)
} }
if token.ModelLimitsEnabled { c.Next()
c.Set("token_model_limit_enabled", true) }
c.Set("token_model_limit", token.GetModelLimitsMap()) }
func SetupContextForToken(c *gin.Context, token *model.Token, parts ...string) error {
if token == nil {
return fmt.Errorf("token is nil")
}
c.Set("id", token.UserId)
c.Set("token_id", token.Id)
c.Set("token_key", token.Key)
c.Set("token_name", token.Name)
c.Set("token_unlimited_quota", token.UnlimitedQuota)
if !token.UnlimitedQuota {
c.Set("token_quota", token.RemainQuota)
}
if token.ModelLimitsEnabled {
c.Set("token_model_limit_enabled", true)
c.Set("token_model_limit", token.GetModelLimitsMap())
} else {
c.Set("token_model_limit_enabled", false)
}
c.Set("allow_ips", token.GetIpLimitsMap())
c.Set("token_group", token.Group)
if len(parts) > 1 {
if model.IsAdmin(token.UserId) {
c.Set("specific_channel_id", parts[1])
} else { } else {
c.Set("token_model_limit_enabled", false) abortWithOpenAiMessage(c, http.StatusForbidden, "普通用户不支持指定渠道")
} return fmt.Errorf("普通用户不支持指定渠道")
c.Set("allow_ips", token.GetIpLimitsMap())
c.Set("token_group", token.Group)
if len(parts) > 1 {
if model.IsAdmin(token.UserId) {
c.Set("specific_channel_id", parts[1])
} else {
abortWithOpenAiMessage(c, http.StatusForbidden, "普通用户不支持指定渠道")
return
}
} }
c.Next()
} }
return nil
} }
...@@ -12,6 +12,7 @@ import ( ...@@ -12,6 +12,7 @@ import (
"one-api/service" "one-api/service"
"one-api/setting" "one-api/setting"
"one-api/setting/ratio_setting" "one-api/setting/ratio_setting"
"one-api/types"
"strconv" "strconv"
"strings" "strings"
"time" "time"
...@@ -21,6 +22,7 @@ import ( ...@@ -21,6 +22,7 @@ import (
type ModelRequest struct { type ModelRequest struct {
Model string `json:"model"` Model string `json:"model"`
Group string `json:"group,omitempty"`
} }
func Distribute() func(c *gin.Context) { func Distribute() func(c *gin.Context) {
...@@ -237,28 +239,47 @@ func getModelRequest(c *gin.Context) (*ModelRequest, bool, error) { ...@@ -237,28 +239,47 @@ func getModelRequest(c *gin.Context) (*ModelRequest, bool, error) {
} }
c.Set("relay_mode", relayMode) c.Set("relay_mode", relayMode)
} }
if strings.HasPrefix(c.Request.URL.Path, "/pg/chat/completions") {
// playground chat completions
err = common.UnmarshalBodyReusable(c, &modelRequest)
if err != nil {
return nil, false, errors.New("无效的请求, " + err.Error())
}
common.SetContextKey(c, constant.ContextKeyTokenGroup, modelRequest.Group)
}
return &modelRequest, shouldSelectChannel, nil return &modelRequest, shouldSelectChannel, nil
} }
func SetupContextForSelectedChannel(c *gin.Context, channel *model.Channel, modelName string) { func SetupContextForSelectedChannel(c *gin.Context, channel *model.Channel, modelName string) *types.NewAPIError {
c.Set("original_model", modelName) // for retry c.Set("original_model", modelName) // for retry
if channel == nil { if channel == nil {
return return types.NewError(errors.New("channel is nil"), types.ErrorCodeGetChannelFailed)
} }
c.Set("channel_id", channel.Id) common.SetContextKey(c, constant.ContextKeyChannelId, channel.Id)
c.Set("channel_name", channel.Name) common.SetContextKey(c, constant.ContextKeyChannelName, channel.Name)
common.SetContextKey(c, constant.ContextKeyChannelType, channel.Type) common.SetContextKey(c, constant.ContextKeyChannelType, channel.Type)
c.Set("channel_create_time", channel.CreatedTime) common.SetContextKey(c, constant.ContextKeyChannelCreateTime, channel.CreatedTime)
common.SetContextKey(c, constant.ContextKeyChannelSetting, channel.GetSetting()) common.SetContextKey(c, constant.ContextKeyChannelSetting, channel.GetSetting())
c.Set("param_override", channel.GetParamOverride()) common.SetContextKey(c, constant.ContextKeyChannelParamOverride, channel.GetParamOverride())
if nil != channel.OpenAIOrganization && "" != *channel.OpenAIOrganization { if nil != channel.OpenAIOrganization && *channel.OpenAIOrganization != "" {
c.Set("channel_organization", *channel.OpenAIOrganization) common.SetContextKey(c, constant.ContextKeyChannelOrganization, *channel.OpenAIOrganization)
} }
c.Set("auto_ban", channel.GetAutoBan()) common.SetContextKey(c, constant.ContextKeyChannelAutoBan, channel.GetAutoBan())
c.Set("model_mapping", channel.GetModelMapping()) common.SetContextKey(c, constant.ContextKeyChannelModelMapping, channel.GetModelMapping())
c.Set("status_code_mapping", channel.GetStatusCodeMapping()) common.SetContextKey(c, constant.ContextKeyChannelStatusCodeMapping, channel.GetStatusCodeMapping())
c.Request.Header.Set("Authorization", fmt.Sprintf("Bearer %s", channel.Key))
common.SetContextKey(c, constant.ContextKeyBaseUrl, channel.GetBaseURL()) key, index, newAPIError := channel.GetNextEnabledKey()
if newAPIError != nil {
return newAPIError
}
if channel.ChannelInfo.IsMultiKey {
common.SetContextKey(c, constant.ContextKeyChannelIsMultiKey, true)
common.SetContextKey(c, constant.ContextKeyChannelMultiKeyIndex, index)
}
// c.Request.Header.Set("Authorization", fmt.Sprintf("Bearer %s", key))
common.SetContextKey(c, constant.ContextKeyChannelKey, key)
common.SetContextKey(c, constant.ContextKeyChannelBaseUrl, channel.GetBaseURL())
// TODO: api_version统一 // TODO: api_version统一
switch channel.Type { switch channel.Type {
case constant.ChannelTypeAzure: case constant.ChannelTypeAzure:
...@@ -278,6 +299,7 @@ func SetupContextForSelectedChannel(c *gin.Context, channel *model.Channel, mode ...@@ -278,6 +299,7 @@ func SetupContextForSelectedChannel(c *gin.Context, channel *model.Channel, mode
case constant.ChannelTypeCoze: case constant.ChannelTypeCoze:
c.Set("bot_id", channel.Other) c.Set("bot_id", channel.Other)
} }
return nil
} }
// extractModelNameFromGeminiPath 从 Gemini API URL 路径中提取模型名 // extractModelNameFromGeminiPath 从 Gemini API URL 路径中提取模型名
......
...@@ -14,8 +14,8 @@ import ( ...@@ -14,8 +14,8 @@ import (
"github.com/gin-gonic/gin" "github.com/gin-gonic/gin"
) )
var group2model2channels map[string]map[string][]*Channel var group2model2channels map[string]map[string][]int // enabled channel
var channelsIDM map[int]*Channel var channelsIDM map[int]*Channel // all channels include disabled
var channelSyncLock sync.RWMutex var channelSyncLock sync.RWMutex
func InitChannelCache() { func InitChannelCache() {
...@@ -24,7 +24,7 @@ func InitChannelCache() { ...@@ -24,7 +24,7 @@ func InitChannelCache() {
} }
newChannelId2channel := make(map[int]*Channel) newChannelId2channel := make(map[int]*Channel)
var channels []*Channel var channels []*Channel
DB.Where("status = ?", common.ChannelStatusEnabled).Find(&channels) DB.Find(&channels)
for _, channel := range channels { for _, channel := range channels {
newChannelId2channel[channel.Id] = channel newChannelId2channel[channel.Id] = channel
} }
...@@ -34,21 +34,22 @@ func InitChannelCache() { ...@@ -34,21 +34,22 @@ func InitChannelCache() {
for _, ability := range abilities { for _, ability := range abilities {
groups[ability.Group] = true groups[ability.Group] = true
} }
newGroup2model2channels := make(map[string]map[string][]*Channel) newGroup2model2channels := make(map[string]map[string][]int)
newChannelsIDM := make(map[int]*Channel)
for group := range groups { for group := range groups {
newGroup2model2channels[group] = make(map[string][]*Channel) newGroup2model2channels[group] = make(map[string][]int)
} }
for _, channel := range channels { for _, channel := range channels {
newChannelsIDM[channel.Id] = channel if channel.Status != common.ChannelStatusEnabled {
continue // skip disabled channels
}
groups := strings.Split(channel.Group, ",") groups := strings.Split(channel.Group, ",")
for _, group := range groups { for _, group := range groups {
models := strings.Split(channel.Models, ",") models := strings.Split(channel.Models, ",")
for _, model := range models { for _, model := range models {
if _, ok := newGroup2model2channels[group][model]; !ok { if _, ok := newGroup2model2channels[group][model]; !ok {
newGroup2model2channels[group][model] = make([]*Channel, 0) newGroup2model2channels[group][model] = make([]int, 0)
} }
newGroup2model2channels[group][model] = append(newGroup2model2channels[group][model], channel) newGroup2model2channels[group][model] = append(newGroup2model2channels[group][model], channel.Id)
} }
} }
} }
...@@ -57,7 +58,7 @@ func InitChannelCache() { ...@@ -57,7 +58,7 @@ func InitChannelCache() {
for group, model2channels := range newGroup2model2channels { for group, model2channels := range newGroup2model2channels {
for model, channels := range model2channels { for model, channels := range model2channels {
sort.Slice(channels, func(i, j int) bool { sort.Slice(channels, func(i, j int) bool {
return channels[i].GetPriority() > channels[j].GetPriority() return newChannelId2channel[channels[i]].GetPriority() > newChannelId2channel[channels[j]].GetPriority()
}) })
newGroup2model2channels[group][model] = channels newGroup2model2channels[group][model] = channels
} }
...@@ -65,7 +66,7 @@ func InitChannelCache() { ...@@ -65,7 +66,7 @@ func InitChannelCache() {
channelSyncLock.Lock() channelSyncLock.Lock()
group2model2channels = newGroup2model2channels group2model2channels = newGroup2model2channels
channelsIDM = newChannelsIDM channelsIDM = newChannelId2channel
channelSyncLock.Unlock() channelSyncLock.Unlock()
common.SysLog("channels synced from database") common.SysLog("channels synced from database")
} }
...@@ -128,16 +129,27 @@ func getRandomSatisfiedChannel(group string, model string, retry int) (*Channel, ...@@ -128,16 +129,27 @@ func getRandomSatisfiedChannel(group string, model string, retry int) (*Channel,
} }
channelSyncLock.RLock() channelSyncLock.RLock()
defer channelSyncLock.RUnlock()
channels := group2model2channels[group][model] channels := group2model2channels[group][model]
channelSyncLock.RUnlock()
if len(channels) == 0 { if len(channels) == 0 {
return nil, errors.New("channel not found") return nil, errors.New("channel not found")
} }
if len(channels) == 1 {
if channel, ok := channelsIDM[channels[0]]; ok {
return channel, nil
}
return nil, fmt.Errorf("数据库一致性错误,渠道# %d 不存在,请联系管理员修复", channels[0])
}
uniquePriorities := make(map[int]bool) uniquePriorities := make(map[int]bool)
for _, channel := range channels { for _, channelId := range channels {
uniquePriorities[int(channel.GetPriority())] = true if channel, ok := channelsIDM[channelId]; ok {
uniquePriorities[int(channel.GetPriority())] = true
} else {
return nil, fmt.Errorf("数据库一致性错误,渠道# %d 不存在,请联系管理员修复", channelId)
}
} }
var sortedUniquePriorities []int var sortedUniquePriorities []int
for priority := range uniquePriorities { for priority := range uniquePriorities {
...@@ -152,9 +164,13 @@ func getRandomSatisfiedChannel(group string, model string, retry int) (*Channel, ...@@ -152,9 +164,13 @@ func getRandomSatisfiedChannel(group string, model string, retry int) (*Channel,
// get the priority for the given retry number // get the priority for the given retry number
var targetChannels []*Channel var targetChannels []*Channel
for _, channel := range channels { for _, channelId := range channels {
if channel.GetPriority() == targetPriority { if channel, ok := channelsIDM[channelId]; ok {
targetChannels = append(targetChannels, channel) if channel.GetPriority() == targetPriority {
targetChannels = append(targetChannels, channel)
}
} else {
return nil, fmt.Errorf("数据库一致性错误,渠道# %d 不存在,请联系管理员修复", channelId)
} }
} }
...@@ -188,11 +204,35 @@ func CacheGetChannel(id int) (*Channel, error) { ...@@ -188,11 +204,35 @@ func CacheGetChannel(id int) (*Channel, error) {
c, ok := channelsIDM[id] c, ok := channelsIDM[id]
if !ok { if !ok {
return nil, errors.New(fmt.Sprintf("当前渠道# %d,已不存在", id)) return nil, fmt.Errorf("渠道# %d,已不存在", id)
}
if c.Status != common.ChannelStatusEnabled {
return nil, fmt.Errorf("渠道# %d,已被禁用", id)
} }
return c, nil return c, nil
} }
func CacheGetChannelInfo(id int) (*ChannelInfo, error) {
if !common.MemoryCacheEnabled {
channel, err := GetChannelById(id, true)
if err != nil {
return nil, err
}
return &channel.ChannelInfo, nil
}
channelSyncLock.RLock()
defer channelSyncLock.RUnlock()
c, ok := channelsIDM[id]
if !ok {
return nil, fmt.Errorf("渠道# %d,已不存在", id)
}
if c.Status != common.ChannelStatusEnabled {
return nil, fmt.Errorf("渠道# %d,已被禁用", id)
}
return &c.ChannelInfo, nil
}
func CacheUpdateChannelStatus(id int, status int) { func CacheUpdateChannelStatus(id int, status int) {
if !common.MemoryCacheEnabled { if !common.MemoryCacheEnabled {
return return
...@@ -203,3 +243,20 @@ func CacheUpdateChannelStatus(id int, status int) { ...@@ -203,3 +243,20 @@ func CacheUpdateChannelStatus(id int, status int) {
channel.Status = status channel.Status = status
} }
} }
func CacheUpdateChannel(channel *Channel) {
if !common.MemoryCacheEnabled {
return
}
channelSyncLock.Lock()
defer channelSyncLock.Unlock()
if channel == nil {
return
}
println("CacheUpdateChannel:", channel.Id, channel.Name, channel.Status, channel.ChannelInfo.MultiKeyPollingIndex)
println("before:", channelsIDM[channel.Id].ChannelInfo.MultiKeyPollingIndex)
channelsIDM[channel.Id] = channel
println("after :", channelsIDM[channel.Id].ChannelInfo.MultiKeyPollingIndex)
}
...@@ -49,7 +49,7 @@ func formatUserLogs(logs []*Log) { ...@@ -49,7 +49,7 @@ func formatUserLogs(logs []*Log) {
for i := range logs { for i := range logs {
logs[i].ChannelName = "" logs[i].ChannelName = ""
var otherMap map[string]interface{} var otherMap map[string]interface{}
otherMap = common.StrToMap(logs[i].Other) otherMap, _ = common.StrToMap(logs[i].Other)
if otherMap != nil { if otherMap != nil {
// delete admin // delete admin
delete(otherMap, "admin_info") delete(otherMap, "admin_info")
......
...@@ -57,7 +57,7 @@ func initCol() { ...@@ -57,7 +57,7 @@ func initCol() {
} }
} }
// log sql type and database type // log sql type and database type
common.SysLog("Using Log SQL Type: " + common.LogSqlType) //common.SysLog("Using Log SQL Type: " + common.LogSqlType)
} }
var DB *gorm.DB var DB *gorm.DB
...@@ -225,12 +225,6 @@ func InitLogDB() (err error) { ...@@ -225,12 +225,6 @@ func InitLogDB() (err error) {
if !common.IsMasterNode { if !common.IsMasterNode {
return nil return nil
} }
//if common.UsingMySQL {
// _, _ = sqlDB.Exec("DROP INDEX idx_channels_key ON channels;") // TODO: delete this line when most users have upgraded
// _, _ = sqlDB.Exec("ALTER TABLE midjourneys MODIFY action VARCHAR(40);") // TODO: delete this line when most users have upgraded
// _, _ = sqlDB.Exec("ALTER TABLE midjourneys MODIFY progress VARCHAR(30);") // TODO: delete this line when most users have upgraded
// _, _ = sqlDB.Exec("ALTER TABLE midjourneys MODIFY status VARCHAR(20);") // TODO: delete this line when most users have upgraded
//}
common.SysLog("database migration started") common.SysLog("database migration started")
err = migrateLOGDB() err = migrateLOGDB()
return err return err
......
package model package model
import ( import (
"encoding/json"
"fmt" "fmt"
"one-api/common" "one-api/common"
"one-api/constant" "one-api/constant"
...@@ -36,7 +35,7 @@ func (user *UserBase) WriteContext(c *gin.Context) { ...@@ -36,7 +35,7 @@ func (user *UserBase) WriteContext(c *gin.Context) {
func (user *UserBase) GetSetting() dto.UserSetting { func (user *UserBase) GetSetting() dto.UserSetting {
setting := dto.UserSetting{} setting := dto.UserSetting{}
if user.Setting != "" { if user.Setting != "" {
err := json.Unmarshal([]byte(user.Setting), &setting) err := common.Unmarshal([]byte(user.Setting), &setting)
if err != nil { if err != nil {
common.SysError("failed to unmarshal setting: " + err.Error()) common.SysError("failed to unmarshal setting: " + err.Error())
} }
......
...@@ -3,7 +3,6 @@ package relay ...@@ -3,7 +3,6 @@ package relay
import ( import (
"errors" "errors"
"fmt" "fmt"
"github.com/gin-gonic/gin"
"net/http" "net/http"
"one-api/common" "one-api/common"
"one-api/dto" "one-api/dto"
...@@ -12,7 +11,10 @@ import ( ...@@ -12,7 +11,10 @@ import (
"one-api/relay/helper" "one-api/relay/helper"
"one-api/service" "one-api/service"
"one-api/setting" "one-api/setting"
"one-api/types"
"strings" "strings"
"github.com/gin-gonic/gin"
) )
func getAndValidAudioRequest(c *gin.Context, info *relaycommon.RelayInfo) (*dto.AudioRequest, error) { func getAndValidAudioRequest(c *gin.Context, info *relaycommon.RelayInfo) (*dto.AudioRequest, error) {
...@@ -54,13 +56,13 @@ func getAndValidAudioRequest(c *gin.Context, info *relaycommon.RelayInfo) (*dto. ...@@ -54,13 +56,13 @@ func getAndValidAudioRequest(c *gin.Context, info *relaycommon.RelayInfo) (*dto.
return audioRequest, nil return audioRequest, nil
} }
func AudioHelper(c *gin.Context) (openaiErr *dto.OpenAIErrorWithStatusCode) { func AudioHelper(c *gin.Context) (newAPIError *types.NewAPIError) {
relayInfo := relaycommon.GenRelayInfoOpenAIAudio(c) relayInfo := relaycommon.GenRelayInfoOpenAIAudio(c)
audioRequest, err := getAndValidAudioRequest(c, relayInfo) audioRequest, err := getAndValidAudioRequest(c, relayInfo)
if err != nil { if err != nil {
common.LogError(c, fmt.Sprintf("getAndValidAudioRequest failed: %s", err.Error())) common.LogError(c, fmt.Sprintf("getAndValidAudioRequest failed: %s", err.Error()))
return service.OpenAIErrorWrapper(err, "invalid_audio_request", http.StatusBadRequest) return types.NewError(err, types.ErrorCodeInvalidRequest)
} }
promptTokens := 0 promptTokens := 0
...@@ -73,7 +75,7 @@ func AudioHelper(c *gin.Context) (openaiErr *dto.OpenAIErrorWithStatusCode) { ...@@ -73,7 +75,7 @@ func AudioHelper(c *gin.Context) (openaiErr *dto.OpenAIErrorWithStatusCode) {
priceData, err := helper.ModelPriceHelper(c, relayInfo, preConsumedTokens, 0) priceData, err := helper.ModelPriceHelper(c, relayInfo, preConsumedTokens, 0)
if err != nil { if err != nil {
return service.OpenAIErrorWrapperLocal(err, "model_price_error", http.StatusInternalServerError) return types.NewError(err, types.ErrorCodeModelPriceError)
} }
preConsumedQuota, userQuota, openaiErr := preConsumeQuota(c, priceData.ShouldPreConsumedQuota, relayInfo) preConsumedQuota, userQuota, openaiErr := preConsumeQuota(c, priceData.ShouldPreConsumedQuota, relayInfo)
...@@ -88,23 +90,23 @@ func AudioHelper(c *gin.Context) (openaiErr *dto.OpenAIErrorWithStatusCode) { ...@@ -88,23 +90,23 @@ func AudioHelper(c *gin.Context) (openaiErr *dto.OpenAIErrorWithStatusCode) {
err = helper.ModelMappedHelper(c, relayInfo, audioRequest) err = helper.ModelMappedHelper(c, relayInfo, audioRequest)
if err != nil { if err != nil {
return service.OpenAIErrorWrapperLocal(err, "model_mapped_error", http.StatusInternalServerError) return types.NewError(err, types.ErrorCodeChannelModelMappedError)
} }
adaptor := GetAdaptor(relayInfo.ApiType) adaptor := GetAdaptor(relayInfo.ApiType)
if adaptor == nil { if adaptor == nil {
return service.OpenAIErrorWrapperLocal(fmt.Errorf("invalid api type: %d", relayInfo.ApiType), "invalid_api_type", http.StatusBadRequest) return types.NewError(fmt.Errorf("invalid api type: %d", relayInfo.ApiType), types.ErrorCodeInvalidApiType)
} }
adaptor.Init(relayInfo) adaptor.Init(relayInfo)
ioReader, err := adaptor.ConvertAudioRequest(c, relayInfo, *audioRequest) ioReader, err := adaptor.ConvertAudioRequest(c, relayInfo, *audioRequest)
if err != nil { if err != nil {
return service.OpenAIErrorWrapperLocal(err, "convert_request_failed", http.StatusInternalServerError) return types.NewError(err, types.ErrorCodeConvertRequestFailed)
} }
resp, err := adaptor.DoRequest(c, relayInfo, ioReader) resp, err := adaptor.DoRequest(c, relayInfo, ioReader)
if err != nil { if err != nil {
return service.OpenAIErrorWrapper(err, "do_request_failed", http.StatusInternalServerError) return types.NewError(err, types.ErrorCodeDoRequestFailed)
} }
statusCodeMappingStr := c.GetString("status_code_mapping") statusCodeMappingStr := c.GetString("status_code_mapping")
...@@ -112,18 +114,18 @@ func AudioHelper(c *gin.Context) (openaiErr *dto.OpenAIErrorWithStatusCode) { ...@@ -112,18 +114,18 @@ func AudioHelper(c *gin.Context) (openaiErr *dto.OpenAIErrorWithStatusCode) {
if resp != nil { if resp != nil {
httpResp = resp.(*http.Response) httpResp = resp.(*http.Response)
if httpResp.StatusCode != http.StatusOK { if httpResp.StatusCode != http.StatusOK {
openaiErr = service.RelayErrorHandler(httpResp, false) newAPIError = service.RelayErrorHandler(httpResp, false)
// reset status code 重置状态码 // reset status code 重置状态码
service.ResetStatusCode(openaiErr, statusCodeMappingStr) service.ResetStatusCode(newAPIError, statusCodeMappingStr)
return openaiErr return newAPIError
} }
} }
usage, openaiErr := adaptor.DoResponse(c, httpResp, relayInfo) usage, newAPIError := adaptor.DoResponse(c, httpResp, relayInfo)
if openaiErr != nil { if newAPIError != nil {
// reset status code 重置状态码 // reset status code 重置状态码
service.ResetStatusCode(openaiErr, statusCodeMappingStr) service.ResetStatusCode(newAPIError, statusCodeMappingStr)
return openaiErr return newAPIError
} }
postConsumeQuota(c, relayInfo, usage.(*dto.Usage), preConsumedQuota, userQuota, priceData, "") postConsumeQuota(c, relayInfo, usage.(*dto.Usage), preConsumedQuota, userQuota, priceData, "")
......
...@@ -5,6 +5,7 @@ import ( ...@@ -5,6 +5,7 @@ import (
"net/http" "net/http"
"one-api/dto" "one-api/dto"
relaycommon "one-api/relay/common" relaycommon "one-api/relay/common"
"one-api/types"
"github.com/gin-gonic/gin" "github.com/gin-gonic/gin"
) )
...@@ -21,7 +22,7 @@ type Adaptor interface { ...@@ -21,7 +22,7 @@ type Adaptor interface {
ConvertImageRequest(c *gin.Context, info *relaycommon.RelayInfo, request dto.ImageRequest) (any, error) ConvertImageRequest(c *gin.Context, info *relaycommon.RelayInfo, request dto.ImageRequest) (any, error)
ConvertOpenAIResponsesRequest(c *gin.Context, info *relaycommon.RelayInfo, request dto.OpenAIResponsesRequest) (any, error) ConvertOpenAIResponsesRequest(c *gin.Context, info *relaycommon.RelayInfo, request dto.OpenAIResponsesRequest) (any, error)
DoRequest(c *gin.Context, info *relaycommon.RelayInfo, requestBody io.Reader) (any, error) DoRequest(c *gin.Context, info *relaycommon.RelayInfo, requestBody io.Reader) (any, error)
DoResponse(c *gin.Context, resp *http.Response, info *relaycommon.RelayInfo) (usage any, err *dto.OpenAIErrorWithStatusCode) DoResponse(c *gin.Context, resp *http.Response, info *relaycommon.RelayInfo) (usage any, err *types.NewAPIError)
GetModelList() []string GetModelList() []string
GetChannelName() string GetChannelName() string
ConvertClaudeRequest(c *gin.Context, info *relaycommon.RelayInfo, request *dto.ClaudeRequest) (any, error) ConvertClaudeRequest(c *gin.Context, info *relaycommon.RelayInfo, request *dto.ClaudeRequest) (any, error)
......
...@@ -10,6 +10,7 @@ import ( ...@@ -10,6 +10,7 @@ import (
"one-api/relay/channel/openai" "one-api/relay/channel/openai"
relaycommon "one-api/relay/common" relaycommon "one-api/relay/common"
"one-api/relay/constant" "one-api/relay/constant"
"one-api/types"
"github.com/gin-gonic/gin" "github.com/gin-gonic/gin"
) )
...@@ -99,7 +100,7 @@ func (a *Adaptor) DoRequest(c *gin.Context, info *relaycommon.RelayInfo, request ...@@ -99,7 +100,7 @@ func (a *Adaptor) DoRequest(c *gin.Context, info *relaycommon.RelayInfo, request
return channel.DoApiRequest(a, c, info, requestBody) return channel.DoApiRequest(a, c, info, requestBody)
} }
func (a *Adaptor) DoResponse(c *gin.Context, resp *http.Response, info *relaycommon.RelayInfo) (usage any, err *dto.OpenAIErrorWithStatusCode) { func (a *Adaptor) DoResponse(c *gin.Context, resp *http.Response, info *relaycommon.RelayInfo) (usage any, err *types.NewAPIError) {
switch info.RelayMode { switch info.RelayMode {
case constant.RelayModeImagesGenerations: case constant.RelayModeImagesGenerations:
err, usage = aliImageHandler(c, resp, info) err, usage = aliImageHandler(c, resp, info)
...@@ -109,9 +110,9 @@ func (a *Adaptor) DoResponse(c *gin.Context, resp *http.Response, info *relaycom ...@@ -109,9 +110,9 @@ func (a *Adaptor) DoResponse(c *gin.Context, resp *http.Response, info *relaycom
err, usage = RerankHandler(c, resp, info) err, usage = RerankHandler(c, resp, info)
default: default:
if info.IsStream { if info.IsStream {
err, usage = openai.OaiStreamHandler(c, resp, info) usage, err = openai.OaiStreamHandler(c, info, resp)
} else { } else {
err, usage = openai.OpenaiHandler(c, resp, info) usage, err = openai.OpenaiHandler(c, info, resp)
} }
} }
return return
......
...@@ -4,15 +4,17 @@ import ( ...@@ -4,15 +4,17 @@ import (
"encoding/json" "encoding/json"
"errors" "errors"
"fmt" "fmt"
"github.com/gin-gonic/gin"
"io" "io"
"net/http" "net/http"
"one-api/common" "one-api/common"
"one-api/dto" "one-api/dto"
relaycommon "one-api/relay/common" relaycommon "one-api/relay/common"
"one-api/service" "one-api/service"
"one-api/types"
"strings" "strings"
"time" "time"
"github.com/gin-gonic/gin"
) )
func oaiImage2Ali(request dto.ImageRequest) *AliImageRequest { func oaiImage2Ali(request dto.ImageRequest) *AliImageRequest {
...@@ -124,49 +126,46 @@ func responseAli2OpenAIImage(c *gin.Context, response *AliResponse, info *relayc ...@@ -124,49 +126,46 @@ func responseAli2OpenAIImage(c *gin.Context, response *AliResponse, info *relayc
return &imageResponse return &imageResponse
} }
func aliImageHandler(c *gin.Context, resp *http.Response, info *relaycommon.RelayInfo) (*dto.OpenAIErrorWithStatusCode, *dto.Usage) { func aliImageHandler(c *gin.Context, resp *http.Response, info *relaycommon.RelayInfo) (*types.NewAPIError, *dto.Usage) {
responseFormat := c.GetString("response_format") responseFormat := c.GetString("response_format")
var aliTaskResponse AliResponse var aliTaskResponse AliResponse
responseBody, err := io.ReadAll(resp.Body) responseBody, err := io.ReadAll(resp.Body)
if err != nil { if err != nil {
return service.OpenAIErrorWrapper(err, "read_response_body_failed", http.StatusInternalServerError), nil return types.NewError(err, types.ErrorCodeReadResponseBodyFailed), nil
} }
common.CloseResponseBodyGracefully(resp) common.CloseResponseBodyGracefully(resp)
err = json.Unmarshal(responseBody, &aliTaskResponse) err = json.Unmarshal(responseBody, &aliTaskResponse)
if err != nil { if err != nil {
return service.OpenAIErrorWrapper(err, "unmarshal_response_body_failed", http.StatusInternalServerError), nil return types.NewError(err, types.ErrorCodeBadResponseBody), nil
} }
if aliTaskResponse.Message != "" { if aliTaskResponse.Message != "" {
common.LogError(c, "ali_async_task_failed: "+aliTaskResponse.Message) common.LogError(c, "ali_async_task_failed: "+aliTaskResponse.Message)
return service.OpenAIErrorWrapper(errors.New(aliTaskResponse.Message), "ali_async_task_failed", http.StatusInternalServerError), nil return types.NewError(errors.New(aliTaskResponse.Message), types.ErrorCodeBadResponse), nil
} }
aliResponse, _, err := asyncTaskWait(info, aliTaskResponse.Output.TaskId) aliResponse, _, err := asyncTaskWait(info, aliTaskResponse.Output.TaskId)
if err != nil { if err != nil {
return service.OpenAIErrorWrapper(err, "ali_async_task_wait_failed", http.StatusInternalServerError), nil return types.NewError(err, types.ErrorCodeBadResponse), nil
} }
if aliResponse.Output.TaskStatus != "SUCCEEDED" { if aliResponse.Output.TaskStatus != "SUCCEEDED" {
return &dto.OpenAIErrorWithStatusCode{ return types.WithOpenAIError(types.OpenAIError{
Error: dto.OpenAIError{ Message: aliResponse.Output.Message,
Message: aliResponse.Output.Message, Type: "ali_error",
Type: "ali_error", Param: "",
Param: "", Code: aliResponse.Output.Code,
Code: aliResponse.Output.Code, }, resp.StatusCode), nil
},
StatusCode: resp.StatusCode,
}, nil
} }
fullTextResponse := responseAli2OpenAIImage(c, aliResponse, info, responseFormat) fullTextResponse := responseAli2OpenAIImage(c, aliResponse, info, responseFormat)
jsonResponse, err := json.Marshal(fullTextResponse) jsonResponse, err := json.Marshal(fullTextResponse)
if err != nil { if err != nil {
return service.OpenAIErrorWrapper(err, "marshal_response_body_failed", http.StatusInternalServerError), nil return types.NewError(err, types.ErrorCodeBadResponseBody), nil
} }
c.Writer.Header().Set("Content-Type", "application/json") c.Writer.Header().Set("Content-Type", "application/json")
c.Writer.WriteHeader(resp.StatusCode) c.Writer.WriteHeader(resp.StatusCode)
_, err = c.Writer.Write(jsonResponse) c.Writer.Write(jsonResponse)
return nil, nil return nil, &dto.Usage{}
} }
...@@ -7,7 +7,7 @@ import ( ...@@ -7,7 +7,7 @@ import (
"one-api/common" "one-api/common"
"one-api/dto" "one-api/dto"
relaycommon "one-api/relay/common" relaycommon "one-api/relay/common"
"one-api/service" "one-api/types"
"github.com/gin-gonic/gin" "github.com/gin-gonic/gin"
) )
...@@ -31,29 +31,26 @@ func ConvertRerankRequest(request dto.RerankRequest) *AliRerankRequest { ...@@ -31,29 +31,26 @@ func ConvertRerankRequest(request dto.RerankRequest) *AliRerankRequest {
} }
} }
func RerankHandler(c *gin.Context, resp *http.Response, info *relaycommon.RelayInfo) (*dto.OpenAIErrorWithStatusCode, *dto.Usage) { func RerankHandler(c *gin.Context, resp *http.Response, info *relaycommon.RelayInfo) (*types.NewAPIError, *dto.Usage) {
responseBody, err := io.ReadAll(resp.Body) responseBody, err := io.ReadAll(resp.Body)
if err != nil { if err != nil {
return service.OpenAIErrorWrapper(err, "read_response_body_failed", http.StatusInternalServerError), nil return types.NewError(err, types.ErrorCodeReadResponseBodyFailed), nil
} }
common.CloseResponseBodyGracefully(resp) common.CloseResponseBodyGracefully(resp)
var aliResponse AliRerankResponse var aliResponse AliRerankResponse
err = json.Unmarshal(responseBody, &aliResponse) err = json.Unmarshal(responseBody, &aliResponse)
if err != nil { if err != nil {
return service.OpenAIErrorWrapper(err, "unmarshal_response_body_failed", http.StatusInternalServerError), nil return types.NewError(err, types.ErrorCodeBadResponseBody), nil
} }
if aliResponse.Code != "" { if aliResponse.Code != "" {
return &dto.OpenAIErrorWithStatusCode{ return types.WithOpenAIError(types.OpenAIError{
Error: dto.OpenAIError{ Message: aliResponse.Message,
Message: aliResponse.Message, Type: aliResponse.Code,
Type: aliResponse.Code, Param: aliResponse.RequestId,
Param: aliResponse.RequestId, Code: aliResponse.Code,
Code: aliResponse.Code, }, resp.StatusCode), nil
},
StatusCode: resp.StatusCode,
}, nil
} }
usage := dto.Usage{ usage := dto.Usage{
...@@ -68,14 +65,10 @@ func RerankHandler(c *gin.Context, resp *http.Response, info *relaycommon.RelayI ...@@ -68,14 +65,10 @@ func RerankHandler(c *gin.Context, resp *http.Response, info *relaycommon.RelayI
jsonResponse, err := json.Marshal(rerankResponse) jsonResponse, err := json.Marshal(rerankResponse)
if err != nil { if err != nil {
return service.OpenAIErrorWrapper(err, "marshal_response_body_failed", http.StatusInternalServerError), nil return types.NewError(err, types.ErrorCodeBadResponseBody), nil
} }
c.Writer.Header().Set("Content-Type", "application/json") c.Writer.Header().Set("Content-Type", "application/json")
c.Writer.WriteHeader(resp.StatusCode) c.Writer.WriteHeader(resp.StatusCode)
_, err = c.Writer.Write(jsonResponse) c.Writer.Write(jsonResponse)
if err != nil {
return service.OpenAIErrorWrapper(err, "write_response_body_failed", http.StatusInternalServerError), nil
}
return nil, &usage return nil, &usage
} }
...@@ -8,9 +8,10 @@ import ( ...@@ -8,9 +8,10 @@ import (
"one-api/common" "one-api/common"
"one-api/dto" "one-api/dto"
"one-api/relay/helper" "one-api/relay/helper"
"one-api/service"
"strings" "strings"
"one-api/types"
"github.com/gin-gonic/gin" "github.com/gin-gonic/gin"
) )
...@@ -38,11 +39,11 @@ func embeddingRequestOpenAI2Ali(request dto.EmbeddingRequest) *AliEmbeddingReque ...@@ -38,11 +39,11 @@ func embeddingRequestOpenAI2Ali(request dto.EmbeddingRequest) *AliEmbeddingReque
} }
} }
func aliEmbeddingHandler(c *gin.Context, resp *http.Response) (*dto.OpenAIErrorWithStatusCode, *dto.Usage) { func aliEmbeddingHandler(c *gin.Context, resp *http.Response) (*types.NewAPIError, *dto.Usage) {
var fullTextResponse dto.OpenAIEmbeddingResponse var fullTextResponse dto.OpenAIEmbeddingResponse
err := json.NewDecoder(resp.Body).Decode(&fullTextResponse) err := json.NewDecoder(resp.Body).Decode(&fullTextResponse)
if err != nil { if err != nil {
return service.OpenAIErrorWrapper(err, "unmarshal_response_body_failed", http.StatusInternalServerError), nil return types.NewError(err, types.ErrorCodeBadResponseBody), nil
} }
common.CloseResponseBodyGracefully(resp) common.CloseResponseBodyGracefully(resp)
...@@ -53,11 +54,11 @@ func aliEmbeddingHandler(c *gin.Context, resp *http.Response) (*dto.OpenAIErrorW ...@@ -53,11 +54,11 @@ func aliEmbeddingHandler(c *gin.Context, resp *http.Response) (*dto.OpenAIErrorW
} }
jsonResponse, err := json.Marshal(fullTextResponse) jsonResponse, err := json.Marshal(fullTextResponse)
if err != nil { if err != nil {
return service.OpenAIErrorWrapper(err, "marshal_response_body_failed", http.StatusInternalServerError), nil return types.NewError(err, types.ErrorCodeBadResponseBody), nil
} }
c.Writer.Header().Set("Content-Type", "application/json") c.Writer.Header().Set("Content-Type", "application/json")
c.Writer.WriteHeader(resp.StatusCode) c.Writer.WriteHeader(resp.StatusCode)
_, err = c.Writer.Write(jsonResponse) c.Writer.Write(jsonResponse)
return nil, &fullTextResponse.Usage return nil, &fullTextResponse.Usage
} }
...@@ -119,7 +120,7 @@ func streamResponseAli2OpenAI(aliResponse *AliResponse) *dto.ChatCompletionsStre ...@@ -119,7 +120,7 @@ func streamResponseAli2OpenAI(aliResponse *AliResponse) *dto.ChatCompletionsStre
return &response return &response
} }
func aliStreamHandler(c *gin.Context, resp *http.Response) (*dto.OpenAIErrorWithStatusCode, *dto.Usage) { func aliStreamHandler(c *gin.Context, resp *http.Response) (*types.NewAPIError, *dto.Usage) {
var usage dto.Usage var usage dto.Usage
scanner := bufio.NewScanner(resp.Body) scanner := bufio.NewScanner(resp.Body)
scanner.Split(bufio.ScanLines) scanner.Split(bufio.ScanLines)
...@@ -174,32 +175,29 @@ func aliStreamHandler(c *gin.Context, resp *http.Response) (*dto.OpenAIErrorWith ...@@ -174,32 +175,29 @@ func aliStreamHandler(c *gin.Context, resp *http.Response) (*dto.OpenAIErrorWith
return nil, &usage return nil, &usage
} }
func aliHandler(c *gin.Context, resp *http.Response) (*dto.OpenAIErrorWithStatusCode, *dto.Usage) { func aliHandler(c *gin.Context, resp *http.Response) (*types.NewAPIError, *dto.Usage) {
var aliResponse AliResponse var aliResponse AliResponse
responseBody, err := io.ReadAll(resp.Body) responseBody, err := io.ReadAll(resp.Body)
if err != nil { if err != nil {
return service.OpenAIErrorWrapper(err, "read_response_body_failed", http.StatusInternalServerError), nil return types.NewError(err, types.ErrorCodeReadResponseBodyFailed), nil
} }
common.CloseResponseBodyGracefully(resp) common.CloseResponseBodyGracefully(resp)
err = json.Unmarshal(responseBody, &aliResponse) err = json.Unmarshal(responseBody, &aliResponse)
if err != nil { if err != nil {
return service.OpenAIErrorWrapper(err, "unmarshal_response_body_failed", http.StatusInternalServerError), nil return types.NewError(err, types.ErrorCodeBadResponseBody), nil
} }
if aliResponse.Code != "" { if aliResponse.Code != "" {
return &dto.OpenAIErrorWithStatusCode{ return types.WithOpenAIError(types.OpenAIError{
Error: dto.OpenAIError{ Message: aliResponse.Message,
Message: aliResponse.Message, Type: "ali_error",
Type: aliResponse.Code, Param: aliResponse.RequestId,
Param: aliResponse.RequestId, Code: aliResponse.Code,
Code: aliResponse.Code, }, resp.StatusCode), nil
},
StatusCode: resp.StatusCode,
}, nil
} }
fullTextResponse := responseAli2OpenAI(&aliResponse) fullTextResponse := responseAli2OpenAI(&aliResponse)
jsonResponse, err := json.Marshal(fullTextResponse) jsonResponse, err := common.Marshal(fullTextResponse)
if err != nil { if err != nil {
return service.OpenAIErrorWrapper(err, "marshal_response_body_failed", http.StatusInternalServerError), nil return types.NewError(err, types.ErrorCodeBadResponseBody), nil
} }
c.Writer.Header().Set("Content-Type", "application/json") c.Writer.Header().Set("Content-Type", "application/json")
c.Writer.WriteHeader(resp.StatusCode) c.Writer.WriteHeader(resp.StatusCode)
......
...@@ -8,6 +8,7 @@ import ( ...@@ -8,6 +8,7 @@ import (
"one-api/relay/channel/claude" "one-api/relay/channel/claude"
relaycommon "one-api/relay/common" relaycommon "one-api/relay/common"
"one-api/setting/model_setting" "one-api/setting/model_setting"
"one-api/types"
"github.com/gin-gonic/gin" "github.com/gin-gonic/gin"
) )
...@@ -84,7 +85,7 @@ func (a *Adaptor) DoRequest(c *gin.Context, info *relaycommon.RelayInfo, request ...@@ -84,7 +85,7 @@ func (a *Adaptor) DoRequest(c *gin.Context, info *relaycommon.RelayInfo, request
return nil, nil return nil, nil
} }
func (a *Adaptor) DoResponse(c *gin.Context, resp *http.Response, info *relaycommon.RelayInfo) (usage any, err *dto.OpenAIErrorWithStatusCode) { 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 = awsStreamHandler(c, resp, info, a.RequestMode) err, usage = awsStreamHandler(c, resp, info, a.RequestMode)
} else { } else {
......
...@@ -3,19 +3,22 @@ package aws ...@@ -3,19 +3,22 @@ package aws
import ( import (
"encoding/json" "encoding/json"
"fmt" "fmt"
"github.com/gin-gonic/gin"
"github.com/pkg/errors"
"net/http" "net/http"
"one-api/common" "one-api/common"
"one-api/dto" "one-api/dto"
"one-api/relay/channel/claude" "one-api/relay/channel/claude"
relaycommon "one-api/relay/common" relaycommon "one-api/relay/common"
"one-api/relay/helper"
"one-api/types"
"strings" "strings"
"github.com/gin-gonic/gin"
"github.com/pkg/errors"
"github.com/aws/aws-sdk-go-v2/aws" "github.com/aws/aws-sdk-go-v2/aws"
"github.com/aws/aws-sdk-go-v2/credentials" "github.com/aws/aws-sdk-go-v2/credentials"
"github.com/aws/aws-sdk-go-v2/service/bedrockruntime" "github.com/aws/aws-sdk-go-v2/service/bedrockruntime"
"github.com/aws/aws-sdk-go-v2/service/bedrockruntime/types" bedrockruntimeTypes "github.com/aws/aws-sdk-go-v2/service/bedrockruntime/types"
) )
func newAwsClient(c *gin.Context, info *relaycommon.RelayInfo) (*bedrockruntime.Client, error) { func newAwsClient(c *gin.Context, info *relaycommon.RelayInfo) (*bedrockruntime.Client, error) {
...@@ -65,24 +68,21 @@ func awsModelCrossRegion(awsModelId, awsRegionPrefix string) string { ...@@ -65,24 +68,21 @@ func awsModelCrossRegion(awsModelId, awsRegionPrefix string) string {
return modelPrefix + "." + awsModelId return modelPrefix + "." + awsModelId
} }
func awsModelID(requestModel string) (string, error) { func awsModelID(requestModel string) string {
if awsModelID, ok := awsModelIDMap[requestModel]; ok { if awsModelID, ok := awsModelIDMap[requestModel]; ok {
return awsModelID, nil return awsModelID
} }
return requestModel, nil return requestModel
} }
func awsHandler(c *gin.Context, info *relaycommon.RelayInfo, requestMode int) (*dto.OpenAIErrorWithStatusCode, *dto.Usage) { func awsHandler(c *gin.Context, info *relaycommon.RelayInfo, requestMode int) (*types.NewAPIError, *dto.Usage) {
awsCli, err := newAwsClient(c, info) awsCli, err := newAwsClient(c, info)
if err != nil { if err != nil {
return wrapErr(errors.Wrap(err, "newAwsClient")), nil return types.NewError(err, types.ErrorCodeChannelAwsClientError), nil
} }
awsModelId, err := awsModelID(c.GetString("request_model")) awsModelId := awsModelID(c.GetString("request_model"))
if err != nil {
return wrapErr(errors.Wrap(err, "awsModelID")), nil
}
awsRegionPrefix := awsRegionPrefix(awsCli.Options().Region) awsRegionPrefix := awsRegionPrefix(awsCli.Options().Region)
canCrossRegion := awsModelCanCrossRegion(awsModelId, awsRegionPrefix) canCrossRegion := awsModelCanCrossRegion(awsModelId, awsRegionPrefix)
...@@ -98,42 +98,42 @@ func awsHandler(c *gin.Context, info *relaycommon.RelayInfo, requestMode int) (* ...@@ -98,42 +98,42 @@ func awsHandler(c *gin.Context, info *relaycommon.RelayInfo, requestMode int) (*
claudeReq_, ok := c.Get("converted_request") claudeReq_, ok := c.Get("converted_request")
if !ok { if !ok {
return wrapErr(errors.New("request not found")), nil return types.NewError(errors.New("aws claude request not found"), types.ErrorCodeInvalidRequest), nil
} }
claudeReq := claudeReq_.(*dto.ClaudeRequest) claudeReq := claudeReq_.(*dto.ClaudeRequest)
awsClaudeReq := copyRequest(claudeReq) awsClaudeReq := copyRequest(claudeReq)
awsReq.Body, err = json.Marshal(awsClaudeReq) awsReq.Body, err = json.Marshal(awsClaudeReq)
if err != nil { if err != nil {
return wrapErr(errors.Wrap(err, "marshal request")), nil return types.NewError(errors.Wrap(err, "marshal request"), types.ErrorCodeBadResponseBody), nil
} }
awsResp, err := awsCli.InvokeModel(c.Request.Context(), awsReq) awsResp, err := awsCli.InvokeModel(c.Request.Context(), awsReq)
if err != nil { if err != nil {
return wrapErr(errors.Wrap(err, "InvokeModel")), nil return types.NewError(errors.Wrap(err, "InvokeModel"), types.ErrorCodeChannelAwsClientError), nil
} }
claudeInfo := &claude.ClaudeResponseInfo{ claudeInfo := &claude.ClaudeResponseInfo{
ResponseId: fmt.Sprintf("chatcmpl-%s", common.GetUUID()), ResponseId: helper.GetResponseID(c),
Created: common.GetTimestamp(), Created: common.GetTimestamp(),
Model: info.UpstreamModelName, Model: info.UpstreamModelName,
ResponseText: strings.Builder{}, ResponseText: strings.Builder{},
Usage: &dto.Usage{}, Usage: &dto.Usage{},
} }
claude.HandleClaudeResponseData(c, info, claudeInfo, awsResp.Body, RequestModeMessage) handlerErr := claude.HandleClaudeResponseData(c, info, claudeInfo, awsResp.Body, RequestModeMessage)
if handlerErr != nil {
return handlerErr, nil
}
return nil, claudeInfo.Usage return nil, claudeInfo.Usage
} }
func awsStreamHandler(c *gin.Context, resp *http.Response, info *relaycommon.RelayInfo, requestMode int) (*dto.OpenAIErrorWithStatusCode, *dto.Usage) { func awsStreamHandler(c *gin.Context, resp *http.Response, info *relaycommon.RelayInfo, requestMode int) (*types.NewAPIError, *dto.Usage) {
awsCli, err := newAwsClient(c, info) awsCli, err := newAwsClient(c, info)
if err != nil { if err != nil {
return wrapErr(errors.Wrap(err, "newAwsClient")), nil return types.NewError(err, types.ErrorCodeChannelAwsClientError), nil
} }
awsModelId, err := awsModelID(c.GetString("request_model")) awsModelId := awsModelID(c.GetString("request_model"))
if err != nil {
return wrapErr(errors.Wrap(err, "awsModelID")), nil
}
awsRegionPrefix := awsRegionPrefix(awsCli.Options().Region) awsRegionPrefix := awsRegionPrefix(awsCli.Options().Region)
canCrossRegion := awsModelCanCrossRegion(awsModelId, awsRegionPrefix) canCrossRegion := awsModelCanCrossRegion(awsModelId, awsRegionPrefix)
...@@ -149,25 +149,25 @@ func awsStreamHandler(c *gin.Context, resp *http.Response, info *relaycommon.Rel ...@@ -149,25 +149,25 @@ func awsStreamHandler(c *gin.Context, resp *http.Response, info *relaycommon.Rel
claudeReq_, ok := c.Get("converted_request") claudeReq_, ok := c.Get("converted_request")
if !ok { if !ok {
return wrapErr(errors.New("request not found")), nil return types.NewError(errors.New("aws claude request not found"), types.ErrorCodeInvalidRequest), nil
} }
claudeReq := claudeReq_.(*dto.ClaudeRequest) claudeReq := claudeReq_.(*dto.ClaudeRequest)
awsClaudeReq := copyRequest(claudeReq) awsClaudeReq := copyRequest(claudeReq)
awsReq.Body, err = json.Marshal(awsClaudeReq) awsReq.Body, err = json.Marshal(awsClaudeReq)
if err != nil { if err != nil {
return wrapErr(errors.Wrap(err, "marshal request")), nil return types.NewError(errors.Wrap(err, "marshal request"), types.ErrorCodeBadResponseBody), nil
} }
awsResp, err := awsCli.InvokeModelWithResponseStream(c.Request.Context(), awsReq) awsResp, err := awsCli.InvokeModelWithResponseStream(c.Request.Context(), awsReq)
if err != nil { if err != nil {
return wrapErr(errors.Wrap(err, "InvokeModelWithResponseStream")), nil return types.NewError(errors.Wrap(err, "InvokeModelWithResponseStream"), types.ErrorCodeChannelAwsClientError), nil
} }
stream := awsResp.GetStream() stream := awsResp.GetStream()
defer stream.Close() defer stream.Close()
claudeInfo := &claude.ClaudeResponseInfo{ claudeInfo := &claude.ClaudeResponseInfo{
ResponseId: fmt.Sprintf("chatcmpl-%s", common.GetUUID()), ResponseId: helper.GetResponseID(c),
Created: common.GetTimestamp(), Created: common.GetTimestamp(),
Model: info.UpstreamModelName, Model: info.UpstreamModelName,
ResponseText: strings.Builder{}, ResponseText: strings.Builder{},
...@@ -176,18 +176,18 @@ func awsStreamHandler(c *gin.Context, resp *http.Response, info *relaycommon.Rel ...@@ -176,18 +176,18 @@ func awsStreamHandler(c *gin.Context, resp *http.Response, info *relaycommon.Rel
for event := range stream.Events() { for event := range stream.Events() {
switch v := event.(type) { switch v := event.(type) {
case *types.ResponseStreamMemberChunk: case *bedrockruntimeTypes.ResponseStreamMemberChunk:
info.SetFirstResponseTime() info.SetFirstResponseTime()
respErr := claude.HandleStreamResponseData(c, info, claudeInfo, string(v.Value.Bytes), RequestModeMessage) respErr := claude.HandleStreamResponseData(c, info, claudeInfo, string(v.Value.Bytes), RequestModeMessage)
if respErr != nil { if respErr != nil {
return respErr, nil return respErr, nil
} }
case *types.UnknownUnionMember: case *bedrockruntimeTypes.UnknownUnionMember:
fmt.Println("unknown tag:", v.Tag) fmt.Println("unknown tag:", v.Tag)
return wrapErr(errors.New("unknown response type")), nil return types.NewError(errors.New("unknown response type"), types.ErrorCodeInvalidRequest), nil
default: default:
fmt.Println("union is nil or unknown type") fmt.Println("union is nil or unknown type")
return wrapErr(errors.New("nil or unknown response type")), nil return types.NewError(errors.New("nil or unknown response type"), types.ErrorCodeInvalidRequest), nil
} }
} }
......
...@@ -9,6 +9,7 @@ import ( ...@@ -9,6 +9,7 @@ import (
"one-api/relay/channel" "one-api/relay/channel"
relaycommon "one-api/relay/common" relaycommon "one-api/relay/common"
"one-api/relay/constant" "one-api/relay/constant"
"one-api/types"
"strings" "strings"
"github.com/gin-gonic/gin" "github.com/gin-gonic/gin"
...@@ -140,15 +141,15 @@ func (a *Adaptor) DoRequest(c *gin.Context, info *relaycommon.RelayInfo, request ...@@ -140,15 +141,15 @@ func (a *Adaptor) DoRequest(c *gin.Context, info *relaycommon.RelayInfo, request
return channel.DoApiRequest(a, c, info, requestBody) return channel.DoApiRequest(a, c, info, requestBody)
} }
func (a *Adaptor) DoResponse(c *gin.Context, resp *http.Response, info *relaycommon.RelayInfo) (usage any, err *dto.OpenAIErrorWithStatusCode) { 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 = baiduStreamHandler(c, resp) err, usage = baiduStreamHandler(c, info, resp)
} else { } else {
switch info.RelayMode { switch info.RelayMode {
case constant.RelayModeEmbeddings: case constant.RelayModeEmbeddings:
err, usage = baiduEmbeddingHandler(c, resp) err, usage = baiduEmbeddingHandler(c, info, resp)
default: default:
err, usage = baiduHandler(c, resp) err, usage = baiduHandler(c, info, resp)
} }
} }
return return
......
package baidu package baidu
import ( import (
"bufio"
"encoding/json" "encoding/json"
"errors" "errors"
"fmt" "fmt"
"github.com/gin-gonic/gin"
"io" "io"
"net/http" "net/http"
"one-api/common" "one-api/common"
"one-api/constant" "one-api/constant"
"one-api/dto" "one-api/dto"
relaycommon "one-api/relay/common"
"one-api/relay/helper" "one-api/relay/helper"
"one-api/service" "one-api/service"
"one-api/types"
"strings" "strings"
"sync" "sync"
"time" "time"
"github.com/gin-gonic/gin"
) )
// https://cloud.baidu.com/doc/WENXINWORKSHOP/s/flfmc9do2 // https://cloud.baidu.com/doc/WENXINWORKSHOP/s/flfmc9do2
...@@ -110,92 +112,49 @@ func embeddingResponseBaidu2OpenAI(response *BaiduEmbeddingResponse) *dto.OpenAI ...@@ -110,92 +112,49 @@ func embeddingResponseBaidu2OpenAI(response *BaiduEmbeddingResponse) *dto.OpenAI
return &openAIEmbeddingResponse return &openAIEmbeddingResponse
} }
func baiduStreamHandler(c *gin.Context, resp *http.Response) (*dto.OpenAIErrorWithStatusCode, *dto.Usage) { func baiduStreamHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Response) (*types.NewAPIError, *dto.Usage) {
var usage dto.Usage usage := &dto.Usage{}
scanner := bufio.NewScanner(resp.Body) helper.StreamScannerHandler(c, resp, info, func(data string) bool {
scanner.Split(func(data []byte, atEOF bool) (advance int, token []byte, err error) { var baiduResponse BaiduChatStreamResponse
if atEOF && len(data) == 0 { err := common.Unmarshal([]byte(data), &baiduResponse)
return 0, nil, nil if err != nil {
} common.SysError("error unmarshalling stream response: " + err.Error())
if i := strings.Index(string(data), "\n"); i >= 0 { return true
return i + 1, data[0:i], nil
}
if atEOF {
return len(data), data, nil
} }
return 0, nil, nil if baiduResponse.Usage.TotalTokens != 0 {
}) usage.TotalTokens = baiduResponse.Usage.TotalTokens
dataChan := make(chan string) usage.PromptTokens = baiduResponse.Usage.PromptTokens
stopChan := make(chan bool) usage.CompletionTokens = baiduResponse.Usage.TotalTokens - baiduResponse.Usage.PromptTokens
go func() {
for scanner.Scan() {
data := scanner.Text()
if len(data) < 6 { // ignore blank line or wrong format
continue
}
data = data[6:]
dataChan <- data
} }
stopChan <- true response := streamResponseBaidu2OpenAI(&baiduResponse)
}() err = helper.ObjectData(c, response)
helper.SetEventStreamHeaders(c) if err != nil {
c.Stream(func(w io.Writer) bool { common.SysError("error sending stream response: " + err.Error())
select {
case data := <-dataChan:
var baiduResponse BaiduChatStreamResponse
err := json.Unmarshal([]byte(data), &baiduResponse)
if err != nil {
common.SysError("error unmarshalling stream response: " + err.Error())
return true
}
if baiduResponse.Usage.TotalTokens != 0 {
usage.TotalTokens = baiduResponse.Usage.TotalTokens
usage.PromptTokens = baiduResponse.Usage.PromptTokens
usage.CompletionTokens = baiduResponse.Usage.TotalTokens - baiduResponse.Usage.PromptTokens
}
response := streamResponseBaidu2OpenAI(&baiduResponse)
jsonResponse, err := json.Marshal(response)
if err != nil {
common.SysError("error marshalling stream response: " + err.Error())
return true
}
c.Render(-1, common.CustomEvent{Data: "data: " + string(jsonResponse)})
return true
case <-stopChan:
c.Render(-1, common.CustomEvent{Data: "data: [DONE]"})
return false
} }
return true
}) })
common.CloseResponseBodyGracefully(resp) common.CloseResponseBodyGracefully(resp)
return nil, &usage return nil, usage
} }
func baiduHandler(c *gin.Context, resp *http.Response) (*dto.OpenAIErrorWithStatusCode, *dto.Usage) { func baiduHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Response) (*types.NewAPIError, *dto.Usage) {
var baiduResponse BaiduChatResponse var baiduResponse BaiduChatResponse
responseBody, err := io.ReadAll(resp.Body) responseBody, err := io.ReadAll(resp.Body)
if err != nil { if err != nil {
return service.OpenAIErrorWrapper(err, "read_response_body_failed", http.StatusInternalServerError), nil return types.NewError(err, types.ErrorCodeBadResponseBody), nil
} }
common.CloseResponseBodyGracefully(resp) common.CloseResponseBodyGracefully(resp)
err = json.Unmarshal(responseBody, &baiduResponse) err = json.Unmarshal(responseBody, &baiduResponse)
if err != nil { if err != nil {
return service.OpenAIErrorWrapper(err, "unmarshal_response_body_failed", http.StatusInternalServerError), nil return types.NewError(err, types.ErrorCodeBadResponseBody), nil
} }
if baiduResponse.ErrorMsg != "" { if baiduResponse.ErrorMsg != "" {
return &dto.OpenAIErrorWithStatusCode{ return types.NewError(fmt.Errorf(baiduResponse.ErrorMsg), types.ErrorCodeBadResponseBody), nil
Error: dto.OpenAIError{
Message: baiduResponse.ErrorMsg,
Type: "baidu_error",
Param: "",
Code: baiduResponse.ErrorCode,
},
StatusCode: resp.StatusCode,
}, nil
} }
fullTextResponse := responseBaidu2OpenAI(&baiduResponse) fullTextResponse := responseBaidu2OpenAI(&baiduResponse)
jsonResponse, err := json.Marshal(fullTextResponse) jsonResponse, err := json.Marshal(fullTextResponse)
if err != nil { if err != nil {
return service.OpenAIErrorWrapper(err, "marshal_response_body_failed", http.StatusInternalServerError), nil return types.NewError(err, types.ErrorCodeBadResponseBody), nil
} }
c.Writer.Header().Set("Content-Type", "application/json") c.Writer.Header().Set("Content-Type", "application/json")
c.Writer.WriteHeader(resp.StatusCode) c.Writer.WriteHeader(resp.StatusCode)
...@@ -203,32 +162,24 @@ func baiduHandler(c *gin.Context, resp *http.Response) (*dto.OpenAIErrorWithStat ...@@ -203,32 +162,24 @@ func baiduHandler(c *gin.Context, resp *http.Response) (*dto.OpenAIErrorWithStat
return nil, &fullTextResponse.Usage return nil, &fullTextResponse.Usage
} }
func baiduEmbeddingHandler(c *gin.Context, resp *http.Response) (*dto.OpenAIErrorWithStatusCode, *dto.Usage) { func baiduEmbeddingHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Response) (*types.NewAPIError, *dto.Usage) {
var baiduResponse BaiduEmbeddingResponse var baiduResponse BaiduEmbeddingResponse
responseBody, err := io.ReadAll(resp.Body) responseBody, err := io.ReadAll(resp.Body)
if err != nil { if err != nil {
return service.OpenAIErrorWrapper(err, "read_response_body_failed", http.StatusInternalServerError), nil return types.NewError(err, types.ErrorCodeBadResponseBody), nil
} }
common.CloseResponseBodyGracefully(resp) common.CloseResponseBodyGracefully(resp)
err = json.Unmarshal(responseBody, &baiduResponse) err = json.Unmarshal(responseBody, &baiduResponse)
if err != nil { if err != nil {
return service.OpenAIErrorWrapper(err, "unmarshal_response_body_failed", http.StatusInternalServerError), nil return types.NewError(err, types.ErrorCodeBadResponseBody), nil
} }
if baiduResponse.ErrorMsg != "" { if baiduResponse.ErrorMsg != "" {
return &dto.OpenAIErrorWithStatusCode{ return types.NewError(fmt.Errorf(baiduResponse.ErrorMsg), types.ErrorCodeBadResponseBody), nil
Error: dto.OpenAIError{
Message: baiduResponse.ErrorMsg,
Type: "baidu_error",
Param: "",
Code: baiduResponse.ErrorCode,
},
StatusCode: resp.StatusCode,
}, nil
} }
fullTextResponse := embeddingResponseBaidu2OpenAI(&baiduResponse) fullTextResponse := embeddingResponseBaidu2OpenAI(&baiduResponse)
jsonResponse, err := json.Marshal(fullTextResponse) jsonResponse, err := json.Marshal(fullTextResponse)
if err != nil { if err != nil {
return service.OpenAIErrorWrapper(err, "marshal_response_body_failed", http.StatusInternalServerError), nil return types.NewError(err, types.ErrorCodeBadResponseBody), nil
} }
c.Writer.Header().Set("Content-Type", "application/json") c.Writer.Header().Set("Content-Type", "application/json")
c.Writer.WriteHeader(resp.StatusCode) c.Writer.WriteHeader(resp.StatusCode)
......
...@@ -9,6 +9,7 @@ import ( ...@@ -9,6 +9,7 @@ import (
"one-api/relay/channel" "one-api/relay/channel"
"one-api/relay/channel/openai" "one-api/relay/channel/openai"
relaycommon "one-api/relay/common" relaycommon "one-api/relay/common"
"one-api/types"
"strings" "strings"
"github.com/gin-gonic/gin" "github.com/gin-gonic/gin"
...@@ -92,11 +93,11 @@ func (a *Adaptor) DoRequest(c *gin.Context, info *relaycommon.RelayInfo, request ...@@ -92,11 +93,11 @@ func (a *Adaptor) DoRequest(c *gin.Context, info *relaycommon.RelayInfo, request
return channel.DoApiRequest(a, c, info, requestBody) return channel.DoApiRequest(a, c, info, requestBody)
} }
func (a *Adaptor) DoResponse(c *gin.Context, resp *http.Response, info *relaycommon.RelayInfo) (usage any, err *dto.OpenAIErrorWithStatusCode) { 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 = openai.OaiStreamHandler(c, resp, info) usage, err = openai.OaiStreamHandler(c, info, resp)
} else { } else {
err, usage = openai.OpenaiHandler(c, resp, info) usage, err = openai.OpenaiHandler(c, info, resp)
} }
return return
} }
......
...@@ -9,6 +9,7 @@ import ( ...@@ -9,6 +9,7 @@ import (
"one-api/relay/channel" "one-api/relay/channel"
relaycommon "one-api/relay/common" relaycommon "one-api/relay/common"
"one-api/setting/model_setting" "one-api/setting/model_setting"
"one-api/types"
"strings" "strings"
"github.com/gin-gonic/gin" "github.com/gin-gonic/gin"
...@@ -94,7 +95,7 @@ func (a *Adaptor) DoRequest(c *gin.Context, info *relaycommon.RelayInfo, request ...@@ -94,7 +95,7 @@ func (a *Adaptor) DoRequest(c *gin.Context, info *relaycommon.RelayInfo, request
return channel.DoApiRequest(a, c, info, requestBody) return channel.DoApiRequest(a, c, info, requestBody)
} }
func (a *Adaptor) DoResponse(c *gin.Context, resp *http.Response, info *relaycommon.RelayInfo) (usage any, err *dto.OpenAIErrorWithStatusCode) { 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) err, usage = ClaudeStreamHandler(c, resp, info, a.RequestMode)
} else { } else {
......
...@@ -12,6 +12,7 @@ import ( ...@@ -12,6 +12,7 @@ import (
"one-api/relay/helper" "one-api/relay/helper"
"one-api/service" "one-api/service"
"one-api/setting/model_setting" "one-api/setting/model_setting"
"one-api/types"
"strings" "strings"
"github.com/gin-gonic/gin" "github.com/gin-gonic/gin"
...@@ -125,7 +126,7 @@ func RequestOpenAI2ClaudeMessage(textRequest dto.GeneralOpenAIRequest) (*dto.Cla ...@@ -125,7 +126,7 @@ func RequestOpenAI2ClaudeMessage(textRequest dto.GeneralOpenAIRequest) (*dto.Cla
if textRequest.Reasoning != nil { if textRequest.Reasoning != nil {
var reasoning openrouter.RequestReasoning var reasoning openrouter.RequestReasoning
if err := common.UnmarshalJson(textRequest.Reasoning, &reasoning); err != nil { if err := common.Unmarshal(textRequest.Reasoning, &reasoning); err != nil {
return nil, err return nil, err
} }
...@@ -517,22 +518,15 @@ func FormatClaudeResponseInfo(requestMode int, claudeResponse *dto.ClaudeRespons ...@@ -517,22 +518,15 @@ func FormatClaudeResponseInfo(requestMode int, claudeResponse *dto.ClaudeRespons
return true return true
} }
func HandleStreamResponseData(c *gin.Context, info *relaycommon.RelayInfo, claudeInfo *ClaudeResponseInfo, data string, requestMode int) *dto.OpenAIErrorWithStatusCode { func HandleStreamResponseData(c *gin.Context, info *relaycommon.RelayInfo, claudeInfo *ClaudeResponseInfo, data string, requestMode int) *types.NewAPIError {
var claudeResponse dto.ClaudeResponse var claudeResponse dto.ClaudeResponse
err := common.UnmarshalJsonStr(data, &claudeResponse) err := common.UnmarshalJsonStr(data, &claudeResponse)
if err != nil { if err != nil {
common.SysError("error unmarshalling stream response: " + err.Error()) common.SysError("error unmarshalling stream response: " + err.Error())
return service.OpenAIErrorWrapper(err, "stream_response_error", http.StatusInternalServerError) return types.NewError(err, types.ErrorCodeBadResponseBody)
} }
if claudeResponse.Error != nil && claudeResponse.Error.Type != "" { if claudeResponse.Error != nil && claudeResponse.Error.Type != "" {
return &dto.OpenAIErrorWithStatusCode{ return types.WithClaudeError(*claudeResponse.Error, http.StatusInternalServerError)
Error: dto.OpenAIError{
Code: "stream_response_error",
Type: claudeResponse.Error.Type,
Message: claudeResponse.Error.Message,
},
StatusCode: http.StatusInternalServerError,
}
} }
if info.RelayFormat == relaycommon.RelayFormatClaude { if info.RelayFormat == relaycommon.RelayFormatClaude {
FormatClaudeResponseInfo(requestMode, &claudeResponse, nil, claudeInfo) FormatClaudeResponseInfo(requestMode, &claudeResponse, nil, claudeInfo)
...@@ -593,15 +587,15 @@ func HandleStreamFinalResponse(c *gin.Context, info *relaycommon.RelayInfo, clau ...@@ -593,15 +587,15 @@ func HandleStreamFinalResponse(c *gin.Context, info *relaycommon.RelayInfo, clau
} }
} }
func ClaudeStreamHandler(c *gin.Context, resp *http.Response, info *relaycommon.RelayInfo, requestMode int) (*dto.OpenAIErrorWithStatusCode, *dto.Usage) { func ClaudeStreamHandler(c *gin.Context, resp *http.Response, info *relaycommon.RelayInfo, requestMode int) (*types.NewAPIError, *dto.Usage) {
claudeInfo := &ClaudeResponseInfo{ claudeInfo := &ClaudeResponseInfo{
ResponseId: fmt.Sprintf("chatcmpl-%s", common.GetUUID()), ResponseId: helper.GetResponseID(c),
Created: common.GetTimestamp(), Created: common.GetTimestamp(),
Model: info.UpstreamModelName, Model: info.UpstreamModelName,
ResponseText: strings.Builder{}, ResponseText: strings.Builder{},
Usage: &dto.Usage{}, Usage: &dto.Usage{},
} }
var err *dto.OpenAIErrorWithStatusCode var err *types.NewAPIError
helper.StreamScannerHandler(c, resp, info, func(data string) bool { helper.StreamScannerHandler(c, resp, info, func(data string) bool {
err = HandleStreamResponseData(c, info, claudeInfo, data, requestMode) err = HandleStreamResponseData(c, info, claudeInfo, data, requestMode)
if err != nil { if err != nil {
...@@ -617,21 +611,14 @@ func ClaudeStreamHandler(c *gin.Context, resp *http.Response, info *relaycommon. ...@@ -617,21 +611,14 @@ func ClaudeStreamHandler(c *gin.Context, resp *http.Response, info *relaycommon.
return nil, claudeInfo.Usage return nil, claudeInfo.Usage
} }
func HandleClaudeResponseData(c *gin.Context, info *relaycommon.RelayInfo, claudeInfo *ClaudeResponseInfo, data []byte, requestMode int) *dto.OpenAIErrorWithStatusCode { func HandleClaudeResponseData(c *gin.Context, info *relaycommon.RelayInfo, claudeInfo *ClaudeResponseInfo, data []byte, requestMode int) *types.NewAPIError {
var claudeResponse dto.ClaudeResponse var claudeResponse dto.ClaudeResponse
err := common.UnmarshalJson(data, &claudeResponse) err := common.Unmarshal(data, &claudeResponse)
if err != nil { if err != nil {
return service.OpenAIErrorWrapper(err, "unmarshal_claude_response_failed", http.StatusInternalServerError) return types.NewError(err, types.ErrorCodeBadResponseBody)
} }
if claudeResponse.Error != nil && claudeResponse.Error.Type != "" { if claudeResponse.Error != nil && claudeResponse.Error.Type != "" {
return &dto.OpenAIErrorWithStatusCode{ return types.WithClaudeError(*claudeResponse.Error, http.StatusInternalServerError)
Error: dto.OpenAIError{
Message: claudeResponse.Error.Message,
Type: claudeResponse.Error.Type,
Code: claudeResponse.Error.Type,
},
StatusCode: http.StatusInternalServerError,
}
} }
if requestMode == RequestModeCompletion { if requestMode == RequestModeCompletion {
completionTokens := service.CountTextToken(claudeResponse.Completion, info.OriginModelName) completionTokens := service.CountTextToken(claudeResponse.Completion, info.OriginModelName)
...@@ -652,7 +639,7 @@ func HandleClaudeResponseData(c *gin.Context, info *relaycommon.RelayInfo, claud ...@@ -652,7 +639,7 @@ func HandleClaudeResponseData(c *gin.Context, info *relaycommon.RelayInfo, claud
openaiResponse.Usage = *claudeInfo.Usage openaiResponse.Usage = *claudeInfo.Usage
responseData, err = json.Marshal(openaiResponse) responseData, err = json.Marshal(openaiResponse)
if err != nil { if err != nil {
return service.OpenAIErrorWrapper(err, "marshal_response_body_failed", http.StatusInternalServerError) return types.NewError(err, types.ErrorCodeBadResponseBody)
} }
case relaycommon.RelayFormatClaude: case relaycommon.RelayFormatClaude:
responseData = data responseData = data
...@@ -662,11 +649,11 @@ func HandleClaudeResponseData(c *gin.Context, info *relaycommon.RelayInfo, claud ...@@ -662,11 +649,11 @@ func HandleClaudeResponseData(c *gin.Context, info *relaycommon.RelayInfo, claud
return nil return nil
} }
func ClaudeHandler(c *gin.Context, resp *http.Response, requestMode int, info *relaycommon.RelayInfo) (*dto.OpenAIErrorWithStatusCode, *dto.Usage) { func ClaudeHandler(c *gin.Context, resp *http.Response, requestMode int, info *relaycommon.RelayInfo) (*types.NewAPIError, *dto.Usage) {
defer common.CloseResponseBodyGracefully(resp) defer common.CloseResponseBodyGracefully(resp)
claudeInfo := &ClaudeResponseInfo{ claudeInfo := &ClaudeResponseInfo{
ResponseId: fmt.Sprintf("chatcmpl-%s", common.GetUUID()), ResponseId: helper.GetResponseID(c),
Created: common.GetTimestamp(), Created: common.GetTimestamp(),
Model: info.UpstreamModelName, Model: info.UpstreamModelName,
ResponseText: strings.Builder{}, ResponseText: strings.Builder{},
...@@ -674,7 +661,7 @@ func ClaudeHandler(c *gin.Context, resp *http.Response, requestMode int, info *r ...@@ -674,7 +661,7 @@ func ClaudeHandler(c *gin.Context, resp *http.Response, requestMode int, info *r
} }
responseBody, err := io.ReadAll(resp.Body) responseBody, err := io.ReadAll(resp.Body)
if err != nil { if err != nil {
return service.OpenAIErrorWrapper(err, "read_response_body_failed", http.StatusInternalServerError), nil return types.NewError(err, types.ErrorCodeBadResponseBody), nil
} }
if common.DebugEnabled { if common.DebugEnabled {
println("responseBody: ", string(responseBody)) println("responseBody: ", string(responseBody))
......
...@@ -10,6 +10,7 @@ import ( ...@@ -10,6 +10,7 @@ import (
"one-api/relay/channel" "one-api/relay/channel"
relaycommon "one-api/relay/common" relaycommon "one-api/relay/common"
"one-api/relay/constant" "one-api/relay/constant"
"one-api/types"
"github.com/gin-gonic/gin" "github.com/gin-gonic/gin"
) )
...@@ -94,20 +95,20 @@ func (a *Adaptor) ConvertImageRequest(c *gin.Context, info *relaycommon.RelayInf ...@@ -94,20 +95,20 @@ func (a *Adaptor) ConvertImageRequest(c *gin.Context, info *relaycommon.RelayInf
return nil, errors.New("not implemented") return nil, errors.New("not implemented")
} }
func (a *Adaptor) DoResponse(c *gin.Context, resp *http.Response, info *relaycommon.RelayInfo) (usage any, err *dto.OpenAIErrorWithStatusCode) { func (a *Adaptor) DoResponse(c *gin.Context, resp *http.Response, info *relaycommon.RelayInfo) (usage any, err *types.NewAPIError) {
switch info.RelayMode { switch info.RelayMode {
case constant.RelayModeEmbeddings: case constant.RelayModeEmbeddings:
fallthrough fallthrough
case constant.RelayModeChatCompletions: case constant.RelayModeChatCompletions:
if info.IsStream { if info.IsStream {
err, usage = cfStreamHandler(c, resp, info) err, usage = cfStreamHandler(c, info, resp)
} else { } else {
err, usage = cfHandler(c, resp, info) err, usage = cfHandler(c, info, resp)
} }
case constant.RelayModeAudioTranslation: case constant.RelayModeAudioTranslation:
fallthrough fallthrough
case constant.RelayModeAudioTranscription: case constant.RelayModeAudioTranscription:
err, usage = cfSTTHandler(c, resp, info) err, usage = cfSTTHandler(c, info, resp)
} }
return return
} }
......
...@@ -3,7 +3,6 @@ package cloudflare ...@@ -3,7 +3,6 @@ package cloudflare
import ( import (
"bufio" "bufio"
"encoding/json" "encoding/json"
"github.com/gin-gonic/gin"
"io" "io"
"net/http" "net/http"
"one-api/common" "one-api/common"
...@@ -11,8 +10,11 @@ import ( ...@@ -11,8 +10,11 @@ import (
relaycommon "one-api/relay/common" relaycommon "one-api/relay/common"
"one-api/relay/helper" "one-api/relay/helper"
"one-api/service" "one-api/service"
"one-api/types"
"strings" "strings"
"time" "time"
"github.com/gin-gonic/gin"
) )
func convertCf2CompletionsRequest(textRequest dto.GeneralOpenAIRequest) *CfRequest { func convertCf2CompletionsRequest(textRequest dto.GeneralOpenAIRequest) *CfRequest {
...@@ -25,7 +27,7 @@ func convertCf2CompletionsRequest(textRequest dto.GeneralOpenAIRequest) *CfReque ...@@ -25,7 +27,7 @@ func convertCf2CompletionsRequest(textRequest dto.GeneralOpenAIRequest) *CfReque
} }
} }
func cfStreamHandler(c *gin.Context, resp *http.Response, info *relaycommon.RelayInfo) (*dto.OpenAIErrorWithStatusCode, *dto.Usage) { func cfStreamHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Response) (*types.NewAPIError, *dto.Usage) {
scanner := bufio.NewScanner(resp.Body) scanner := bufio.NewScanner(resp.Body)
scanner.Split(bufio.ScanLines) scanner.Split(bufio.ScanLines)
...@@ -86,16 +88,16 @@ func cfStreamHandler(c *gin.Context, resp *http.Response, info *relaycommon.Rela ...@@ -86,16 +88,16 @@ func cfStreamHandler(c *gin.Context, resp *http.Response, info *relaycommon.Rela
return nil, usage return nil, usage
} }
func cfHandler(c *gin.Context, resp *http.Response, info *relaycommon.RelayInfo) (*dto.OpenAIErrorWithStatusCode, *dto.Usage) { func cfHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Response) (*types.NewAPIError, *dto.Usage) {
responseBody, err := io.ReadAll(resp.Body) responseBody, err := io.ReadAll(resp.Body)
if err != nil { if err != nil {
return service.OpenAIErrorWrapper(err, "read_response_body_failed", http.StatusInternalServerError), nil return types.NewError(err, types.ErrorCodeBadResponseBody), nil
} }
common.CloseResponseBodyGracefully(resp) common.CloseResponseBodyGracefully(resp)
var response dto.TextResponse var response dto.TextResponse
err = json.Unmarshal(responseBody, &response) err = json.Unmarshal(responseBody, &response)
if err != nil { if err != nil {
return service.OpenAIErrorWrapper(err, "unmarshal_response_body_failed", http.StatusInternalServerError), nil return types.NewError(err, types.ErrorCodeBadResponseBody), nil
} }
response.Model = info.UpstreamModelName response.Model = info.UpstreamModelName
var responseText string var responseText string
...@@ -107,7 +109,7 @@ func cfHandler(c *gin.Context, resp *http.Response, info *relaycommon.RelayInfo) ...@@ -107,7 +109,7 @@ func cfHandler(c *gin.Context, resp *http.Response, info *relaycommon.RelayInfo)
response.Id = helper.GetResponseID(c) response.Id = helper.GetResponseID(c)
jsonResponse, err := json.Marshal(response) jsonResponse, err := json.Marshal(response)
if err != nil { if err != nil {
return service.OpenAIErrorWrapper(err, "marshal_response_body_failed", http.StatusInternalServerError), nil return types.NewError(err, types.ErrorCodeBadResponseBody), nil
} }
c.Writer.Header().Set("Content-Type", "application/json") c.Writer.Header().Set("Content-Type", "application/json")
c.Writer.WriteHeader(resp.StatusCode) c.Writer.WriteHeader(resp.StatusCode)
...@@ -115,16 +117,16 @@ func cfHandler(c *gin.Context, resp *http.Response, info *relaycommon.RelayInfo) ...@@ -115,16 +117,16 @@ func cfHandler(c *gin.Context, resp *http.Response, info *relaycommon.RelayInfo)
return nil, usage return nil, usage
} }
func cfSTTHandler(c *gin.Context, resp *http.Response, info *relaycommon.RelayInfo) (*dto.OpenAIErrorWithStatusCode, *dto.Usage) { func cfSTTHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Response) (*types.NewAPIError, *dto.Usage) {
var cfResp CfAudioResponse var cfResp CfAudioResponse
responseBody, err := io.ReadAll(resp.Body) responseBody, err := io.ReadAll(resp.Body)
if err != nil { if err != nil {
return service.OpenAIErrorWrapper(err, "read_response_body_failed", http.StatusInternalServerError), nil return types.NewError(err, types.ErrorCodeBadResponseBody), nil
} }
common.CloseResponseBodyGracefully(resp) common.CloseResponseBodyGracefully(resp)
err = json.Unmarshal(responseBody, &cfResp) err = json.Unmarshal(responseBody, &cfResp)
if err != nil { if err != nil {
return service.OpenAIErrorWrapper(err, "unmarshal_response_body_failed", http.StatusInternalServerError), nil return types.NewError(err, types.ErrorCodeBadResponseBody), nil
} }
audioResp := &dto.AudioResponse{ audioResp := &dto.AudioResponse{
...@@ -133,7 +135,7 @@ func cfSTTHandler(c *gin.Context, resp *http.Response, info *relaycommon.RelayIn ...@@ -133,7 +135,7 @@ func cfSTTHandler(c *gin.Context, resp *http.Response, info *relaycommon.RelayIn
jsonResponse, err := json.Marshal(audioResp) jsonResponse, err := json.Marshal(audioResp)
if err != nil { if err != nil {
return service.OpenAIErrorWrapper(err, "marshal_response_body_failed", http.StatusInternalServerError), nil return types.NewError(err, types.ErrorCodeBadResponseBody), nil
} }
c.Writer.Header().Set("Content-Type", "application/json") c.Writer.Header().Set("Content-Type", "application/json")
c.Writer.WriteHeader(resp.StatusCode) c.Writer.WriteHeader(resp.StatusCode)
......
...@@ -9,6 +9,7 @@ import ( ...@@ -9,6 +9,7 @@ import (
"one-api/relay/channel" "one-api/relay/channel"
relaycommon "one-api/relay/common" relaycommon "one-api/relay/common"
"one-api/relay/constant" "one-api/relay/constant"
"one-api/types"
"github.com/gin-gonic/gin" "github.com/gin-gonic/gin"
) )
...@@ -71,14 +72,14 @@ func (a *Adaptor) ConvertEmbeddingRequest(c *gin.Context, info *relaycommon.Rela ...@@ -71,14 +72,14 @@ func (a *Adaptor) ConvertEmbeddingRequest(c *gin.Context, info *relaycommon.Rela
return nil, errors.New("not implemented") return nil, errors.New("not implemented")
} }
func (a *Adaptor) DoResponse(c *gin.Context, resp *http.Response, info *relaycommon.RelayInfo) (usage any, err *dto.OpenAIErrorWithStatusCode) { func (a *Adaptor) DoResponse(c *gin.Context, resp *http.Response, info *relaycommon.RelayInfo) (usage any, err *types.NewAPIError) {
if info.RelayMode == constant.RelayModeRerank { if info.RelayMode == constant.RelayModeRerank {
err, usage = cohereRerankHandler(c, resp, info) usage, err = cohereRerankHandler(c, resp, info)
} else { } else {
if info.IsStream { if info.IsStream {
err, usage = cohereStreamHandler(c, resp, info) usage, err = cohereStreamHandler(c, info, resp) // TODO: fix this
} else { } else {
err, usage = cohereHandler(c, resp, info.UpstreamModelName, info.PromptTokens) usage, err = cohereHandler(c, info, resp)
} }
} }
return return
......
...@@ -3,7 +3,6 @@ package cohere ...@@ -3,7 +3,6 @@ package cohere
import ( import (
"bufio" "bufio"
"encoding/json" "encoding/json"
"github.com/gin-gonic/gin"
"io" "io"
"net/http" "net/http"
"one-api/common" "one-api/common"
...@@ -11,8 +10,11 @@ import ( ...@@ -11,8 +10,11 @@ import (
relaycommon "one-api/relay/common" relaycommon "one-api/relay/common"
"one-api/relay/helper" "one-api/relay/helper"
"one-api/service" "one-api/service"
"one-api/types"
"strings" "strings"
"time" "time"
"github.com/gin-gonic/gin"
) )
func requestOpenAI2Cohere(textRequest dto.GeneralOpenAIRequest) *CohereRequest { func requestOpenAI2Cohere(textRequest dto.GeneralOpenAIRequest) *CohereRequest {
...@@ -76,7 +78,7 @@ func stopReasonCohere2OpenAI(reason string) string { ...@@ -76,7 +78,7 @@ func stopReasonCohere2OpenAI(reason string) string {
} }
} }
func cohereStreamHandler(c *gin.Context, resp *http.Response, info *relaycommon.RelayInfo) (*dto.OpenAIErrorWithStatusCode, *dto.Usage) { func cohereStreamHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Response) (*dto.Usage, *types.NewAPIError) {
responseId := helper.GetResponseID(c) responseId := helper.GetResponseID(c)
createdTime := common.GetTimestamp() createdTime := common.GetTimestamp()
usage := &dto.Usage{} usage := &dto.Usage{}
...@@ -164,20 +166,20 @@ func cohereStreamHandler(c *gin.Context, resp *http.Response, info *relaycommon. ...@@ -164,20 +166,20 @@ func cohereStreamHandler(c *gin.Context, resp *http.Response, info *relaycommon.
if usage.PromptTokens == 0 { if usage.PromptTokens == 0 {
usage = service.ResponseText2Usage(responseText, info.UpstreamModelName, info.PromptTokens) usage = service.ResponseText2Usage(responseText, info.UpstreamModelName, info.PromptTokens)
} }
return nil, usage return usage, nil
} }
func cohereHandler(c *gin.Context, resp *http.Response, modelName string, promptTokens int) (*dto.OpenAIErrorWithStatusCode, *dto.Usage) { func cohereHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Response) (*dto.Usage, *types.NewAPIError) {
createdTime := common.GetTimestamp() createdTime := common.GetTimestamp()
responseBody, err := io.ReadAll(resp.Body) responseBody, err := io.ReadAll(resp.Body)
if err != nil { if err != nil {
return service.OpenAIErrorWrapper(err, "read_response_body_failed", http.StatusInternalServerError), nil return nil, types.NewError(err, types.ErrorCodeBadResponseBody)
} }
common.CloseResponseBodyGracefully(resp) common.CloseResponseBodyGracefully(resp)
var cohereResp CohereResponseResult var cohereResp CohereResponseResult
err = json.Unmarshal(responseBody, &cohereResp) err = json.Unmarshal(responseBody, &cohereResp)
if err != nil { if err != nil {
return service.OpenAIErrorWrapper(err, "unmarshal_response_body_failed", http.StatusInternalServerError), nil return nil, types.NewError(err, types.ErrorCodeBadResponseBody)
} }
usage := dto.Usage{} usage := dto.Usage{}
usage.PromptTokens = cohereResp.Meta.BilledUnits.InputTokens usage.PromptTokens = cohereResp.Meta.BilledUnits.InputTokens
...@@ -188,7 +190,7 @@ func cohereHandler(c *gin.Context, resp *http.Response, modelName string, prompt ...@@ -188,7 +190,7 @@ func cohereHandler(c *gin.Context, resp *http.Response, modelName string, prompt
openaiResp.Id = cohereResp.ResponseId openaiResp.Id = cohereResp.ResponseId
openaiResp.Created = createdTime openaiResp.Created = createdTime
openaiResp.Object = "chat.completion" openaiResp.Object = "chat.completion"
openaiResp.Model = modelName openaiResp.Model = info.UpstreamModelName
openaiResp.Usage = usage openaiResp.Usage = usage
openaiResp.Choices = []dto.OpenAITextResponseChoice{ openaiResp.Choices = []dto.OpenAITextResponseChoice{
...@@ -201,24 +203,24 @@ func cohereHandler(c *gin.Context, resp *http.Response, modelName string, prompt ...@@ -201,24 +203,24 @@ func cohereHandler(c *gin.Context, resp *http.Response, modelName string, prompt
jsonResponse, err := json.Marshal(openaiResp) jsonResponse, err := json.Marshal(openaiResp)
if err != nil { if err != nil {
return service.OpenAIErrorWrapper(err, "marshal_response_body_failed", http.StatusInternalServerError), nil return nil, types.NewError(err, types.ErrorCodeBadResponseBody)
} }
c.Writer.Header().Set("Content-Type", "application/json") c.Writer.Header().Set("Content-Type", "application/json")
c.Writer.WriteHeader(resp.StatusCode) c.Writer.WriteHeader(resp.StatusCode)
_, err = c.Writer.Write(jsonResponse) _, _ = c.Writer.Write(jsonResponse)
return nil, &usage return &usage, nil
} }
func cohereRerankHandler(c *gin.Context, resp *http.Response, info *relaycommon.RelayInfo) (*dto.OpenAIErrorWithStatusCode, *dto.Usage) { func cohereRerankHandler(c *gin.Context, resp *http.Response, info *relaycommon.RelayInfo) (*dto.Usage, *types.NewAPIError) {
responseBody, err := io.ReadAll(resp.Body) responseBody, err := io.ReadAll(resp.Body)
if err != nil { if err != nil {
return service.OpenAIErrorWrapper(err, "read_response_body_failed", http.StatusInternalServerError), nil return nil, types.NewError(err, types.ErrorCodeBadResponseBody)
} }
common.CloseResponseBodyGracefully(resp) common.CloseResponseBodyGracefully(resp)
var cohereResp CohereRerankResponseResult var cohereResp CohereRerankResponseResult
err = json.Unmarshal(responseBody, &cohereResp) err = json.Unmarshal(responseBody, &cohereResp)
if err != nil { if err != nil {
return service.OpenAIErrorWrapper(err, "unmarshal_response_body_failed", http.StatusInternalServerError), nil return nil, types.NewError(err, types.ErrorCodeBadResponseBody)
} }
usage := dto.Usage{} usage := dto.Usage{}
if cohereResp.Meta.BilledUnits.InputTokens == 0 { if cohereResp.Meta.BilledUnits.InputTokens == 0 {
...@@ -237,10 +239,10 @@ func cohereRerankHandler(c *gin.Context, resp *http.Response, info *relaycommon. ...@@ -237,10 +239,10 @@ func cohereRerankHandler(c *gin.Context, resp *http.Response, info *relaycommon.
jsonResponse, err := json.Marshal(rerankResp) jsonResponse, err := json.Marshal(rerankResp)
if err != nil { if err != nil {
return service.OpenAIErrorWrapper(err, "marshal_response_body_failed", http.StatusInternalServerError), nil return nil, types.NewError(err, types.ErrorCodeBadResponseBody)
} }
c.Writer.Header().Set("Content-Type", "application/json") c.Writer.Header().Set("Content-Type", "application/json")
c.Writer.WriteHeader(resp.StatusCode) c.Writer.WriteHeader(resp.StatusCode)
_, err = c.Writer.Write(jsonResponse) _, err = c.Writer.Write(jsonResponse)
return nil, &usage return &usage, nil
} }
...@@ -9,6 +9,7 @@ import ( ...@@ -9,6 +9,7 @@ import (
"one-api/dto" "one-api/dto"
"one-api/relay/channel" "one-api/relay/channel"
"one-api/relay/common" "one-api/relay/common"
"one-api/types"
"time" "time"
"github.com/gin-gonic/gin" "github.com/gin-gonic/gin"
...@@ -95,11 +96,11 @@ func (a *Adaptor) DoRequest(c *gin.Context, info *common.RelayInfo, requestBody ...@@ -95,11 +96,11 @@ func (a *Adaptor) DoRequest(c *gin.Context, info *common.RelayInfo, requestBody
} }
// DoResponse implements channel.Adaptor. // DoResponse implements channel.Adaptor.
func (a *Adaptor) DoResponse(c *gin.Context, resp *http.Response, info *common.RelayInfo) (usage any, err *dto.OpenAIErrorWithStatusCode) { func (a *Adaptor) DoResponse(c *gin.Context, resp *http.Response, info *common.RelayInfo) (usage any, err *types.NewAPIError) {
if info.IsStream { if info.IsStream {
err, usage = cozeChatStreamHandler(c, resp, info) usage, err = cozeChatStreamHandler(c, info, resp)
} else { } else {
err, usage = cozeChatHandler(c, resp, info) usage, err = cozeChatHandler(c, info, resp)
} }
return return
} }
......
...@@ -12,6 +12,7 @@ import ( ...@@ -12,6 +12,7 @@ import (
relaycommon "one-api/relay/common" relaycommon "one-api/relay/common"
"one-api/relay/helper" "one-api/relay/helper"
"one-api/service" "one-api/service"
"one-api/types"
"strings" "strings"
"github.com/gin-gonic/gin" "github.com/gin-gonic/gin"
...@@ -43,10 +44,10 @@ func convertCozeChatRequest(c *gin.Context, request dto.GeneralOpenAIRequest) *C ...@@ -43,10 +44,10 @@ func convertCozeChatRequest(c *gin.Context, request dto.GeneralOpenAIRequest) *C
return cozeRequest return cozeRequest
} }
func cozeChatHandler(c *gin.Context, resp *http.Response, info *relaycommon.RelayInfo) (*dto.OpenAIErrorWithStatusCode, *dto.Usage) { func cozeChatHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Response) (*dto.Usage, *types.NewAPIError) {
responseBody, err := io.ReadAll(resp.Body) responseBody, err := io.ReadAll(resp.Body)
if err != nil { if err != nil {
return service.OpenAIErrorWrapper(err, "read_response_body_failed", http.StatusInternalServerError), nil return nil, types.NewError(err, types.ErrorCodeBadResponseBody)
} }
common.CloseResponseBodyGracefully(resp) common.CloseResponseBodyGracefully(resp)
// convert coze response to openai response // convert coze response to openai response
...@@ -55,10 +56,10 @@ func cozeChatHandler(c *gin.Context, resp *http.Response, info *relaycommon.Rela ...@@ -55,10 +56,10 @@ func cozeChatHandler(c *gin.Context, resp *http.Response, info *relaycommon.Rela
response.Model = info.UpstreamModelName response.Model = info.UpstreamModelName
err = json.Unmarshal(responseBody, &cozeResponse) err = json.Unmarshal(responseBody, &cozeResponse)
if err != nil { if err != nil {
return service.OpenAIErrorWrapper(err, "unmarshal_response_body_failed", http.StatusInternalServerError), nil return nil, types.NewError(err, types.ErrorCodeBadResponseBody)
} }
if cozeResponse.Code != 0 { if cozeResponse.Code != 0 {
return service.OpenAIErrorWrapper(errors.New(cozeResponse.Msg), fmt.Sprintf("%d", cozeResponse.Code), http.StatusInternalServerError), nil return nil, types.NewError(errors.New(cozeResponse.Msg), types.ErrorCodeBadResponseBody)
} }
// 从上下文获取 usage // 从上下文获取 usage
var usage dto.Usage var usage dto.Usage
...@@ -85,16 +86,16 @@ func cozeChatHandler(c *gin.Context, resp *http.Response, info *relaycommon.Rela ...@@ -85,16 +86,16 @@ func cozeChatHandler(c *gin.Context, resp *http.Response, info *relaycommon.Rela
} }
jsonResponse, err := json.Marshal(response) jsonResponse, err := json.Marshal(response)
if err != nil { if err != nil {
return service.OpenAIErrorWrapper(err, "marshal_response_body_failed", http.StatusInternalServerError), nil return nil, types.NewError(err, types.ErrorCodeBadResponseBody)
} }
c.Writer.Header().Set("Content-Type", "application/json") c.Writer.Header().Set("Content-Type", "application/json")
c.Writer.WriteHeader(resp.StatusCode) c.Writer.WriteHeader(resp.StatusCode)
_, _ = c.Writer.Write(jsonResponse) _, _ = c.Writer.Write(jsonResponse)
return nil, &usage return &usage, nil
} }
func cozeChatStreamHandler(c *gin.Context, resp *http.Response, info *relaycommon.RelayInfo) (*dto.OpenAIErrorWithStatusCode, *dto.Usage) { func cozeChatStreamHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Response) (*dto.Usage, *types.NewAPIError) {
scanner := bufio.NewScanner(resp.Body) scanner := bufio.NewScanner(resp.Body)
scanner.Split(bufio.ScanLines) scanner.Split(bufio.ScanLines)
helper.SetEventStreamHeaders(c) helper.SetEventStreamHeaders(c)
...@@ -135,7 +136,7 @@ func cozeChatStreamHandler(c *gin.Context, resp *http.Response, info *relaycommo ...@@ -135,7 +136,7 @@ func cozeChatStreamHandler(c *gin.Context, resp *http.Response, info *relaycommo
} }
if err := scanner.Err(); err != nil { if err := scanner.Err(); err != nil {
return service.OpenAIErrorWrapper(err, "stream_scanner_error", http.StatusInternalServerError), nil return nil, types.NewError(err, types.ErrorCodeBadResponseBody)
} }
helper.Done(c) helper.Done(c)
...@@ -143,7 +144,7 @@ func cozeChatStreamHandler(c *gin.Context, resp *http.Response, info *relaycommo ...@@ -143,7 +144,7 @@ func cozeChatStreamHandler(c *gin.Context, resp *http.Response, info *relaycommo
usage = service.ResponseText2Usage(responseText, info.UpstreamModelName, c.GetInt("coze_input_count")) usage = service.ResponseText2Usage(responseText, info.UpstreamModelName, c.GetInt("coze_input_count"))
} }
return nil, usage return usage, nil
} }
func handleCozeEvent(c *gin.Context, event string, data string, responseText *string, usage *dto.Usage, id string, info *relaycommon.RelayInfo) { func handleCozeEvent(c *gin.Context, event string, data string, responseText *string, usage *dto.Usage, id string, info *relaycommon.RelayInfo) {
......
...@@ -10,6 +10,7 @@ import ( ...@@ -10,6 +10,7 @@ import (
"one-api/relay/channel/openai" "one-api/relay/channel/openai"
relaycommon "one-api/relay/common" relaycommon "one-api/relay/common"
"one-api/relay/constant" "one-api/relay/constant"
"one-api/types"
"strings" "strings"
"github.com/gin-gonic/gin" "github.com/gin-gonic/gin"
...@@ -81,11 +82,11 @@ func (a *Adaptor) DoRequest(c *gin.Context, info *relaycommon.RelayInfo, request ...@@ -81,11 +82,11 @@ func (a *Adaptor) DoRequest(c *gin.Context, info *relaycommon.RelayInfo, request
return channel.DoApiRequest(a, c, info, requestBody) return channel.DoApiRequest(a, c, info, requestBody)
} }
func (a *Adaptor) DoResponse(c *gin.Context, resp *http.Response, info *relaycommon.RelayInfo) (usage any, err *dto.OpenAIErrorWithStatusCode) { 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 = openai.OaiStreamHandler(c, resp, info) usage, err = openai.OaiStreamHandler(c, info, resp)
} else { } else {
err, usage = openai.OpenaiHandler(c, resp, info) usage, err = openai.OpenaiHandler(c, info, resp)
} }
return return
} }
......
...@@ -8,6 +8,7 @@ import ( ...@@ -8,6 +8,7 @@ import (
"one-api/dto" "one-api/dto"
"one-api/relay/channel" "one-api/relay/channel"
relaycommon "one-api/relay/common" relaycommon "one-api/relay/common"
"one-api/types"
"github.com/gin-gonic/gin" "github.com/gin-gonic/gin"
) )
...@@ -96,11 +97,11 @@ func (a *Adaptor) DoRequest(c *gin.Context, info *relaycommon.RelayInfo, request ...@@ -96,11 +97,11 @@ func (a *Adaptor) DoRequest(c *gin.Context, info *relaycommon.RelayInfo, request
return channel.DoApiRequest(a, c, info, requestBody) return channel.DoApiRequest(a, c, info, requestBody)
} }
func (a *Adaptor) DoResponse(c *gin.Context, resp *http.Response, info *relaycommon.RelayInfo) (usage any, err *dto.OpenAIErrorWithStatusCode) { 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 = difyStreamHandler(c, resp, info) return difyStreamHandler(c, info, resp)
} else { } else {
err, usage = difyHandler(c, resp, info) return difyHandler(c, info, resp)
} }
return return
} }
......
...@@ -14,6 +14,7 @@ import ( ...@@ -14,6 +14,7 @@ import (
relaycommon "one-api/relay/common" relaycommon "one-api/relay/common"
"one-api/relay/helper" "one-api/relay/helper"
"one-api/service" "one-api/service"
"one-api/types"
"os" "os"
"strings" "strings"
...@@ -209,7 +210,7 @@ func streamResponseDify2OpenAI(difyResponse DifyChunkChatCompletionResponse) *dt ...@@ -209,7 +210,7 @@ func streamResponseDify2OpenAI(difyResponse DifyChunkChatCompletionResponse) *dt
return &response return &response
} }
func difyStreamHandler(c *gin.Context, resp *http.Response, info *relaycommon.RelayInfo) (*dto.OpenAIErrorWithStatusCode, *dto.Usage) { func difyStreamHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Response) (*dto.Usage, *types.NewAPIError) {
var responseText string var responseText string
usage := &dto.Usage{} usage := &dto.Usage{}
var nodeToken int var nodeToken int
...@@ -247,20 +248,20 @@ func difyStreamHandler(c *gin.Context, resp *http.Response, info *relaycommon.Re ...@@ -247,20 +248,20 @@ func difyStreamHandler(c *gin.Context, resp *http.Response, info *relaycommon.Re
usage = service.ResponseText2Usage(responseText, info.UpstreamModelName, info.PromptTokens) usage = service.ResponseText2Usage(responseText, info.UpstreamModelName, info.PromptTokens)
} }
usage.CompletionTokens += nodeToken usage.CompletionTokens += nodeToken
return nil, usage return usage, nil
} }
func difyHandler(c *gin.Context, resp *http.Response, info *relaycommon.RelayInfo) (*dto.OpenAIErrorWithStatusCode, *dto.Usage) { func difyHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Response) (*dto.Usage, *types.NewAPIError) {
var difyResponse DifyChatCompletionResponse var difyResponse DifyChatCompletionResponse
responseBody, err := io.ReadAll(resp.Body) responseBody, err := io.ReadAll(resp.Body)
if err != nil { if err != nil {
return service.OpenAIErrorWrapper(err, "read_response_body_failed", http.StatusInternalServerError), nil return nil, types.NewError(err, types.ErrorCodeBadResponseBody)
} }
common.CloseResponseBodyGracefully(resp) common.CloseResponseBodyGracefully(resp)
err = json.Unmarshal(responseBody, &difyResponse) err = json.Unmarshal(responseBody, &difyResponse)
if err != nil { if err != nil {
return service.OpenAIErrorWrapper(err, "unmarshal_response_body_failed", http.StatusInternalServerError), nil return nil, types.NewError(err, types.ErrorCodeBadResponseBody)
} }
fullTextResponse := dto.OpenAITextResponse{ fullTextResponse := dto.OpenAITextResponse{
Id: difyResponse.ConversationId, Id: difyResponse.ConversationId,
...@@ -279,10 +280,10 @@ func difyHandler(c *gin.Context, resp *http.Response, info *relaycommon.RelayInf ...@@ -279,10 +280,10 @@ func difyHandler(c *gin.Context, resp *http.Response, info *relaycommon.RelayInf
fullTextResponse.Choices = append(fullTextResponse.Choices, choice) fullTextResponse.Choices = append(fullTextResponse.Choices, choice)
jsonResponse, err := json.Marshal(fullTextResponse) jsonResponse, err := json.Marshal(fullTextResponse)
if err != nil { if err != nil {
return service.OpenAIErrorWrapper(err, "marshal_response_body_failed", http.StatusInternalServerError), nil return nil, types.NewError(err, types.ErrorCodeBadResponseBody)
} }
c.Writer.Header().Set("Content-Type", "application/json") c.Writer.Header().Set("Content-Type", "application/json")
c.Writer.WriteHeader(resp.StatusCode) c.Writer.WriteHeader(resp.StatusCode)
_, err = c.Writer.Write(jsonResponse) c.Writer.Write(jsonResponse)
return nil, &difyResponse.MetaData.Usage return &difyResponse.MetaData.Usage, nil
} }
...@@ -11,8 +11,8 @@ import ( ...@@ -11,8 +11,8 @@ import (
"one-api/relay/channel" "one-api/relay/channel"
relaycommon "one-api/relay/common" relaycommon "one-api/relay/common"
"one-api/relay/constant" "one-api/relay/constant"
"one-api/service"
"one-api/setting/model_setting" "one-api/setting/model_setting"
"one-api/types"
"strings" "strings"
"github.com/gin-gonic/gin" "github.com/gin-gonic/gin"
...@@ -168,30 +168,30 @@ func (a *Adaptor) DoRequest(c *gin.Context, info *relaycommon.RelayInfo, request ...@@ -168,30 +168,30 @@ func (a *Adaptor) DoRequest(c *gin.Context, info *relaycommon.RelayInfo, request
return channel.DoApiRequest(a, c, info, requestBody) return channel.DoApiRequest(a, c, info, requestBody)
} }
func (a *Adaptor) DoResponse(c *gin.Context, resp *http.Response, info *relaycommon.RelayInfo) (usage any, err *dto.OpenAIErrorWithStatusCode) { func (a *Adaptor) DoResponse(c *gin.Context, resp *http.Response, info *relaycommon.RelayInfo) (usage any, err *types.NewAPIError) {
if info.RelayMode == constant.RelayModeGemini { if info.RelayMode == constant.RelayModeGemini {
if info.IsStream { if info.IsStream {
return GeminiTextGenerationStreamHandler(c, resp, info) return GeminiTextGenerationStreamHandler(c, info, resp)
} else { } else {
return GeminiTextGenerationHandler(c, resp, info) return GeminiTextGenerationHandler(c, info, resp)
} }
} }
if strings.HasPrefix(info.UpstreamModelName, "imagen") { if strings.HasPrefix(info.UpstreamModelName, "imagen") {
return GeminiImageHandler(c, resp, info) return GeminiImageHandler(c, info, resp)
} }
// check if the model is an embedding model // check if the model is an embedding model
if strings.HasPrefix(info.UpstreamModelName, "text-embedding") || if strings.HasPrefix(info.UpstreamModelName, "text-embedding") ||
strings.HasPrefix(info.UpstreamModelName, "embedding") || strings.HasPrefix(info.UpstreamModelName, "embedding") ||
strings.HasPrefix(info.UpstreamModelName, "gemini-embedding") { strings.HasPrefix(info.UpstreamModelName, "gemini-embedding") {
return GeminiEmbeddingHandler(c, resp, info) return GeminiEmbeddingHandler(c, info, resp)
} }
if info.IsStream { if info.IsStream {
err, usage = GeminiChatStreamHandler(c, resp, info) return GeminiChatStreamHandler(c, info, resp)
} else { } else {
err, usage = GeminiChatHandler(c, resp, info) return GeminiChatHandler(c, info, resp)
} }
//if usage.(*dto.Usage).CompletionTokenDetails.ReasoningTokens > 100 { //if usage.(*dto.Usage).CompletionTokenDetails.ReasoningTokens > 100 {
...@@ -205,23 +205,23 @@ func (a *Adaptor) DoResponse(c *gin.Context, resp *http.Response, info *relaycom ...@@ -205,23 +205,23 @@ func (a *Adaptor) DoResponse(c *gin.Context, resp *http.Response, info *relaycom
// } // }
//} //}
return return nil, types.NewError(errors.New("not implemented"), types.ErrorCodeBadResponseBody)
} }
func GeminiImageHandler(c *gin.Context, resp *http.Response, info *relaycommon.RelayInfo) (usage any, err *dto.OpenAIErrorWithStatusCode) { func GeminiImageHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Response) (*dto.Usage, *types.NewAPIError) {
responseBody, readErr := io.ReadAll(resp.Body) responseBody, readErr := io.ReadAll(resp.Body)
if readErr != nil { if readErr != nil {
return nil, service.OpenAIErrorWrapper(readErr, "read_response_body_failed", http.StatusInternalServerError) return nil, types.NewError(readErr, types.ErrorCodeBadResponseBody)
} }
_ = resp.Body.Close() _ = resp.Body.Close()
var geminiResponse GeminiImageResponse var geminiResponse GeminiImageResponse
if jsonErr := json.Unmarshal(responseBody, &geminiResponse); jsonErr != nil { if jsonErr := json.Unmarshal(responseBody, &geminiResponse); jsonErr != nil {
return nil, service.OpenAIErrorWrapper(jsonErr, "unmarshal_response_body_failed", http.StatusInternalServerError) return nil, types.NewError(jsonErr, types.ErrorCodeBadResponseBody)
} }
if len(geminiResponse.Predictions) == 0 { if len(geminiResponse.Predictions) == 0 {
return nil, service.OpenAIErrorWrapper(errors.New("no images generated"), "no_images", http.StatusBadRequest) return nil, types.NewError(errors.New("no images generated"), types.ErrorCodeBadResponseBody)
} }
// convert to openai format response // convert to openai format response
...@@ -241,7 +241,7 @@ func GeminiImageHandler(c *gin.Context, resp *http.Response, info *relaycommon.R ...@@ -241,7 +241,7 @@ func GeminiImageHandler(c *gin.Context, resp *http.Response, info *relaycommon.R
jsonResponse, jsonErr := json.Marshal(openAIResponse) jsonResponse, jsonErr := json.Marshal(openAIResponse)
if jsonErr != nil { if jsonErr != nil {
return nil, service.OpenAIErrorWrapper(jsonErr, "marshal_response_failed", http.StatusInternalServerError) return nil, types.NewError(jsonErr, types.ErrorCodeBadResponseBody)
} }
c.Writer.Header().Set("Content-Type", "application/json") c.Writer.Header().Set("Content-Type", "application/json")
...@@ -253,7 +253,7 @@ func GeminiImageHandler(c *gin.Context, resp *http.Response, info *relaycommon.R ...@@ -253,7 +253,7 @@ func GeminiImageHandler(c *gin.Context, resp *http.Response, info *relaycommon.R
const imageTokens = 258 const imageTokens = 258
generatedImages := len(openAIResponse.Data) generatedImages := len(openAIResponse.Data)
usage = &dto.Usage{ usage := &dto.Usage{
PromptTokens: imageTokens * generatedImages, // each generated image has fixed 258 tokens PromptTokens: imageTokens * generatedImages, // each generated image has fixed 258 tokens
CompletionTokens: 0, // image generation does not calculate completion tokens CompletionTokens: 0, // image generation does not calculate completion tokens
TotalTokens: imageTokens * generatedImages, TotalTokens: imageTokens * generatedImages,
......
...@@ -8,18 +8,19 @@ import ( ...@@ -8,18 +8,19 @@ import (
relaycommon "one-api/relay/common" relaycommon "one-api/relay/common"
"one-api/relay/helper" "one-api/relay/helper"
"one-api/service" "one-api/service"
"one-api/types"
"strings" "strings"
"github.com/gin-gonic/gin" "github.com/gin-gonic/gin"
) )
func GeminiTextGenerationHandler(c *gin.Context, resp *http.Response, info *relaycommon.RelayInfo) (*dto.Usage, *dto.OpenAIErrorWithStatusCode) { func GeminiTextGenerationHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Response) (*dto.Usage, *types.NewAPIError) {
defer common.CloseResponseBodyGracefully(resp) defer common.CloseResponseBodyGracefully(resp)
// 读取响应体 // 读取响应体
responseBody, err := io.ReadAll(resp.Body) responseBody, err := io.ReadAll(resp.Body)
if err != nil { if err != nil {
return nil, service.OpenAIErrorWrapper(err, "read_response_body_failed", http.StatusInternalServerError) return nil, types.NewError(err, types.ErrorCodeBadResponseBody)
} }
if common.DebugEnabled { if common.DebugEnabled {
...@@ -28,9 +29,9 @@ func GeminiTextGenerationHandler(c *gin.Context, resp *http.Response, info *rela ...@@ -28,9 +29,9 @@ func GeminiTextGenerationHandler(c *gin.Context, resp *http.Response, info *rela
// 解析为 Gemini 原生响应格式 // 解析为 Gemini 原生响应格式
var geminiResponse GeminiChatResponse var geminiResponse GeminiChatResponse
err = common.UnmarshalJson(responseBody, &geminiResponse) err = common.Unmarshal(responseBody, &geminiResponse)
if err != nil { if err != nil {
return nil, service.OpenAIErrorWrapper(err, "unmarshal_response_body_failed", http.StatusInternalServerError) return nil, types.NewError(err, types.ErrorCodeBadResponseBody)
} }
// 计算使用量(基于 UsageMetadata) // 计算使用量(基于 UsageMetadata)
...@@ -51,9 +52,9 @@ func GeminiTextGenerationHandler(c *gin.Context, resp *http.Response, info *rela ...@@ -51,9 +52,9 @@ func GeminiTextGenerationHandler(c *gin.Context, resp *http.Response, info *rela
} }
// 直接返回 Gemini 原生格式的 JSON 响应 // 直接返回 Gemini 原生格式的 JSON 响应
jsonResponse, err := common.EncodeJson(geminiResponse) jsonResponse, err := common.Marshal(geminiResponse)
if err != nil { if err != nil {
return nil, service.OpenAIErrorWrapper(err, "marshal_response_body_failed", http.StatusInternalServerError) return nil, types.NewError(err, types.ErrorCodeBadResponseBody)
} }
common.IOCopyBytesGracefully(c, resp, jsonResponse) common.IOCopyBytesGracefully(c, resp, jsonResponse)
...@@ -61,7 +62,7 @@ func GeminiTextGenerationHandler(c *gin.Context, resp *http.Response, info *rela ...@@ -61,7 +62,7 @@ func GeminiTextGenerationHandler(c *gin.Context, resp *http.Response, info *rela
return &usage, nil return &usage, nil
} }
func GeminiTextGenerationStreamHandler(c *gin.Context, resp *http.Response, info *relaycommon.RelayInfo) (*dto.Usage, *dto.OpenAIErrorWithStatusCode) { func GeminiTextGenerationStreamHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Response) (*dto.Usage, *types.NewAPIError) {
var usage = &dto.Usage{} var usage = &dto.Usage{}
var imageCount int var imageCount int
......
...@@ -2,6 +2,7 @@ package gemini ...@@ -2,6 +2,7 @@ package gemini
import ( import (
"encoding/json" "encoding/json"
"errors"
"fmt" "fmt"
"io" "io"
"net/http" "net/http"
...@@ -12,6 +13,7 @@ import ( ...@@ -12,6 +13,7 @@ import (
"one-api/relay/helper" "one-api/relay/helper"
"one-api/service" "one-api/service"
"one-api/setting/model_setting" "one-api/setting/model_setting"
"one-api/types"
"strconv" "strconv"
"strings" "strings"
"unicode/utf8" "unicode/utf8"
...@@ -792,7 +794,7 @@ func streamResponseGeminiChat2OpenAI(geminiResponse *GeminiChatResponse) (*dto.C ...@@ -792,7 +794,7 @@ func streamResponseGeminiChat2OpenAI(geminiResponse *GeminiChatResponse) (*dto.C
return &response, isStop, hasImage return &response, isStop, hasImage
} }
func GeminiChatStreamHandler(c *gin.Context, resp *http.Response, info *relaycommon.RelayInfo) (*dto.OpenAIErrorWithStatusCode, *dto.Usage) { func GeminiChatStreamHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Response) (*dto.Usage, *types.NewAPIError) {
// responseText := "" // responseText := ""
id := helper.GetResponseID(c) id := helper.GetResponseID(c)
createAt := common.GetTimestamp() createAt := common.GetTimestamp()
...@@ -858,33 +860,25 @@ func GeminiChatStreamHandler(c *gin.Context, resp *http.Response, info *relaycom ...@@ -858,33 +860,25 @@ func GeminiChatStreamHandler(c *gin.Context, resp *http.Response, info *relaycom
} }
helper.Done(c) helper.Done(c)
//resp.Body.Close() //resp.Body.Close()
return nil, usage return usage, nil
} }
func GeminiChatHandler(c *gin.Context, resp *http.Response, info *relaycommon.RelayInfo) (*dto.OpenAIErrorWithStatusCode, *dto.Usage) { func GeminiChatHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Response) (*dto.Usage, *types.NewAPIError) {
responseBody, err := io.ReadAll(resp.Body) responseBody, err := io.ReadAll(resp.Body)
if err != nil { if err != nil {
return service.OpenAIErrorWrapper(err, "read_response_body_failed", http.StatusInternalServerError), nil return nil, types.NewError(err, types.ErrorCodeBadResponseBody)
} }
common.CloseResponseBodyGracefully(resp) common.CloseResponseBodyGracefully(resp)
if common.DebugEnabled { if common.DebugEnabled {
println(string(responseBody)) println(string(responseBody))
} }
var geminiResponse GeminiChatResponse var geminiResponse GeminiChatResponse
err = common.UnmarshalJson(responseBody, &geminiResponse) err = common.Unmarshal(responseBody, &geminiResponse)
if err != nil { if err != nil {
return service.OpenAIErrorWrapper(err, "unmarshal_response_body_failed", http.StatusInternalServerError), nil return nil, types.NewError(err, types.ErrorCodeBadResponseBody)
} }
if len(geminiResponse.Candidates) == 0 { if len(geminiResponse.Candidates) == 0 {
return &dto.OpenAIErrorWithStatusCode{ return nil, types.NewError(errors.New("no candidates returned"), types.ErrorCodeBadResponseBody)
Error: dto.OpenAIError{
Message: "No candidates returned",
Type: "server_error",
Param: "",
Code: 500,
},
StatusCode: resp.StatusCode,
}, nil
} }
fullTextResponse := responseGeminiChat2OpenAI(c, &geminiResponse) fullTextResponse := responseGeminiChat2OpenAI(c, &geminiResponse)
fullTextResponse.Model = info.UpstreamModelName fullTextResponse.Model = info.UpstreamModelName
...@@ -908,25 +902,25 @@ func GeminiChatHandler(c *gin.Context, resp *http.Response, info *relaycommon.Re ...@@ -908,25 +902,25 @@ func GeminiChatHandler(c *gin.Context, resp *http.Response, info *relaycommon.Re
fullTextResponse.Usage = usage fullTextResponse.Usage = usage
jsonResponse, err := json.Marshal(fullTextResponse) jsonResponse, err := json.Marshal(fullTextResponse)
if err != nil { if err != nil {
return service.OpenAIErrorWrapper(err, "marshal_response_body_failed", http.StatusInternalServerError), nil return nil, types.NewError(err, types.ErrorCodeBadResponseBody)
} }
c.Writer.Header().Set("Content-Type", "application/json") c.Writer.Header().Set("Content-Type", "application/json")
c.Writer.WriteHeader(resp.StatusCode) c.Writer.WriteHeader(resp.StatusCode)
_, err = c.Writer.Write(jsonResponse) c.Writer.Write(jsonResponse)
return nil, &usage return &usage, nil
} }
func GeminiEmbeddingHandler(c *gin.Context, resp *http.Response, info *relaycommon.RelayInfo) (usage any, err *dto.OpenAIErrorWithStatusCode) { func GeminiEmbeddingHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Response) (*dto.Usage, *types.NewAPIError) {
defer common.CloseResponseBodyGracefully(resp) defer common.CloseResponseBodyGracefully(resp)
responseBody, readErr := io.ReadAll(resp.Body) responseBody, readErr := io.ReadAll(resp.Body)
if readErr != nil { if readErr != nil {
return nil, service.OpenAIErrorWrapper(readErr, "read_response_body_failed", http.StatusInternalServerError) return nil, types.NewError(readErr, types.ErrorCodeBadResponseBody)
} }
var geminiResponse GeminiEmbeddingResponse var geminiResponse GeminiEmbeddingResponse
if jsonErr := json.Unmarshal(responseBody, &geminiResponse); jsonErr != nil { if jsonErr := common.Unmarshal(responseBody, &geminiResponse); jsonErr != nil {
return nil, service.OpenAIErrorWrapper(jsonErr, "unmarshal_response_body_failed", http.StatusInternalServerError) return nil, types.NewError(jsonErr, types.ErrorCodeBadResponseBody)
} }
// convert to openai format response // convert to openai format response
...@@ -947,16 +941,16 @@ func GeminiEmbeddingHandler(c *gin.Context, resp *http.Response, info *relaycomm ...@@ -947,16 +941,16 @@ func GeminiEmbeddingHandler(c *gin.Context, resp *http.Response, info *relaycomm
// Google has not yet clarified how embedding models will be billed // Google has not yet clarified how embedding models will be billed
// refer to openai billing method to use input tokens billing // refer to openai billing method to use input tokens billing
// https://platform.openai.com/docs/guides/embeddings#what-are-embeddings // https://platform.openai.com/docs/guides/embeddings#what-are-embeddings
usage = &dto.Usage{ usage := &dto.Usage{
PromptTokens: info.PromptTokens, PromptTokens: info.PromptTokens,
CompletionTokens: 0, CompletionTokens: 0,
TotalTokens: info.PromptTokens, TotalTokens: info.PromptTokens,
} }
openAIResponse.Usage = *usage.(*dto.Usage) openAIResponse.Usage = *usage
jsonResponse, jsonErr := common.EncodeJson(openAIResponse) jsonResponse, jsonErr := common.Marshal(openAIResponse)
if jsonErr != nil { if jsonErr != nil {
return nil, service.OpenAIErrorWrapper(jsonErr, "marshal_response_failed", http.StatusInternalServerError) return nil, types.NewError(jsonErr, types.ErrorCodeBadResponseBody)
} }
common.IOCopyBytesGracefully(c, resp, jsonResponse) common.IOCopyBytesGracefully(c, resp, jsonResponse)
......
...@@ -11,6 +11,7 @@ import ( ...@@ -11,6 +11,7 @@ import (
relaycommon "one-api/relay/common" relaycommon "one-api/relay/common"
"one-api/relay/common_handler" "one-api/relay/common_handler"
"one-api/relay/constant" "one-api/relay/constant"
"one-api/types"
"github.com/gin-gonic/gin" "github.com/gin-gonic/gin"
) )
...@@ -73,11 +74,11 @@ func (a *Adaptor) ConvertEmbeddingRequest(c *gin.Context, info *relaycommon.Rela ...@@ -73,11 +74,11 @@ func (a *Adaptor) ConvertEmbeddingRequest(c *gin.Context, info *relaycommon.Rela
return request, nil return request, nil
} }
func (a *Adaptor) DoResponse(c *gin.Context, resp *http.Response, info *relaycommon.RelayInfo) (usage any, err *dto.OpenAIErrorWithStatusCode) { func (a *Adaptor) DoResponse(c *gin.Context, resp *http.Response, info *relaycommon.RelayInfo) (usage any, err *types.NewAPIError) {
if info.RelayMode == constant.RelayModeRerank { if info.RelayMode == constant.RelayModeRerank {
err, usage = common_handler.RerankHandler(c, info, resp) usage, err = common_handler.RerankHandler(c, info, resp)
} else if info.RelayMode == constant.RelayModeEmbeddings { } else if info.RelayMode == constant.RelayModeEmbeddings {
err, usage = openai.OpenaiHandler(c, resp, info) usage, err = openai.OpenaiHandler(c, info, resp)
} }
return return
} }
......
...@@ -8,6 +8,7 @@ import ( ...@@ -8,6 +8,7 @@ import (
"one-api/relay/channel" "one-api/relay/channel"
"one-api/relay/channel/openai" "one-api/relay/channel/openai"
relaycommon "one-api/relay/common" relaycommon "one-api/relay/common"
"one-api/types"
"github.com/gin-gonic/gin" "github.com/gin-gonic/gin"
) )
...@@ -69,11 +70,11 @@ func (a *Adaptor) DoRequest(c *gin.Context, info *relaycommon.RelayInfo, request ...@@ -69,11 +70,11 @@ func (a *Adaptor) DoRequest(c *gin.Context, info *relaycommon.RelayInfo, request
return channel.DoApiRequest(a, c, info, requestBody) return channel.DoApiRequest(a, c, info, requestBody)
} }
func (a *Adaptor) DoResponse(c *gin.Context, resp *http.Response, info *relaycommon.RelayInfo) (usage any, err *dto.OpenAIErrorWithStatusCode) { 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 = openai.OaiStreamHandler(c, resp, info) usage, err = openai.OaiStreamHandler(c, info, resp)
} else { } else {
err, usage = openai.OpenaiHandler(c, resp, info) usage, err = openai.OpenaiHandler(c, info, resp)
} }
return return
} }
......
...@@ -9,6 +9,7 @@ import ( ...@@ -9,6 +9,7 @@ import (
"one-api/relay/channel" "one-api/relay/channel"
relaycommon "one-api/relay/common" relaycommon "one-api/relay/common"
"one-api/relay/constant" "one-api/relay/constant"
"one-api/types"
"strings" "strings"
"github.com/gin-gonic/gin" "github.com/gin-gonic/gin"
...@@ -84,11 +85,11 @@ func (a *Adaptor) DoRequest(c *gin.Context, info *relaycommon.RelayInfo, request ...@@ -84,11 +85,11 @@ func (a *Adaptor) DoRequest(c *gin.Context, info *relaycommon.RelayInfo, request
return channel.DoApiRequest(a, c, info, requestBody) return channel.DoApiRequest(a, c, info, requestBody)
} }
func (a *Adaptor) DoResponse(c *gin.Context, resp *http.Response, info *relaycommon.RelayInfo) (usage any, err *dto.OpenAIErrorWithStatusCode) { func (a *Adaptor) DoResponse(c *gin.Context, resp *http.Response, info *relaycommon.RelayInfo) (usage any, err *types.NewAPIError) {
switch info.RelayMode { switch info.RelayMode {
case constant.RelayModeEmbeddings: case constant.RelayModeEmbeddings:
err, usage = mokaEmbeddingHandler(c, resp) return mokaEmbeddingHandler(c, info, resp)
default: default:
// err, usage = mokaHandler(c, resp) // err, usage = mokaHandler(c, resp)
......
...@@ -2,12 +2,14 @@ package mokaai ...@@ -2,12 +2,14 @@ package mokaai
import ( import (
"encoding/json" "encoding/json"
"github.com/gin-gonic/gin"
"io" "io"
"net/http" "net/http"
"one-api/common" "one-api/common"
"one-api/dto" "one-api/dto"
"one-api/service" relaycommon "one-api/relay/common"
"one-api/types"
"github.com/gin-gonic/gin"
) )
func embeddingRequestOpenAI2Moka(request dto.GeneralOpenAIRequest) *dto.EmbeddingRequest { func embeddingRequestOpenAI2Moka(request dto.GeneralOpenAIRequest) *dto.EmbeddingRequest {
...@@ -48,16 +50,16 @@ func embeddingResponseMoka2OpenAI(response *dto.EmbeddingResponse) *dto.OpenAIEm ...@@ -48,16 +50,16 @@ func embeddingResponseMoka2OpenAI(response *dto.EmbeddingResponse) *dto.OpenAIEm
return &openAIEmbeddingResponse return &openAIEmbeddingResponse
} }
func mokaEmbeddingHandler(c *gin.Context, resp *http.Response) (*dto.OpenAIErrorWithStatusCode, *dto.Usage) { func mokaEmbeddingHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Response) (*dto.Usage, *types.NewAPIError) {
var baiduResponse dto.EmbeddingResponse var baiduResponse dto.EmbeddingResponse
responseBody, err := io.ReadAll(resp.Body) responseBody, err := io.ReadAll(resp.Body)
if err != nil { if err != nil {
return service.OpenAIErrorWrapper(err, "read_response_body_failed", http.StatusInternalServerError), nil return nil, types.NewError(err, types.ErrorCodeBadResponseBody)
} }
common.CloseResponseBodyGracefully(resp) common.CloseResponseBodyGracefully(resp)
err = json.Unmarshal(responseBody, &baiduResponse) err = json.Unmarshal(responseBody, &baiduResponse)
if err != nil { if err != nil {
return service.OpenAIErrorWrapper(err, "unmarshal_response_body_failed", http.StatusInternalServerError), nil return nil, types.NewError(err, types.ErrorCodeBadResponseBody)
} }
// if baiduResponse.ErrorMsg != "" { // if baiduResponse.ErrorMsg != "" {
// return &dto.OpenAIErrorWithStatusCode{ // return &dto.OpenAIErrorWithStatusCode{
...@@ -69,12 +71,12 @@ func mokaEmbeddingHandler(c *gin.Context, resp *http.Response) (*dto.OpenAIError ...@@ -69,12 +71,12 @@ func mokaEmbeddingHandler(c *gin.Context, resp *http.Response) (*dto.OpenAIError
// }, nil // }, nil
// } // }
fullTextResponse := embeddingResponseMoka2OpenAI(&baiduResponse) fullTextResponse := embeddingResponseMoka2OpenAI(&baiduResponse)
jsonResponse, err := json.Marshal(fullTextResponse) jsonResponse, err := common.Marshal(fullTextResponse)
if err != nil { if err != nil {
return service.OpenAIErrorWrapper(err, "marshal_response_body_failed", http.StatusInternalServerError), nil return nil, types.NewError(err, types.ErrorCodeBadResponseBody)
} }
c.Writer.Header().Set("Content-Type", "application/json") c.Writer.Header().Set("Content-Type", "application/json")
c.Writer.WriteHeader(resp.StatusCode) c.Writer.WriteHeader(resp.StatusCode)
_, err = c.Writer.Write(jsonResponse) common.IOCopyBytesGracefully(c, resp, jsonResponse)
return nil, &fullTextResponse.Usage return &fullTextResponse.Usage, nil
} }
...@@ -9,6 +9,7 @@ import ( ...@@ -9,6 +9,7 @@ import (
"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"
"one-api/types"
"github.com/gin-gonic/gin" "github.com/gin-gonic/gin"
) )
...@@ -74,14 +75,14 @@ func (a *Adaptor) DoRequest(c *gin.Context, info *relaycommon.RelayInfo, request ...@@ -74,14 +75,14 @@ func (a *Adaptor) DoRequest(c *gin.Context, info *relaycommon.RelayInfo, request
return channel.DoApiRequest(a, c, info, requestBody) return channel.DoApiRequest(a, c, info, requestBody)
} }
func (a *Adaptor) DoResponse(c *gin.Context, resp *http.Response, info *relaycommon.RelayInfo) (usage any, err *dto.OpenAIErrorWithStatusCode) { 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 = openai.OaiStreamHandler(c, resp, info) usage, err = openai.OaiStreamHandler(c, info, resp)
} else { } else {
if info.RelayMode == relayconstant.RelayModeEmbeddings { if info.RelayMode == relayconstant.RelayModeEmbeddings {
err, usage = ollamaEmbeddingHandler(c, resp, info.PromptTokens, info.UpstreamModelName, info.RelayMode) usage, err = ollamaEmbeddingHandler(c, info, resp)
} else { } else {
err, usage = openai.OpenaiHandler(c, resp, info) usage, err = openai.OpenaiHandler(c, info, resp)
} }
} }
return return
......
package ollama package ollama
import ( import (
"encoding/json"
"fmt" "fmt"
"github.com/gin-gonic/gin"
"io" "io"
"net/http" "net/http"
"one-api/common" "one-api/common"
"one-api/dto" "one-api/dto"
relaycommon "one-api/relay/common"
"one-api/service" "one-api/service"
"one-api/types"
"strings" "strings"
"github.com/gin-gonic/gin"
) )
func requestOpenAI2Ollama(request dto.GeneralOpenAIRequest) (*OllamaRequest, error) { func requestOpenAI2Ollama(request dto.GeneralOpenAIRequest) (*OllamaRequest, error) {
...@@ -82,19 +84,19 @@ func requestOpenAI2Embeddings(request dto.EmbeddingRequest) *OllamaEmbeddingRequ ...@@ -82,19 +84,19 @@ func requestOpenAI2Embeddings(request dto.EmbeddingRequest) *OllamaEmbeddingRequ
} }
} }
func ollamaEmbeddingHandler(c *gin.Context, resp *http.Response, promptTokens int, model string, relayMode int) (*dto.OpenAIErrorWithStatusCode, *dto.Usage) { func ollamaEmbeddingHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Response) (*dto.Usage, *types.NewAPIError) {
var ollamaEmbeddingResponse OllamaEmbeddingResponse var ollamaEmbeddingResponse OllamaEmbeddingResponse
responseBody, err := io.ReadAll(resp.Body) responseBody, err := io.ReadAll(resp.Body)
if err != nil { if err != nil {
return service.OpenAIErrorWrapper(err, "read_response_body_failed", http.StatusInternalServerError), nil return nil, types.NewError(err, types.ErrorCodeBadResponseBody)
} }
common.CloseResponseBodyGracefully(resp) common.CloseResponseBodyGracefully(resp)
err = json.Unmarshal(responseBody, &ollamaEmbeddingResponse) err = common.Unmarshal(responseBody, &ollamaEmbeddingResponse)
if err != nil { if err != nil {
return service.OpenAIErrorWrapper(err, "unmarshal_response_body_failed", http.StatusInternalServerError), nil return nil, types.NewError(err, types.ErrorCodeBadResponseBody)
} }
if ollamaEmbeddingResponse.Error != "" { if ollamaEmbeddingResponse.Error != "" {
return service.OpenAIErrorWrapper(err, "ollama_error", resp.StatusCode), nil return nil, types.NewError(fmt.Errorf("ollama error: %s", ollamaEmbeddingResponse.Error), types.ErrorCodeBadResponseBody)
} }
flattenedEmbeddings := flattenEmbeddings(ollamaEmbeddingResponse.Embedding) flattenedEmbeddings := flattenEmbeddings(ollamaEmbeddingResponse.Embedding)
data := make([]dto.OpenAIEmbeddingResponseItem, 0, 1) data := make([]dto.OpenAIEmbeddingResponseItem, 0, 1)
...@@ -103,22 +105,22 @@ func ollamaEmbeddingHandler(c *gin.Context, resp *http.Response, promptTokens in ...@@ -103,22 +105,22 @@ func ollamaEmbeddingHandler(c *gin.Context, resp *http.Response, promptTokens in
Object: "embedding", Object: "embedding",
}) })
usage := &dto.Usage{ usage := &dto.Usage{
TotalTokens: promptTokens, TotalTokens: info.PromptTokens,
CompletionTokens: 0, CompletionTokens: 0,
PromptTokens: promptTokens, PromptTokens: info.PromptTokens,
} }
embeddingResponse := &dto.OpenAIEmbeddingResponse{ embeddingResponse := &dto.OpenAIEmbeddingResponse{
Object: "list", Object: "list",
Data: data, Data: data,
Model: model, Model: info.UpstreamModelName,
Usage: *usage, Usage: *usage,
} }
doResponseBody, err := json.Marshal(embeddingResponse) doResponseBody, err := common.Marshal(embeddingResponse)
if err != nil { if err != nil {
return service.OpenAIErrorWrapper(err, "marshal_response_body_failed", http.StatusInternalServerError), nil return nil, types.NewError(err, types.ErrorCodeBadResponseBody)
} }
common.IOCopyBytesGracefully(c, resp, doResponseBody) common.IOCopyBytesGracefully(c, resp, doResponseBody)
return nil, usage return usage, nil
} }
func flattenEmbeddings(embeddings [][]float64) []float64 { func flattenEmbeddings(embeddings [][]float64) []float64 {
......
...@@ -22,6 +22,7 @@ import ( ...@@ -22,6 +22,7 @@ import (
"one-api/relay/common_handler" "one-api/relay/common_handler"
relayconstant "one-api/relay/constant" relayconstant "one-api/relay/constant"
"one-api/service" "one-api/service"
"one-api/types"
"path/filepath" "path/filepath"
"strings" "strings"
...@@ -34,9 +35,9 @@ type Adaptor struct { ...@@ -34,9 +35,9 @@ type Adaptor struct {
} }
func (a *Adaptor) ConvertClaudeRequest(c *gin.Context, info *relaycommon.RelayInfo, request *dto.ClaudeRequest) (any, error) { func (a *Adaptor) ConvertClaudeRequest(c *gin.Context, info *relaycommon.RelayInfo, request *dto.ClaudeRequest) (any, error) {
if !strings.Contains(request.Model, "claude") { //if !strings.Contains(request.Model, "claude") {
return nil, fmt.Errorf("you are using openai channel type with path /v1/messages, only claude model supported convert, but got %s", request.Model) // return nil, fmt.Errorf("you are using openai channel type with path /v1/messages, only claude model supported convert, but got %s", request.Model)
} //}
aiRequest, err := service.ClaudeToOpenAIRequest(*request, info) aiRequest, err := service.ClaudeToOpenAIRequest(*request, info)
if err != nil { if err != nil {
return nil, err return nil, err
...@@ -421,31 +422,31 @@ func (a *Adaptor) DoRequest(c *gin.Context, info *relaycommon.RelayInfo, request ...@@ -421,31 +422,31 @@ 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 *dto.OpenAIErrorWithStatusCode) { func (a *Adaptor) DoResponse(c *gin.Context, resp *http.Response, info *relaycommon.RelayInfo) (usage any, err *types.NewAPIError) {
switch info.RelayMode { switch info.RelayMode {
case relayconstant.RelayModeRealtime: case relayconstant.RelayModeRealtime:
err, usage = OpenaiRealtimeHandler(c, info) err, usage = OpenaiRealtimeHandler(c, info)
case relayconstant.RelayModeAudioSpeech: case relayconstant.RelayModeAudioSpeech:
err, usage = OpenaiTTSHandler(c, resp, info) usage = OpenaiTTSHandler(c, resp, info)
case relayconstant.RelayModeAudioTranslation: case relayconstant.RelayModeAudioTranslation:
fallthrough fallthrough
case relayconstant.RelayModeAudioTranscription: case relayconstant.RelayModeAudioTranscription:
err, usage = OpenaiSTTHandler(c, resp, info, a.ResponseFormat) err, usage = OpenaiSTTHandler(c, resp, info, a.ResponseFormat)
case relayconstant.RelayModeImagesGenerations, relayconstant.RelayModeImagesEdits: case relayconstant.RelayModeImagesGenerations, relayconstant.RelayModeImagesEdits:
err, usage = OpenaiHandlerWithUsage(c, resp, info) usage, err = OpenaiHandlerWithUsage(c, info, resp)
case relayconstant.RelayModeRerank: case relayconstant.RelayModeRerank:
err, usage = common_handler.RerankHandler(c, info, resp) usage, err = common_handler.RerankHandler(c, info, resp)
case relayconstant.RelayModeResponses: case relayconstant.RelayModeResponses:
if info.IsStream { if info.IsStream {
err, usage = OaiResponsesStreamHandler(c, resp, info) usage, err = OaiResponsesStreamHandler(c, info, resp)
} else { } else {
err, usage = OaiResponsesHandler(c, resp, info) usage, err = OaiResponsesHandler(c, info, resp)
} }
default: default:
if info.IsStream { if info.IsStream {
err, usage = OaiStreamHandler(c, resp, info) usage, err = OaiStreamHandler(c, info, resp)
} else { } else {
err, usage = OpenaiHandler(c, resp, info) usage, err = OpenaiHandler(c, info, resp)
} }
} }
return return
......
...@@ -17,6 +17,8 @@ import ( ...@@ -17,6 +17,8 @@ import (
"path/filepath" "path/filepath"
"strings" "strings"
"one-api/types"
"github.com/bytedance/gopkg/util/gopool" "github.com/bytedance/gopkg/util/gopool"
"github.com/gin-gonic/gin" "github.com/gin-gonic/gin"
"github.com/gorilla/websocket" "github.com/gorilla/websocket"
...@@ -104,10 +106,10 @@ func sendStreamData(c *gin.Context, info *relaycommon.RelayInfo, data string, fo ...@@ -104,10 +106,10 @@ func sendStreamData(c *gin.Context, info *relaycommon.RelayInfo, data string, fo
return helper.ObjectData(c, lastStreamResponse) return helper.ObjectData(c, lastStreamResponse)
} }
func OaiStreamHandler(c *gin.Context, resp *http.Response, info *relaycommon.RelayInfo) (*dto.OpenAIErrorWithStatusCode, *dto.Usage) { func OaiStreamHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Response) (*dto.Usage, *types.NewAPIError) {
if resp == nil || resp.Body == nil { if resp == nil || resp.Body == nil {
common.LogError(c, "invalid response or response body") common.LogError(c, "invalid response or response body")
return service.OpenAIErrorWrapper(fmt.Errorf("invalid response"), "invalid_response", http.StatusInternalServerError), nil return nil, types.NewError(fmt.Errorf("invalid response"), types.ErrorCodeBadResponse)
} }
defer common.CloseResponseBodyGracefully(resp) defer common.CloseResponseBodyGracefully(resp)
...@@ -177,26 +179,23 @@ func OaiStreamHandler(c *gin.Context, resp *http.Response, info *relaycommon.Rel ...@@ -177,26 +179,23 @@ func OaiStreamHandler(c *gin.Context, resp *http.Response, info *relaycommon.Rel
handleFinalResponse(c, info, lastStreamData, responseId, createAt, model, systemFingerprint, usage, containStreamUsage) handleFinalResponse(c, info, lastStreamData, responseId, createAt, model, systemFingerprint, usage, containStreamUsage)
return nil, usage return usage, nil
} }
func OpenaiHandler(c *gin.Context, resp *http.Response, info *relaycommon.RelayInfo) (*dto.OpenAIErrorWithStatusCode, *dto.Usage) { func OpenaiHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Response) (*dto.Usage, *types.NewAPIError) {
defer common.CloseResponseBodyGracefully(resp) defer common.CloseResponseBodyGracefully(resp)
var simpleResponse dto.OpenAITextResponse var simpleResponse dto.OpenAITextResponse
responseBody, err := io.ReadAll(resp.Body) responseBody, err := io.ReadAll(resp.Body)
if err != nil { if err != nil {
return service.OpenAIErrorWrapper(err, "read_response_body_failed", http.StatusInternalServerError), nil return nil, types.NewError(err, types.ErrorCodeReadResponseBodyFailed)
} }
err = common.UnmarshalJson(responseBody, &simpleResponse) err = common.Unmarshal(responseBody, &simpleResponse)
if err != nil { if err != nil {
return service.OpenAIErrorWrapper(err, "unmarshal_response_body_failed", http.StatusInternalServerError), nil return nil, types.NewError(err, types.ErrorCodeBadResponseBody)
} }
if simpleResponse.Error != nil && simpleResponse.Error.Type != "" { if simpleResponse.Error != nil && simpleResponse.Error.Type != "" {
return &dto.OpenAIErrorWithStatusCode{ return nil, types.WithOpenAIError(*simpleResponse.Error, resp.StatusCode)
Error: *simpleResponse.Error,
StatusCode: resp.StatusCode,
}, nil
} }
forceFormat := false forceFormat := false
...@@ -220,28 +219,28 @@ func OpenaiHandler(c *gin.Context, resp *http.Response, info *relaycommon.RelayI ...@@ -220,28 +219,28 @@ func OpenaiHandler(c *gin.Context, resp *http.Response, info *relaycommon.RelayI
switch info.RelayFormat { switch info.RelayFormat {
case relaycommon.RelayFormatOpenAI: case relaycommon.RelayFormatOpenAI:
if forceFormat { if forceFormat {
responseBody, err = common.EncodeJson(simpleResponse) responseBody, err = common.Marshal(simpleResponse)
if err != nil { if err != nil {
return service.OpenAIErrorWrapper(err, "marshal_response_body_failed", http.StatusInternalServerError), nil return nil, types.NewError(err, types.ErrorCodeBadResponseBody)
} }
} else { } else {
break break
} }
case relaycommon.RelayFormatClaude: case relaycommon.RelayFormatClaude:
claudeResp := service.ResponseOpenAI2Claude(&simpleResponse, info) claudeResp := service.ResponseOpenAI2Claude(&simpleResponse, info)
claudeRespStr, err := common.EncodeJson(claudeResp) claudeRespStr, err := common.Marshal(claudeResp)
if err != nil { if err != nil {
return service.OpenAIErrorWrapper(err, "marshal_response_body_failed", http.StatusInternalServerError), nil return nil, types.NewError(err, types.ErrorCodeBadResponseBody)
} }
responseBody = claudeRespStr responseBody = claudeRespStr
} }
common.IOCopyBytesGracefully(c, resp, responseBody) common.IOCopyBytesGracefully(c, resp, responseBody)
return nil, &simpleResponse.Usage return &simpleResponse.Usage, nil
} }
func OpenaiTTSHandler(c *gin.Context, resp *http.Response, info *relaycommon.RelayInfo) (*dto.OpenAIErrorWithStatusCode, *dto.Usage) { func OpenaiTTSHandler(c *gin.Context, resp *http.Response, info *relaycommon.RelayInfo) *dto.Usage {
// the status code has been judged before, if there is a body reading failure, // the status code has been judged before, if there is a body reading failure,
// it should be regarded as a non-recoverable error, so it should not return err for external retry. // it should be regarded as a non-recoverable error, so it should not return err for external retry.
// Analogous to nginx's load balancing, it will only retry if it can't be requested or // Analogous to nginx's load balancing, it will only retry if it can't be requested or
...@@ -261,20 +260,20 @@ func OpenaiTTSHandler(c *gin.Context, resp *http.Response, info *relaycommon.Rel ...@@ -261,20 +260,20 @@ func OpenaiTTSHandler(c *gin.Context, resp *http.Response, info *relaycommon.Rel
if err != nil { if err != nil {
common.LogError(c, err.Error()) common.LogError(c, err.Error())
} }
return nil, usage return usage
} }
func OpenaiSTTHandler(c *gin.Context, resp *http.Response, info *relaycommon.RelayInfo, responseFormat string) (*dto.OpenAIErrorWithStatusCode, *dto.Usage) { func OpenaiSTTHandler(c *gin.Context, resp *http.Response, info *relaycommon.RelayInfo, responseFormat string) (*types.NewAPIError, *dto.Usage) {
defer common.CloseResponseBodyGracefully(resp) defer common.CloseResponseBodyGracefully(resp)
// count tokens by audio file duration // count tokens by audio file duration
audioTokens, err := countAudioTokens(c) audioTokens, err := countAudioTokens(c)
if err != nil { if err != nil {
return service.OpenAIErrorWrapper(err, "count_audio_tokens_failed", http.StatusInternalServerError), nil return types.NewError(err, types.ErrorCodeCountTokenFailed), nil
} }
responseBody, err := io.ReadAll(resp.Body) responseBody, err := io.ReadAll(resp.Body)
if err != nil { if err != nil {
return service.OpenAIErrorWrapper(err, "read_response_body_failed", http.StatusInternalServerError), nil return types.NewError(err, types.ErrorCodeReadResponseBodyFailed), nil
} }
// 写入新的 response body // 写入新的 response body
common.IOCopyBytesGracefully(c, resp, responseBody) common.IOCopyBytesGracefully(c, resp, responseBody)
...@@ -328,9 +327,9 @@ func countAudioTokens(c *gin.Context) (int, error) { ...@@ -328,9 +327,9 @@ func countAudioTokens(c *gin.Context) (int, error) {
return int(math.Round(math.Ceil(duration) / 60.0 * 1000)), nil // 1 minute 相当于 1k tokens return int(math.Round(math.Ceil(duration) / 60.0 * 1000)), nil // 1 minute 相当于 1k tokens
} }
func OpenaiRealtimeHandler(c *gin.Context, info *relaycommon.RelayInfo) (*dto.OpenAIErrorWithStatusCode, *dto.RealtimeUsage) { func OpenaiRealtimeHandler(c *gin.Context, info *relaycommon.RelayInfo) (*types.NewAPIError, *dto.RealtimeUsage) {
if info == nil || info.ClientWs == nil || info.TargetWs == nil { if info == nil || info.ClientWs == nil || info.TargetWs == nil {
return service.OpenAIErrorWrapper(fmt.Errorf("invalid websocket connection"), "invalid_connection", http.StatusBadRequest), nil return types.NewError(fmt.Errorf("invalid websocket connection"), types.ErrorCodeBadResponse), nil
} }
info.IsStream = true info.IsStream = true
...@@ -368,7 +367,7 @@ func OpenaiRealtimeHandler(c *gin.Context, info *relaycommon.RelayInfo) (*dto.Op ...@@ -368,7 +367,7 @@ func OpenaiRealtimeHandler(c *gin.Context, info *relaycommon.RelayInfo) (*dto.Op
} }
realtimeEvent := &dto.RealtimeEvent{} realtimeEvent := &dto.RealtimeEvent{}
err = common.UnmarshalJson(message, realtimeEvent) err = common.Unmarshal(message, realtimeEvent)
if err != nil { if err != nil {
errChan <- fmt.Errorf("error unmarshalling message: %v", err) errChan <- fmt.Errorf("error unmarshalling message: %v", err)
return return
...@@ -428,7 +427,7 @@ func OpenaiRealtimeHandler(c *gin.Context, info *relaycommon.RelayInfo) (*dto.Op ...@@ -428,7 +427,7 @@ func OpenaiRealtimeHandler(c *gin.Context, info *relaycommon.RelayInfo) (*dto.Op
} }
info.SetFirstResponseTime() info.SetFirstResponseTime()
realtimeEvent := &dto.RealtimeEvent{} realtimeEvent := &dto.RealtimeEvent{}
err = common.UnmarshalJson(message, realtimeEvent) err = common.Unmarshal(message, realtimeEvent)
if err != nil { if err != nil {
errChan <- fmt.Errorf("error unmarshalling message: %v", err) errChan <- fmt.Errorf("error unmarshalling message: %v", err)
return return
...@@ -553,18 +552,18 @@ func preConsumeUsage(ctx *gin.Context, info *relaycommon.RelayInfo, usage *dto.R ...@@ -553,18 +552,18 @@ func preConsumeUsage(ctx *gin.Context, info *relaycommon.RelayInfo, usage *dto.R
return err return err
} }
func OpenaiHandlerWithUsage(c *gin.Context, resp *http.Response, info *relaycommon.RelayInfo) (*dto.OpenAIErrorWithStatusCode, *dto.Usage) { func OpenaiHandlerWithUsage(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Response) (*dto.Usage, *types.NewAPIError) {
defer common.CloseResponseBodyGracefully(resp) defer common.CloseResponseBodyGracefully(resp)
responseBody, err := io.ReadAll(resp.Body) responseBody, err := io.ReadAll(resp.Body)
if err != nil { if err != nil {
return service.OpenAIErrorWrapper(err, "read_response_body_failed", http.StatusInternalServerError), nil return nil, types.NewError(err, types.ErrorCodeReadResponseBodyFailed)
} }
var usageResp dto.SimpleResponse var usageResp dto.SimpleResponse
err = common.UnmarshalJson(responseBody, &usageResp) err = common.Unmarshal(responseBody, &usageResp)
if err != nil { if err != nil {
return service.OpenAIErrorWrapper(err, "parse_response_body_failed", http.StatusInternalServerError), nil return nil, types.NewError(err, types.ErrorCodeBadResponseBody)
} }
// 写入新的 response body // 写入新的 response body
...@@ -584,5 +583,5 @@ func OpenaiHandlerWithUsage(c *gin.Context, resp *http.Response, info *relaycomm ...@@ -584,5 +583,5 @@ func OpenaiHandlerWithUsage(c *gin.Context, resp *http.Response, info *relaycomm
usageResp.PromptTokensDetails.ImageTokens += usageResp.InputTokensDetails.ImageTokens usageResp.PromptTokensDetails.ImageTokens += usageResp.InputTokensDetails.ImageTokens
usageResp.PromptTokensDetails.TextTokens += usageResp.InputTokensDetails.TextTokens usageResp.PromptTokensDetails.TextTokens += usageResp.InputTokensDetails.TextTokens
} }
return nil, &usageResp.Usage return &usageResp.Usage, nil
} }
...@@ -9,33 +9,27 @@ import ( ...@@ -9,33 +9,27 @@ import (
relaycommon "one-api/relay/common" relaycommon "one-api/relay/common"
"one-api/relay/helper" "one-api/relay/helper"
"one-api/service" "one-api/service"
"one-api/types"
"strings" "strings"
"github.com/gin-gonic/gin" "github.com/gin-gonic/gin"
) )
func OaiResponsesHandler(c *gin.Context, resp *http.Response, info *relaycommon.RelayInfo) (*dto.OpenAIErrorWithStatusCode, *dto.Usage) { func OaiResponsesHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Response) (*dto.Usage, *types.NewAPIError) {
defer common.CloseResponseBodyGracefully(resp) defer common.CloseResponseBodyGracefully(resp)
// read response body // read response body
var responsesResponse dto.OpenAIResponsesResponse var responsesResponse dto.OpenAIResponsesResponse
responseBody, err := io.ReadAll(resp.Body) responseBody, err := io.ReadAll(resp.Body)
if err != nil { if err != nil {
return service.OpenAIErrorWrapper(err, "read_response_body_failed", http.StatusInternalServerError), nil return nil, types.NewError(err, types.ErrorCodeReadResponseBodyFailed)
} }
err = common.UnmarshalJson(responseBody, &responsesResponse) err = common.Unmarshal(responseBody, &responsesResponse)
if err != nil { if err != nil {
return service.OpenAIErrorWrapper(err, "unmarshal_response_body_failed", http.StatusInternalServerError), nil return nil, types.NewError(err, types.ErrorCodeBadResponseBody)
} }
if responsesResponse.Error != nil { if responsesResponse.Error != nil {
return &dto.OpenAIErrorWithStatusCode{ return nil, types.WithOpenAIError(*responsesResponse.Error, resp.StatusCode)
Error: dto.OpenAIError{
Message: responsesResponse.Error.Message,
Type: "openai_error",
Code: responsesResponse.Error.Code,
},
StatusCode: resp.StatusCode,
}, nil
} }
// 写入新的 response body // 写入新的 response body
...@@ -50,13 +44,13 @@ func OaiResponsesHandler(c *gin.Context, resp *http.Response, info *relaycommon. ...@@ -50,13 +44,13 @@ func OaiResponsesHandler(c *gin.Context, resp *http.Response, info *relaycommon.
for _, tool := range responsesResponse.Tools { for _, tool := range responsesResponse.Tools {
info.ResponsesUsageInfo.BuiltInTools[tool.Type].CallCount++ info.ResponsesUsageInfo.BuiltInTools[tool.Type].CallCount++
} }
return nil, &usage return &usage, nil
} }
func OaiResponsesStreamHandler(c *gin.Context, resp *http.Response, info *relaycommon.RelayInfo) (*dto.OpenAIErrorWithStatusCode, *dto.Usage) { func OaiResponsesStreamHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Response) (*dto.Usage, *types.NewAPIError) {
if resp == nil || resp.Body == nil { if resp == nil || resp.Body == nil {
common.LogError(c, "invalid response or response body") common.LogError(c, "invalid response or response body")
return service.OpenAIErrorWrapper(fmt.Errorf("invalid response"), "invalid_response", http.StatusInternalServerError), nil return nil, types.NewError(fmt.Errorf("invalid response"), types.ErrorCodeBadResponse)
} }
var usage = &dto.Usage{} var usage = &dto.Usage{}
...@@ -99,5 +93,5 @@ func OaiResponsesStreamHandler(c *gin.Context, resp *http.Response, info *relayc ...@@ -99,5 +93,5 @@ func OaiResponsesStreamHandler(c *gin.Context, resp *http.Response, info *relayc
} }
} }
return nil, usage return usage, nil
} }
...@@ -9,6 +9,7 @@ import ( ...@@ -9,6 +9,7 @@ import (
"one-api/relay/channel" "one-api/relay/channel"
relaycommon "one-api/relay/common" relaycommon "one-api/relay/common"
"one-api/service" "one-api/service"
"one-api/types"
"github.com/gin-gonic/gin" "github.com/gin-gonic/gin"
) )
...@@ -70,13 +71,13 @@ func (a *Adaptor) DoRequest(c *gin.Context, info *relaycommon.RelayInfo, request ...@@ -70,13 +71,13 @@ func (a *Adaptor) DoRequest(c *gin.Context, info *relaycommon.RelayInfo, request
return channel.DoApiRequest(a, c, info, requestBody) return channel.DoApiRequest(a, c, info, requestBody)
} }
func (a *Adaptor) DoResponse(c *gin.Context, resp *http.Response, info *relaycommon.RelayInfo) (usage any, err *dto.OpenAIErrorWithStatusCode) { func (a *Adaptor) DoResponse(c *gin.Context, resp *http.Response, info *relaycommon.RelayInfo) (usage any, err *types.NewAPIError) {
if info.IsStream { if info.IsStream {
var responseText string var responseText string
err, responseText = palmStreamHandler(c, resp) err, responseText = palmStreamHandler(c, resp)
usage = service.ResponseText2Usage(responseText, info.UpstreamModelName, info.PromptTokens) usage = service.ResponseText2Usage(responseText, info.UpstreamModelName, info.PromptTokens)
} else { } else {
err, usage = palmHandler(c, resp, info.PromptTokens, info.UpstreamModelName) usage, err = palmHandler(c, info, resp)
} }
return return
} }
......
...@@ -2,14 +2,17 @@ package palm ...@@ -2,14 +2,17 @@ package palm
import ( import (
"encoding/json" "encoding/json"
"github.com/gin-gonic/gin"
"io" "io"
"net/http" "net/http"
"one-api/common" "one-api/common"
"one-api/constant" "one-api/constant"
"one-api/dto" "one-api/dto"
relaycommon "one-api/relay/common"
"one-api/relay/helper" "one-api/relay/helper"
"one-api/service" "one-api/service"
"one-api/types"
"github.com/gin-gonic/gin"
) )
// https://developers.generativeai.google/api/rest/generativelanguage/models/generateMessage#request-body // https://developers.generativeai.google/api/rest/generativelanguage/models/generateMessage#request-body
...@@ -70,7 +73,7 @@ func streamResponsePaLM2OpenAI(palmResponse *PaLMChatResponse) *dto.ChatCompleti ...@@ -70,7 +73,7 @@ func streamResponsePaLM2OpenAI(palmResponse *PaLMChatResponse) *dto.ChatCompleti
return &response return &response
} }
func palmStreamHandler(c *gin.Context, resp *http.Response) (*dto.OpenAIErrorWithStatusCode, string) { func palmStreamHandler(c *gin.Context, resp *http.Response) (*types.NewAPIError, string) {
responseText := "" responseText := ""
responseId := helper.GetResponseID(c) responseId := helper.GetResponseID(c)
createdTime := common.GetTimestamp() createdTime := common.GetTimestamp()
...@@ -121,42 +124,39 @@ func palmStreamHandler(c *gin.Context, resp *http.Response) (*dto.OpenAIErrorWit ...@@ -121,42 +124,39 @@ func palmStreamHandler(c *gin.Context, resp *http.Response) (*dto.OpenAIErrorWit
return nil, responseText return nil, responseText
} }
func palmHandler(c *gin.Context, resp *http.Response, promptTokens int, model string) (*dto.OpenAIErrorWithStatusCode, *dto.Usage) { func palmHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Response) (*dto.Usage, *types.NewAPIError) {
responseBody, err := io.ReadAll(resp.Body) responseBody, err := io.ReadAll(resp.Body)
if err != nil { if err != nil {
return service.OpenAIErrorWrapper(err, "read_response_body_failed", http.StatusInternalServerError), nil return nil, types.NewError(err, types.ErrorCodeReadResponseBodyFailed)
} }
common.CloseResponseBodyGracefully(resp) common.CloseResponseBodyGracefully(resp)
var palmResponse PaLMChatResponse var palmResponse PaLMChatResponse
err = json.Unmarshal(responseBody, &palmResponse) err = json.Unmarshal(responseBody, &palmResponse)
if err != nil { if err != nil {
return service.OpenAIErrorWrapper(err, "unmarshal_response_body_failed", http.StatusInternalServerError), nil return nil, types.NewError(err, types.ErrorCodeBadResponseBody)
} }
if palmResponse.Error.Code != 0 || len(palmResponse.Candidates) == 0 { if palmResponse.Error.Code != 0 || len(palmResponse.Candidates) == 0 {
return &dto.OpenAIErrorWithStatusCode{ return nil, types.WithOpenAIError(types.OpenAIError{
Error: dto.OpenAIError{ Message: palmResponse.Error.Message,
Message: palmResponse.Error.Message, Type: palmResponse.Error.Status,
Type: palmResponse.Error.Status, Param: "",
Param: "", Code: palmResponse.Error.Code,
Code: palmResponse.Error.Code, }, resp.StatusCode)
},
StatusCode: resp.StatusCode,
}, nil
} }
fullTextResponse := responsePaLM2OpenAI(&palmResponse) fullTextResponse := responsePaLM2OpenAI(&palmResponse)
completionTokens := service.CountTextToken(palmResponse.Candidates[0].Content, model) completionTokens := service.CountTextToken(palmResponse.Candidates[0].Content, info.UpstreamModelName)
usage := dto.Usage{ usage := dto.Usage{
PromptTokens: promptTokens, PromptTokens: info.PromptTokens,
CompletionTokens: completionTokens, CompletionTokens: completionTokens,
TotalTokens: promptTokens + completionTokens, TotalTokens: info.PromptTokens + completionTokens,
} }
fullTextResponse.Usage = usage fullTextResponse.Usage = usage
jsonResponse, err := json.Marshal(fullTextResponse) jsonResponse, err := common.Marshal(fullTextResponse)
if err != nil { if err != nil {
return service.OpenAIErrorWrapper(err, "marshal_response_body_failed", http.StatusInternalServerError), nil return nil, types.NewError(err, types.ErrorCodeBadResponseBody)
} }
c.Writer.Header().Set("Content-Type", "application/json") c.Writer.Header().Set("Content-Type", "application/json")
c.Writer.WriteHeader(resp.StatusCode) c.Writer.WriteHeader(resp.StatusCode)
_, err = c.Writer.Write(jsonResponse) common.IOCopyBytesGracefully(c, resp, jsonResponse)
return nil, &usage return &usage, nil
} }
...@@ -9,6 +9,7 @@ import ( ...@@ -9,6 +9,7 @@ import (
"one-api/relay/channel" "one-api/relay/channel"
"one-api/relay/channel/openai" "one-api/relay/channel/openai"
relaycommon "one-api/relay/common" relaycommon "one-api/relay/common"
"one-api/types"
"github.com/gin-gonic/gin" "github.com/gin-gonic/gin"
) )
...@@ -73,11 +74,11 @@ func (a *Adaptor) DoRequest(c *gin.Context, info *relaycommon.RelayInfo, request ...@@ -73,11 +74,11 @@ func (a *Adaptor) DoRequest(c *gin.Context, info *relaycommon.RelayInfo, request
return channel.DoApiRequest(a, c, info, requestBody) return channel.DoApiRequest(a, c, info, requestBody)
} }
func (a *Adaptor) DoResponse(c *gin.Context, resp *http.Response, info *relaycommon.RelayInfo) (usage any, err *dto.OpenAIErrorWithStatusCode) { 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 = openai.OaiStreamHandler(c, resp, info) usage, err = openai.OaiStreamHandler(c, info, resp)
} else { } else {
err, usage = openai.OpenaiHandler(c, resp, info) usage, err = openai.OpenaiHandler(c, info, resp)
} }
return return
} }
......
...@@ -10,6 +10,7 @@ import ( ...@@ -10,6 +10,7 @@ import (
"one-api/relay/channel/openai" "one-api/relay/channel/openai"
relaycommon "one-api/relay/common" relaycommon "one-api/relay/common"
"one-api/relay/constant" "one-api/relay/constant"
"one-api/types"
"github.com/gin-gonic/gin" "github.com/gin-gonic/gin"
) )
...@@ -76,20 +77,20 @@ func (a *Adaptor) ConvertEmbeddingRequest(c *gin.Context, info *relaycommon.Rela ...@@ -76,20 +77,20 @@ func (a *Adaptor) ConvertEmbeddingRequest(c *gin.Context, info *relaycommon.Rela
return request, nil return request, nil
} }
func (a *Adaptor) DoResponse(c *gin.Context, resp *http.Response, info *relaycommon.RelayInfo) (usage any, err *dto.OpenAIErrorWithStatusCode) { func (a *Adaptor) DoResponse(c *gin.Context, resp *http.Response, info *relaycommon.RelayInfo) (usage any, err *types.NewAPIError) {
switch info.RelayMode { switch info.RelayMode {
case constant.RelayModeRerank: case constant.RelayModeRerank:
err, usage = siliconflowRerankHandler(c, resp) usage, err = siliconflowRerankHandler(c, info, resp)
case constant.RelayModeCompletions: case constant.RelayModeCompletions:
fallthrough fallthrough
case constant.RelayModeChatCompletions: case constant.RelayModeChatCompletions:
if info.IsStream { if info.IsStream {
err, usage = openai.OaiStreamHandler(c, resp, info) usage, err = openai.OaiStreamHandler(c, info, resp)
} else { } else {
err, usage = openai.OpenaiHandler(c, resp, info) usage, err = openai.OpenaiHandler(c, info, resp)
} }
case constant.RelayModeEmbeddings: case constant.RelayModeEmbeddings:
err, usage = openai.OpenaiHandler(c, resp, info) usage, err = openai.OpenaiHandler(c, info, resp)
} }
return return
} }
......
...@@ -2,24 +2,26 @@ package siliconflow ...@@ -2,24 +2,26 @@ package siliconflow
import ( import (
"encoding/json" "encoding/json"
"github.com/gin-gonic/gin"
"io" "io"
"net/http" "net/http"
"one-api/common" "one-api/common"
"one-api/dto" "one-api/dto"
"one-api/service" relaycommon "one-api/relay/common"
"one-api/types"
"github.com/gin-gonic/gin"
) )
func siliconflowRerankHandler(c *gin.Context, resp *http.Response) (*dto.OpenAIErrorWithStatusCode, *dto.Usage) { func siliconflowRerankHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Response) (*dto.Usage, *types.NewAPIError) {
responseBody, err := io.ReadAll(resp.Body) responseBody, err := io.ReadAll(resp.Body)
if err != nil { if err != nil {
return service.OpenAIErrorWrapper(err, "read_response_body_failed", http.StatusInternalServerError), nil return nil, types.NewError(err, types.ErrorCodeReadResponseBodyFailed)
} }
common.CloseResponseBodyGracefully(resp) common.CloseResponseBodyGracefully(resp)
var siliconflowResp SFRerankResponse var siliconflowResp SFRerankResponse
err = json.Unmarshal(responseBody, &siliconflowResp) err = json.Unmarshal(responseBody, &siliconflowResp)
if err != nil { if err != nil {
return service.OpenAIErrorWrapper(err, "unmarshal_response_body_failed", http.StatusInternalServerError), nil return nil, types.NewError(err, types.ErrorCodeBadResponseBody)
} }
usage := &dto.Usage{ usage := &dto.Usage{
PromptTokens: siliconflowResp.Meta.Tokens.InputTokens, PromptTokens: siliconflowResp.Meta.Tokens.InputTokens,
...@@ -33,10 +35,10 @@ func siliconflowRerankHandler(c *gin.Context, resp *http.Response) (*dto.OpenAIE ...@@ -33,10 +35,10 @@ func siliconflowRerankHandler(c *gin.Context, resp *http.Response) (*dto.OpenAIE
jsonResponse, err := json.Marshal(rerankResp) jsonResponse, err := json.Marshal(rerankResp)
if err != nil { if err != nil {
return service.OpenAIErrorWrapper(err, "marshal_response_body_failed", http.StatusInternalServerError), nil return nil, types.NewError(err, types.ErrorCodeBadResponseBody)
} }
c.Writer.Header().Set("Content-Type", "application/json") c.Writer.Header().Set("Content-Type", "application/json")
c.Writer.WriteHeader(resp.StatusCode) c.Writer.WriteHeader(resp.StatusCode)
_, err = c.Writer.Write(jsonResponse) common.IOCopyBytesGracefully(c, resp, jsonResponse)
return nil, usage return usage, nil
} }
...@@ -6,10 +6,11 @@ import ( ...@@ -6,10 +6,11 @@ import (
"io" "io"
"net/http" "net/http"
"one-api/common" "one-api/common"
"one-api/constant"
"one-api/dto" "one-api/dto"
"one-api/relay/channel" "one-api/relay/channel"
relaycommon "one-api/relay/common" relaycommon "one-api/relay/common"
"one-api/service" "one-api/types"
"strconv" "strconv"
"strings" "strings"
...@@ -63,7 +64,7 @@ func (a *Adaptor) ConvertOpenAIRequest(c *gin.Context, info *relaycommon.RelayIn ...@@ -63,7 +64,7 @@ func (a *Adaptor) ConvertOpenAIRequest(c *gin.Context, info *relaycommon.RelayIn
if request == nil { if request == nil {
return nil, errors.New("request is nil") return nil, errors.New("request is nil")
} }
apiKey := c.Request.Header.Get("Authorization") apiKey := common.GetContextKeyString(c, constant.ContextKeyChannelKey)
apiKey = strings.TrimPrefix(apiKey, "Bearer ") apiKey = strings.TrimPrefix(apiKey, "Bearer ")
appId, secretId, secretKey, err := parseTencentConfig(apiKey) appId, secretId, secretKey, err := parseTencentConfig(apiKey)
a.AppID = appId a.AppID = appId
...@@ -94,13 +95,11 @@ func (a *Adaptor) DoRequest(c *gin.Context, info *relaycommon.RelayInfo, request ...@@ -94,13 +95,11 @@ func (a *Adaptor) DoRequest(c *gin.Context, info *relaycommon.RelayInfo, request
return channel.DoApiRequest(a, c, info, requestBody) return channel.DoApiRequest(a, c, info, requestBody)
} }
func (a *Adaptor) DoResponse(c *gin.Context, resp *http.Response, info *relaycommon.RelayInfo) (usage any, err *dto.OpenAIErrorWithStatusCode) { func (a *Adaptor) DoResponse(c *gin.Context, resp *http.Response, info *relaycommon.RelayInfo) (usage any, err *types.NewAPIError) {
if info.IsStream { if info.IsStream {
var responseText string usage, err = tencentStreamHandler(c, info, resp)
err, responseText = tencentStreamHandler(c, resp)
usage = service.ResponseText2Usage(responseText, info.UpstreamModelName, info.PromptTokens)
} else { } else {
err, usage = tencentHandler(c, resp) usage, err = tencentHandler(c, info, resp)
} }
return return
} }
......
...@@ -8,17 +8,20 @@ import ( ...@@ -8,17 +8,20 @@ import (
"encoding/json" "encoding/json"
"errors" "errors"
"fmt" "fmt"
"github.com/gin-gonic/gin"
"io" "io"
"net/http" "net/http"
"one-api/common" "one-api/common"
"one-api/constant" "one-api/constant"
"one-api/dto" "one-api/dto"
relaycommon "one-api/relay/common"
"one-api/relay/helper" "one-api/relay/helper"
"one-api/service" "one-api/service"
"one-api/types"
"strconv" "strconv"
"strings" "strings"
"time" "time"
"github.com/gin-gonic/gin"
) )
// https://cloud.tencent.com/document/product/1729/97732 // https://cloud.tencent.com/document/product/1729/97732
...@@ -86,7 +89,7 @@ func streamResponseTencent2OpenAI(TencentResponse *TencentChatResponse) *dto.Cha ...@@ -86,7 +89,7 @@ func streamResponseTencent2OpenAI(TencentResponse *TencentChatResponse) *dto.Cha
return &response return &response
} }
func tencentStreamHandler(c *gin.Context, resp *http.Response) (*dto.OpenAIErrorWithStatusCode, string) { func tencentStreamHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Response) (*dto.Usage, *types.NewAPIError) {
var responseText string var responseText string
scanner := bufio.NewScanner(resp.Body) scanner := bufio.NewScanner(resp.Body)
scanner.Split(bufio.ScanLines) scanner.Split(bufio.ScanLines)
...@@ -126,38 +129,35 @@ func tencentStreamHandler(c *gin.Context, resp *http.Response) (*dto.OpenAIError ...@@ -126,38 +129,35 @@ func tencentStreamHandler(c *gin.Context, resp *http.Response) (*dto.OpenAIError
common.CloseResponseBodyGracefully(resp) common.CloseResponseBodyGracefully(resp)
return nil, responseText return service.ResponseText2Usage(responseText, info.UpstreamModelName, info.PromptTokens), nil
} }
func tencentHandler(c *gin.Context, resp *http.Response) (*dto.OpenAIErrorWithStatusCode, *dto.Usage) { func tencentHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Response) (*dto.Usage, *types.NewAPIError) {
var tencentSb TencentChatResponseSB var tencentSb TencentChatResponseSB
responseBody, err := io.ReadAll(resp.Body) responseBody, err := io.ReadAll(resp.Body)
if err != nil { if err != nil {
return service.OpenAIErrorWrapper(err, "read_response_body_failed", http.StatusInternalServerError), nil return nil, types.NewError(err, types.ErrorCodeReadResponseBodyFailed)
} }
common.CloseResponseBodyGracefully(resp) common.CloseResponseBodyGracefully(resp)
err = json.Unmarshal(responseBody, &tencentSb) err = json.Unmarshal(responseBody, &tencentSb)
if err != nil { if err != nil {
return service.OpenAIErrorWrapper(err, "unmarshal_response_body_failed", http.StatusInternalServerError), nil return nil, types.NewError(err, types.ErrorCodeBadResponseBody)
} }
if tencentSb.Response.Error.Code != 0 { if tencentSb.Response.Error.Code != 0 {
return &dto.OpenAIErrorWithStatusCode{ return nil, types.WithOpenAIError(types.OpenAIError{
Error: dto.OpenAIError{ Message: tencentSb.Response.Error.Message,
Message: tencentSb.Response.Error.Message, Code: tencentSb.Response.Error.Code,
Code: tencentSb.Response.Error.Code, }, resp.StatusCode)
},
StatusCode: resp.StatusCode,
}, nil
} }
fullTextResponse := responseTencent2OpenAI(&tencentSb.Response) fullTextResponse := responseTencent2OpenAI(&tencentSb.Response)
jsonResponse, err := json.Marshal(fullTextResponse) jsonResponse, err := common.Marshal(fullTextResponse)
if err != nil { if err != nil {
return service.OpenAIErrorWrapper(err, "marshal_response_body_failed", http.StatusInternalServerError), nil return nil, types.NewError(err, types.ErrorCodeBadResponseBody)
} }
c.Writer.Header().Set("Content-Type", "application/json") c.Writer.Header().Set("Content-Type", "application/json")
c.Writer.WriteHeader(resp.StatusCode) c.Writer.WriteHeader(resp.StatusCode)
_, err = c.Writer.Write(jsonResponse) common.IOCopyBytesGracefully(c, resp, jsonResponse)
return nil, &fullTextResponse.Usage return &fullTextResponse.Usage, nil
} }
func parseTencentConfig(config string) (appId int64, secretId string, secretKey string, err error) { func parseTencentConfig(config string) (appId int64, secretId string, secretKey string, err error) {
......
...@@ -14,6 +14,7 @@ import ( ...@@ -14,6 +14,7 @@ import (
relaycommon "one-api/relay/common" relaycommon "one-api/relay/common"
"one-api/relay/constant" "one-api/relay/constant"
"one-api/setting/model_setting" "one-api/setting/model_setting"
"one-api/types"
"strings" "strings"
"github.com/gin-gonic/gin" "github.com/gin-gonic/gin"
...@@ -208,19 +209,19 @@ func (a *Adaptor) DoRequest(c *gin.Context, info *relaycommon.RelayInfo, request ...@@ -208,19 +209,19 @@ func (a *Adaptor) DoRequest(c *gin.Context, info *relaycommon.RelayInfo, request
return channel.DoApiRequest(a, c, info, requestBody) return channel.DoApiRequest(a, c, info, requestBody)
} }
func (a *Adaptor) DoResponse(c *gin.Context, resp *http.Response, info *relaycommon.RelayInfo) (usage any, err *dto.OpenAIErrorWithStatusCode) { func (a *Adaptor) DoResponse(c *gin.Context, resp *http.Response, info *relaycommon.RelayInfo) (usage any, err *types.NewAPIError) {
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) err, usage = 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, resp, info) usage, err = gemini.GeminiTextGenerationStreamHandler(c, info, resp)
} else { } else {
err, usage = gemini.GeminiChatStreamHandler(c, resp, info) usage, err = gemini.GeminiChatStreamHandler(c, info, resp)
} }
case RequestModeLlama: case RequestModeLlama:
err, usage = openai.OaiStreamHandler(c, resp, info) usage, err = openai.OaiStreamHandler(c, info, resp)
} }
} else { } else {
switch a.RequestMode { switch a.RequestMode {
...@@ -228,12 +229,12 @@ func (a *Adaptor) DoResponse(c *gin.Context, resp *http.Response, info *relaycom ...@@ -228,12 +229,12 @@ func (a *Adaptor) DoResponse(c *gin.Context, resp *http.Response, info *relaycom
err, usage = claude.ClaudeHandler(c, resp, claude.RequestModeMessage, info) err, usage = claude.ClaudeHandler(c, resp, claude.RequestModeMessage, info)
case RequestModeGemini: case RequestModeGemini:
if info.RelayMode == constant.RelayModeGemini { if info.RelayMode == constant.RelayModeGemini {
usage, err = gemini.GeminiTextGenerationHandler(c, resp, info) usage, err = gemini.GeminiTextGenerationHandler(c, info, resp)
} else { } else {
err, usage = gemini.GeminiChatHandler(c, resp, info) usage, err = gemini.GeminiChatHandler(c, info, resp)
} }
case RequestModeLlama: case RequestModeLlama:
err, usage = openai.OpenaiHandler(c, resp, info) usage, err = openai.OpenaiHandler(c, info, resp)
} }
} }
return return
......
...@@ -4,8 +4,11 @@ import "one-api/common" ...@@ -4,8 +4,11 @@ import "one-api/common"
func GetModelRegion(other string, localModelName string) string { func GetModelRegion(other string, localModelName string) string {
// if other is json string // if other is json string
if common.IsJsonStr(other) { if common.IsJsonObject(other) {
m := common.StrToMap(other) m, err := common.StrToMap(other)
if err != nil {
return other // return original if parsing fails
}
if m[localModelName] != nil { if m[localModelName] != nil {
return m[localModelName].(string) return m[localModelName].(string)
} else { } else {
......
...@@ -13,6 +13,7 @@ import ( ...@@ -13,6 +13,7 @@ import (
"one-api/relay/channel/openai" "one-api/relay/channel/openai"
relaycommon "one-api/relay/common" relaycommon "one-api/relay/common"
"one-api/relay/constant" "one-api/relay/constant"
"one-api/types"
"path/filepath" "path/filepath"
"strings" "strings"
...@@ -225,18 +226,18 @@ func (a *Adaptor) DoRequest(c *gin.Context, info *relaycommon.RelayInfo, request ...@@ -225,18 +226,18 @@ func (a *Adaptor) DoRequest(c *gin.Context, info *relaycommon.RelayInfo, request
return channel.DoApiRequest(a, c, info, requestBody) return channel.DoApiRequest(a, c, info, requestBody)
} }
func (a *Adaptor) DoResponse(c *gin.Context, resp *http.Response, info *relaycommon.RelayInfo) (usage any, err *dto.OpenAIErrorWithStatusCode) { func (a *Adaptor) DoResponse(c *gin.Context, resp *http.Response, info *relaycommon.RelayInfo) (usage any, err *types.NewAPIError) {
switch info.RelayMode { switch info.RelayMode {
case constant.RelayModeChatCompletions: case constant.RelayModeChatCompletions:
if info.IsStream { if info.IsStream {
err, usage = openai.OaiStreamHandler(c, resp, info) usage, err = openai.OaiStreamHandler(c, info, resp)
} else { } else {
err, usage = openai.OpenaiHandler(c, resp, info) usage, err = openai.OpenaiHandler(c, info, resp)
} }
case constant.RelayModeEmbeddings: case constant.RelayModeEmbeddings:
err, usage = openai.OpenaiHandler(c, resp, info) usage, err = openai.OpenaiHandler(c, info, resp)
case constant.RelayModeImagesGenerations, constant.RelayModeImagesEdits: case constant.RelayModeImagesGenerations, constant.RelayModeImagesEdits:
err, usage = openai.OpenaiHandlerWithUsage(c, resp, info) usage, err = openai.OpenaiHandlerWithUsage(c, info, resp)
} }
return return
} }
......
...@@ -8,6 +8,7 @@ import ( ...@@ -8,6 +8,7 @@ import (
"one-api/relay/channel" "one-api/relay/channel"
"one-api/relay/channel/openai" "one-api/relay/channel/openai"
relaycommon "one-api/relay/common" relaycommon "one-api/relay/common"
"one-api/types"
"strings" "strings"
"one-api/relay/constant" "one-api/relay/constant"
...@@ -56,6 +57,15 @@ func (a *Adaptor) ConvertOpenAIRequest(c *gin.Context, info *relaycommon.RelayIn ...@@ -56,6 +57,15 @@ func (a *Adaptor) ConvertOpenAIRequest(c *gin.Context, info *relaycommon.RelayIn
if request == nil { if request == nil {
return nil, errors.New("request is nil") return nil, errors.New("request is nil")
} }
if strings.HasSuffix(info.UpstreamModelName, "-search") {
info.UpstreamModelName = strings.TrimSuffix(info.UpstreamModelName, "-search")
request.Model = info.UpstreamModelName
toMap := request.ToMap()
toMap["search_parameters"] = map[string]any{
"mode": "on",
}
return toMap, nil
}
if strings.HasPrefix(request.Model, "grok-3-mini") { if strings.HasPrefix(request.Model, "grok-3-mini") {
if request.MaxCompletionTokens == 0 && request.MaxTokens != 0 { if request.MaxCompletionTokens == 0 && request.MaxTokens != 0 {
request.MaxCompletionTokens = request.MaxTokens request.MaxCompletionTokens = request.MaxTokens
...@@ -95,15 +105,15 @@ func (a *Adaptor) DoRequest(c *gin.Context, info *relaycommon.RelayInfo, request ...@@ -95,15 +105,15 @@ func (a *Adaptor) DoRequest(c *gin.Context, info *relaycommon.RelayInfo, request
return channel.DoApiRequest(a, c, info, requestBody) return channel.DoApiRequest(a, c, info, requestBody)
} }
func (a *Adaptor) DoResponse(c *gin.Context, resp *http.Response, info *relaycommon.RelayInfo) (usage any, err *dto.OpenAIErrorWithStatusCode) { func (a *Adaptor) DoResponse(c *gin.Context, resp *http.Response, info *relaycommon.RelayInfo) (usage any, err *types.NewAPIError) {
switch info.RelayMode { switch info.RelayMode {
case constant.RelayModeImagesGenerations, constant.RelayModeImagesEdits: case constant.RelayModeImagesGenerations, constant.RelayModeImagesEdits:
err, usage = openai.OpenaiHandlerWithUsage(c, resp, info) usage, err = openai.OpenaiHandlerWithUsage(c, info, resp)
default: default:
if info.IsStream { if info.IsStream {
err, usage = xAIStreamHandler(c, resp, info) usage, err = xAIStreamHandler(c, info, resp)
} else { } else {
err, usage = xAIHandler(c, resp, info) usage, err = xAIHandler(c, info, resp)
} }
} }
return return
......
This diff is collapsed. Click to expand it.
This source diff could not be displayed because it is too large. You can view the blob instead.
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