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 {
// 兼容
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 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"
ContextKeyChannelType ContextKey = "channel_type"
ContextKeyChannelId ContextKey = "channel_id" ContextKeyChannelId ContextKey = "channel_id"
ContextKeyChannelName ContextKey = "channel_name"
ContextKeyChannelCreateTime ContextKey = "channel_create_time"
ContextKeyChannelBaseUrl ContextKey = "base_url"
ContextKeyChannelType ContextKey = "channel_type"
ContextKeyChannelSetting ContextKey = "channel_setting" ContextKeyChannelSetting ContextKey = "channel_setting"
ContextKeyParamOverride ContextKey = "param_override" 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,10 +174,7 @@ func UpdateRedemption(c *gin.Context) { ...@@ -229,10 +174,7 @@ 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{
......
...@@ -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,7 +16,7 @@ type OpenAIErrorWithStatusCode struct { ...@@ -14,7 +16,7 @@ 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"`
......
...@@ -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"
...@@ -25,7 +27,7 @@ type RealtimeEvent struct { ...@@ -25,7 +27,7 @@ type RealtimeEvent struct {
//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"`
......
package middleware package middleware
import ( import (
"fmt"
"net/http" "net/http"
"one-api/common" "one-api/common"
"one-api/model" "one-api/model"
...@@ -233,6 +234,18 @@ func TokenAuth() func(c *gin.Context) { ...@@ -233,6 +234,18 @@ func TokenAuth() func(c *gin.Context) {
userCache.WriteContext(c) userCache.WriteContext(c)
err = SetupContextForToken(c, token, parts...)
if err != nil {
return
}
c.Next()
}
}
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("id", token.UserId)
c.Set("token_id", token.Id) c.Set("token_id", token.Id)
c.Set("token_key", token.Key) c.Set("token_key", token.Key)
...@@ -254,9 +267,8 @@ func TokenAuth() func(c *gin.Context) { ...@@ -254,9 +267,8 @@ func TokenAuth() func(c *gin.Context) {
c.Set("specific_channel_id", parts[1]) c.Set("specific_channel_id", parts[1])
} else { } else {
abortWithOpenAiMessage(c, http.StatusForbidden, "普通用户不支持指定渠道") abortWithOpenAiMessage(c, http.StatusForbidden, "普通用户不支持指定渠道")
return return fmt.Errorf("普通用户不支持指定渠道")
} }
} }
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 {
if channel, ok := channelsIDM[channelId]; ok {
uniquePriorities[int(channel.GetPriority())] = true 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,10 +164,14 @@ func getRandomSatisfiedChannel(group string, model string, retry int) (*Channel, ...@@ -152,10 +164,14 @@ 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, ok := channelsIDM[channelId]; ok {
if channel.GetPriority() == targetPriority { if channel.GetPriority() == targetPriority {
targetChannels = append(targetChannels, channel) 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: aliResponse.Code, Type: "ali_error",
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,40 +112,11 @@ func embeddingResponseBaidu2OpenAI(response *BaiduEmbeddingResponse) *dto.OpenAI ...@@ -110,40 +112,11 @@ 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) {
if atEOF && len(data) == 0 {
return 0, nil, nil
}
if i := strings.Index(string(data), "\n"); i >= 0 {
return i + 1, data[0:i], nil
}
if atEOF {
return len(data), data, nil
}
return 0, nil, nil
})
dataChan := make(chan string)
stopChan := make(chan bool)
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
}()
helper.SetEventStreamHeaders(c)
c.Stream(func(w io.Writer) bool {
select {
case data := <-dataChan:
var baiduResponse BaiduChatStreamResponse var baiduResponse BaiduChatStreamResponse
err := json.Unmarshal([]byte(data), &baiduResponse) err := common.Unmarshal([]byte(data), &baiduResponse)
if err != nil { if err != nil {
common.SysError("error unmarshalling stream response: " + err.Error()) common.SysError("error unmarshalling stream response: " + err.Error())
return true return true
...@@ -154,48 +127,34 @@ func baiduStreamHandler(c *gin.Context, resp *http.Response) (*dto.OpenAIErrorWi ...@@ -154,48 +127,34 @@ func baiduStreamHandler(c *gin.Context, resp *http.Response) (*dto.OpenAIErrorWi
usage.CompletionTokens = baiduResponse.Usage.TotalTokens - baiduResponse.Usage.PromptTokens usage.CompletionTokens = baiduResponse.Usage.TotalTokens - baiduResponse.Usage.PromptTokens
} }
response := streamResponseBaidu2OpenAI(&baiduResponse) response := streamResponseBaidu2OpenAI(&baiduResponse)
jsonResponse, err := json.Marshal(response) err = helper.ObjectData(c, response)
if err != nil { if err != nil {
common.SysError("error marshalling stream response: " + err.Error()) common.SysError("error sending stream response: " + err.Error())
return true
} }
c.Render(-1, common.CustomEvent{Data: "data: " + string(jsonResponse)})
return true return true
case <-stopChan:
c.Render(-1, common.CustomEvent{Data: "data: [DONE]"})
return false
}
}) })
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
} }
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