Commit 97ea8b65 by CaIon

refactor: Introduce pre-consume quota and unify relay handlers

This commit introduces a major architectural refactoring to improve quota management, centralize logging, and streamline the relay handling logic.

Key changes:
- **Pre-consume Quota:** Implements a new mechanism to check and reserve user quota *before* making the request to the upstream provider. This ensures more accurate quota deduction and prevents users from exceeding their limits due to concurrent requests.

- **Unified Relay Handlers:** Refactors the relay logic to use generic handlers (e.g., `ChatHandler`, `ImageHandler`) instead of provider-specific implementations. This significantly reduces code duplication and simplifies adding new channels.

- **Centralized Logger:** A new dedicated `logger` package is introduced, and all system logging calls are migrated to use it, moving this responsibility out of the `common` package.

- **Code Reorganization:** DTOs are generalized (e.g., `dalle.go` -> `openai_image.go`) and utility code is moved to more appropriate packages (e.g., `common/http.go` -> `service/http.go`) for better code structure.
parent c7281a35
...@@ -5,7 +5,7 @@ import ( ...@@ -5,7 +5,7 @@ import (
_ "embed" _ "embed"
"fmt" "fmt"
"github.com/go-redis/redis/v8" "github.com/go-redis/redis/v8"
"one-api/common" "one-api/logger"
"sync" "sync"
) )
...@@ -27,7 +27,7 @@ func New(ctx context.Context, r *redis.Client) *RedisLimiter { ...@@ -27,7 +27,7 @@ func New(ctx context.Context, r *redis.Client) *RedisLimiter {
// 预加载脚本 // 预加载脚本
limitSHA, err := r.ScriptLoad(ctx, rateLimitScript).Result() limitSHA, err := r.ScriptLoad(ctx, rateLimitScript).Result()
if err != nil { if err != nil {
common.SysLog(fmt.Sprintf("Failed to load rate limit script: %v", err)) logger.SysLog(fmt.Sprintf("Failed to load rate limit script: %v", err))
} }
instance = &RedisLimiter{ instance = &RedisLimiter{
client: r, client: r,
......
package common package common
import ( import (
"context"
"encoding/json"
"fmt" "fmt"
"github.com/bytedance/gopkg/util/gopool"
"github.com/gin-gonic/gin" "github.com/gin-gonic/gin"
"io"
"log"
"os" "os"
"path/filepath"
"sync"
"time" "time"
) )
const (
loggerINFO = "INFO"
loggerWarn = "WARN"
loggerError = "ERR"
)
const maxLogCount = 1000000
var logCount int
var setupLogLock sync.Mutex
var setupLogWorking bool
func SetupLogger() {
if *LogDir != "" {
ok := setupLogLock.TryLock()
if !ok {
log.Println("setup log is already working")
return
}
defer func() {
setupLogLock.Unlock()
setupLogWorking = false
}()
logPath := filepath.Join(*LogDir, fmt.Sprintf("oneapi-%s.log", time.Now().Format("20060102150405")))
fd, err := os.OpenFile(logPath, os.O_APPEND|os.O_CREATE|os.O_WRONLY, 0644)
if err != nil {
log.Fatal("failed to open log file")
}
gin.DefaultWriter = io.MultiWriter(os.Stdout, fd)
gin.DefaultErrorWriter = io.MultiWriter(os.Stderr, fd)
}
}
func SysLog(s string) { func SysLog(s string) {
t := time.Now() t := time.Now()
_, _ = fmt.Fprintf(gin.DefaultWriter, "[SYS] %v | %s \n", t.Format("2006/01/02 - 15:04:05"), s) _, _ = fmt.Fprintf(gin.DefaultWriter, "[SYS] %v | %s \n", t.Format("2006/01/02 - 15:04:05"), s)
...@@ -57,67 +17,8 @@ func SysError(s string) { ...@@ -57,67 +17,8 @@ func SysError(s string) {
_, _ = fmt.Fprintf(gin.DefaultErrorWriter, "[SYS] %v | %s \n", t.Format("2006/01/02 - 15:04:05"), s) _, _ = fmt.Fprintf(gin.DefaultErrorWriter, "[SYS] %v | %s \n", t.Format("2006/01/02 - 15:04:05"), s)
} }
func LogInfo(ctx context.Context, msg string) {
logHelper(ctx, loggerINFO, msg)
}
func LogWarn(ctx context.Context, msg string) {
logHelper(ctx, loggerWarn, msg)
}
func LogError(ctx context.Context, msg string) {
logHelper(ctx, loggerError, msg)
}
func logHelper(ctx context.Context, level string, msg string) {
writer := gin.DefaultErrorWriter
if level == loggerINFO {
writer = gin.DefaultWriter
}
id := ctx.Value(RequestIdKey)
if id == nil {
id = "SYSTEM"
}
now := time.Now()
_, _ = fmt.Fprintf(writer, "[%s] %v | %s | %s \n", level, now.Format("2006/01/02 - 15:04:05"), id, msg)
logCount++ // we don't need accurate count, so no lock here
if logCount > maxLogCount && !setupLogWorking {
logCount = 0
setupLogWorking = true
gopool.Go(func() {
SetupLogger()
})
}
}
func FatalLog(v ...any) { func FatalLog(v ...any) {
t := time.Now() t := time.Now()
_, _ = fmt.Fprintf(gin.DefaultErrorWriter, "[FATAL] %v | %v \n", t.Format("2006/01/02 - 15:04:05"), v) _, _ = fmt.Fprintf(gin.DefaultErrorWriter, "[FATAL] %v | %v \n", t.Format("2006/01/02 - 15:04:05"), v)
os.Exit(1) os.Exit(1)
} }
func LogQuota(quota int) string {
if DisplayInCurrencyEnabled {
return fmt.Sprintf("$%.6f 额度", float64(quota)/QuotaPerUnit)
} else {
return fmt.Sprintf("%d 点额度", quota)
}
}
func FormatQuota(quota int) string {
if DisplayInCurrencyEnabled {
return fmt.Sprintf("$%.6f", float64(quota)/QuotaPerUnit)
} else {
return fmt.Sprintf("%d", quota)
}
}
// LogJson 仅供测试使用 only for test
func LogJson(ctx context.Context, msg string, obj any) {
jsonStr, err := json.Marshal(obj)
if err != nil {
LogError(ctx, fmt.Sprintf("json marshal failed: %s", err.Error()))
return
}
LogInfo(ctx, fmt.Sprintf("%s | %s", msg, string(jsonStr)))
}
...@@ -3,6 +3,8 @@ package constant ...@@ -3,6 +3,8 @@ package constant
type ContextKey string type ContextKey string
const ( const (
ContextKeyPromptTokens ContextKey = "prompt_tokens"
ContextKeyOriginalModel ContextKey = "original_model" ContextKeyOriginalModel ContextKey = "original_model"
ContextKeyRequestStartTime ContextKey = "request_start_time" ContextKeyRequestStartTime ContextKey = "request_start_time"
......
...@@ -8,6 +8,7 @@ import ( ...@@ -8,6 +8,7 @@ import (
"net/http" "net/http"
"one-api/common" "one-api/common"
"one-api/constant" "one-api/constant"
"one-api/logger"
"one-api/model" "one-api/model"
"one-api/service" "one-api/service"
"one-api/setting" "one-api/setting"
...@@ -485,8 +486,8 @@ func UpdateAllChannelsBalance(c *gin.Context) { ...@@ -485,8 +486,8 @@ func UpdateAllChannelsBalance(c *gin.Context) {
func AutomaticallyUpdateChannels(frequency int) { func AutomaticallyUpdateChannels(frequency int) {
for { for {
time.Sleep(time.Duration(frequency) * time.Minute) time.Sleep(time.Duration(frequency) * time.Minute)
common.SysLog("updating all channels") logger.SysLog("updating all channels")
_ = updateAllChannelsBalance() _ = updateAllChannelsBalance()
common.SysLog("channels update done") logger.SysLog("channels update done")
} }
} }
...@@ -13,6 +13,7 @@ import ( ...@@ -13,6 +13,7 @@ import (
"one-api/common" "one-api/common"
"one-api/constant" "one-api/constant"
"one-api/dto" "one-api/dto"
"one-api/logger"
"one-api/middleware" "one-api/middleware"
"one-api/model" "one-api/model"
"one-api/relay" "one-api/relay"
...@@ -159,7 +160,7 @@ func testChannel(channel *model.Channel, testModel string) testResult { ...@@ -159,7 +160,7 @@ func testChannel(channel *model.Channel, testModel string) testResult {
// 创建一个用于日志的 info 副本,移除 ApiKey // 创建一个用于日志的 info 副本,移除 ApiKey
logInfo := *info logInfo := *info
logInfo.ApiKey = "" logInfo.ApiKey = ""
common.SysLog(fmt.Sprintf("testing channel %d with model %s , info %+v ", channel.Id, testModel, logInfo)) logger.SysLog(fmt.Sprintf("testing channel %d with model %s , info %+v ", channel.Id, testModel, logInfo))
priceData, err := helper.ModelPriceHelper(c, info, 0, int(request.GetMaxTokens())) priceData, err := helper.ModelPriceHelper(c, info, 0, int(request.GetMaxTokens()))
if err != nil { if err != nil {
...@@ -279,7 +280,7 @@ func testChannel(channel *model.Channel, testModel string) testResult { ...@@ -279,7 +280,7 @@ func testChannel(channel *model.Channel, testModel string) testResult {
Group: info.UsingGroup, Group: info.UsingGroup,
Other: other, Other: other,
}) })
common.SysLog(fmt.Sprintf("testing channel #%d, response: \n%s", channel.Id, string(respBody))) logger.SysLog(fmt.Sprintf("testing channel #%d, response: \n%s", channel.Id, string(respBody)))
return testResult{ return testResult{
context: c, context: c,
localErr: nil, localErr: nil,
...@@ -461,13 +462,13 @@ func TestAllChannels(c *gin.Context) { ...@@ -461,13 +462,13 @@ func TestAllChannels(c *gin.Context) {
func AutomaticallyTestChannels(frequency int) { func AutomaticallyTestChannels(frequency int) {
if frequency <= 0 { if frequency <= 0 {
common.SysLog("CHANNEL_TEST_FREQUENCY is not set or invalid, skipping automatic channel test") logger.SysLog("CHANNEL_TEST_FREQUENCY is not set or invalid, skipping automatic channel test")
return return
} }
for { for {
time.Sleep(time.Duration(frequency) * time.Minute) time.Sleep(time.Duration(frequency) * time.Minute)
common.SysLog("testing all channels") logger.SysLog("testing all channels")
_ = testAllChannels(false) _ = testAllChannels(false)
common.SysLog("channel test finished") logger.SysLog("channel test finished")
} }
} }
...@@ -3,101 +3,101 @@ ...@@ -3,101 +3,101 @@
package controller package controller
import ( import (
"encoding/json" "encoding/json"
"net/http" "github.com/gin-gonic/gin"
"one-api/common" "net/http"
"one-api/model" "one-api/logger"
"github.com/gin-gonic/gin" "one-api/model"
) )
// MigrateConsoleSetting 迁移旧的控制台相关配置到 console_setting.* // MigrateConsoleSetting 迁移旧的控制台相关配置到 console_setting.*
func MigrateConsoleSetting(c *gin.Context) { func MigrateConsoleSetting(c *gin.Context) {
// 读取全部 option // 读取全部 option
opts, err := model.AllOption() opts, err := model.AllOption()
if err != nil { if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"success": false, "message": err.Error()}) c.JSON(http.StatusInternalServerError, gin.H{"success": false, "message": err.Error()})
return return
} }
// 建立 map // 建立 map
valMap := map[string]string{} valMap := map[string]string{}
for _, o := range opts { for _, o := range opts {
valMap[o.Key] = o.Value valMap[o.Key] = o.Value
} }
// 处理 APIInfo // 处理 APIInfo
if v := valMap["ApiInfo"]; v != "" { if v := valMap["ApiInfo"]; v != "" {
var arr []map[string]interface{} var arr []map[string]interface{}
if err := json.Unmarshal([]byte(v), &arr); err == nil { if err := json.Unmarshal([]byte(v), &arr); err == nil {
if len(arr) > 50 { if len(arr) > 50 {
arr = arr[:50] arr = arr[:50]
} }
bytes, _ := json.Marshal(arr) bytes, _ := json.Marshal(arr)
model.UpdateOption("console_setting.api_info", string(bytes)) model.UpdateOption("console_setting.api_info", string(bytes))
} }
model.UpdateOption("ApiInfo", "") model.UpdateOption("ApiInfo", "")
} }
// Announcements 直接搬 // Announcements 直接搬
if v := valMap["Announcements"]; v != "" { if v := valMap["Announcements"]; v != "" {
model.UpdateOption("console_setting.announcements", v) model.UpdateOption("console_setting.announcements", v)
model.UpdateOption("Announcements", "") model.UpdateOption("Announcements", "")
} }
// FAQ 转换 // FAQ 转换
if v := valMap["FAQ"]; v != "" { if v := valMap["FAQ"]; v != "" {
var arr []map[string]interface{} var arr []map[string]interface{}
if err := json.Unmarshal([]byte(v), &arr); err == nil { if err := json.Unmarshal([]byte(v), &arr); err == nil {
out := []map[string]interface{}{} out := []map[string]interface{}{}
for _, item := range arr { for _, item := range arr {
q, _ := item["question"].(string) q, _ := item["question"].(string)
if q == "" { if q == "" {
q, _ = item["title"].(string) q, _ = item["title"].(string)
} }
a, _ := item["answer"].(string) a, _ := item["answer"].(string)
if a == "" { if a == "" {
a, _ = item["content"].(string) a, _ = item["content"].(string)
} }
if q != "" && a != "" { if q != "" && a != "" {
out = append(out, map[string]interface{}{"question": q, "answer": a}) out = append(out, map[string]interface{}{"question": q, "answer": a})
} }
} }
if len(out) > 50 { if len(out) > 50 {
out = out[:50] out = out[:50]
} }
bytes, _ := json.Marshal(out) bytes, _ := json.Marshal(out)
model.UpdateOption("console_setting.faq", string(bytes)) model.UpdateOption("console_setting.faq", string(bytes))
} }
model.UpdateOption("FAQ", "") model.UpdateOption("FAQ", "")
} }
// Uptime Kuma 迁移到新的 groups 结构(console_setting.uptime_kuma_groups) // Uptime Kuma 迁移到新的 groups 结构(console_setting.uptime_kuma_groups)
url := valMap["UptimeKumaUrl"] url := valMap["UptimeKumaUrl"]
slug := valMap["UptimeKumaSlug"] slug := valMap["UptimeKumaSlug"]
if url != "" && slug != "" { if url != "" && slug != "" {
// 仅当同时存在 URL 与 Slug 时才进行迁移 // 仅当同时存在 URL 与 Slug 时才进行迁移
groups := []map[string]interface{}{ groups := []map[string]interface{}{
{ {
"id": 1, "id": 1,
"categoryName": "old", "categoryName": "old",
"url": url, "url": url,
"slug": slug, "slug": slug,
"description": "", "description": "",
}, },
} }
bytes, _ := json.Marshal(groups) bytes, _ := json.Marshal(groups)
model.UpdateOption("console_setting.uptime_kuma_groups", string(bytes)) model.UpdateOption("console_setting.uptime_kuma_groups", string(bytes))
} }
// 清空旧键内容 // 清空旧键内容
if url != "" { if url != "" {
model.UpdateOption("UptimeKumaUrl", "") model.UpdateOption("UptimeKumaUrl", "")
} }
if slug != "" { if slug != "" {
model.UpdateOption("UptimeKumaSlug", "") model.UpdateOption("UptimeKumaSlug", "")
} }
// 删除旧键记录 // 删除旧键记录
oldKeys := []string{"ApiInfo", "Announcements", "FAQ", "UptimeKumaUrl", "UptimeKumaSlug"} oldKeys := []string{"ApiInfo", "Announcements", "FAQ", "UptimeKumaUrl", "UptimeKumaSlug"}
model.DB.Where("key IN ?", oldKeys).Delete(&model.Option{}) model.DB.Where("key IN ?", oldKeys).Delete(&model.Option{})
// 重新加载 OptionMap // 重新加载 OptionMap
model.InitOptionMap() model.InitOptionMap()
common.SysLog("console setting migrated") logger.SysLog("console setting migrated")
c.JSON(http.StatusOK, gin.H{"success": true, "message": "migrated"}) c.JSON(http.StatusOK, gin.H{"success": true, "message": "migrated"})
} }
\ No newline at end of file
...@@ -7,6 +7,7 @@ import ( ...@@ -7,6 +7,7 @@ import (
"fmt" "fmt"
"net/http" "net/http"
"one-api/common" "one-api/common"
"one-api/logger"
"one-api/model" "one-api/model"
"strconv" "strconv"
"time" "time"
...@@ -47,7 +48,7 @@ func getGitHubUserInfoByCode(code string) (*GitHubUser, error) { ...@@ -47,7 +48,7 @@ func getGitHubUserInfoByCode(code string) (*GitHubUser, error) {
} }
res, err := client.Do(req) res, err := client.Do(req)
if err != nil { if err != nil {
common.SysLog(err.Error()) logger.SysLog(err.Error())
return nil, errors.New("无法连接至 GitHub 服务器,请稍后重试!") return nil, errors.New("无法连接至 GitHub 服务器,请稍后重试!")
} }
defer res.Body.Close() defer res.Body.Close()
...@@ -63,7 +64,7 @@ func getGitHubUserInfoByCode(code string) (*GitHubUser, error) { ...@@ -63,7 +64,7 @@ func getGitHubUserInfoByCode(code string) (*GitHubUser, error) {
req.Header.Set("Authorization", fmt.Sprintf("Bearer %s", oAuthResponse.AccessToken)) req.Header.Set("Authorization", fmt.Sprintf("Bearer %s", oAuthResponse.AccessToken))
res2, err := client.Do(req) res2, err := client.Do(req)
if err != nil { if err != nil {
common.SysLog(err.Error()) logger.SysLog(err.Error())
return nil, errors.New("无法连接至 GitHub 服务器,请稍后重试!") return nil, errors.New("无法连接至 GitHub 服务器,请稍后重试!")
} }
defer res2.Body.Close() defer res2.Body.Close()
......
...@@ -9,6 +9,7 @@ import ( ...@@ -9,6 +9,7 @@ import (
"net/http" "net/http"
"one-api/common" "one-api/common"
"one-api/dto" "one-api/dto"
"one-api/logger"
"one-api/model" "one-api/model"
"one-api/service" "one-api/service"
"one-api/setting" "one-api/setting"
...@@ -28,7 +29,7 @@ func UpdateMidjourneyTaskBulk() { ...@@ -28,7 +29,7 @@ func UpdateMidjourneyTaskBulk() {
continue continue
} }
common.LogInfo(ctx, fmt.Sprintf("检测到未完成的任务数有: %v", len(tasks))) logger.LogInfo(ctx, fmt.Sprintf("检测到未完成的任务数有: %v", len(tasks)))
taskChannelM := make(map[int][]string) taskChannelM := make(map[int][]string)
taskM := make(map[string]*model.Midjourney) taskM := make(map[string]*model.Midjourney)
nullTaskIds := make([]int, 0) nullTaskIds := make([]int, 0)
...@@ -47,9 +48,9 @@ func UpdateMidjourneyTaskBulk() { ...@@ -47,9 +48,9 @@ func UpdateMidjourneyTaskBulk() {
"progress": "100%", "progress": "100%",
}) })
if err != nil { if err != nil {
common.LogError(ctx, fmt.Sprintf("Fix null mj_id task error: %v", err)) logger.LogError(ctx, fmt.Sprintf("Fix null mj_id task error: %v", err))
} else { } else {
common.LogInfo(ctx, fmt.Sprintf("Fix null mj_id task success: %v", nullTaskIds)) logger.LogInfo(ctx, fmt.Sprintf("Fix null mj_id task success: %v", nullTaskIds))
} }
} }
if len(taskChannelM) == 0 { if len(taskChannelM) == 0 {
...@@ -57,20 +58,20 @@ func UpdateMidjourneyTaskBulk() { ...@@ -57,20 +58,20 @@ func UpdateMidjourneyTaskBulk() {
} }
for channelId, taskIds := range taskChannelM { for channelId, taskIds := range taskChannelM {
common.LogInfo(ctx, fmt.Sprintf("渠道 #%d 未完成的任务有: %d", channelId, len(taskIds))) logger.LogInfo(ctx, fmt.Sprintf("渠道 #%d 未完成的任务有: %d", channelId, len(taskIds)))
if len(taskIds) == 0 { if len(taskIds) == 0 {
continue continue
} }
midjourneyChannel, err := model.CacheGetChannel(channelId) midjourneyChannel, err := model.CacheGetChannel(channelId)
if err != nil { if err != nil {
common.LogError(ctx, fmt.Sprintf("CacheGetChannel: %v", err)) logger.LogError(ctx, fmt.Sprintf("CacheGetChannel: %v", err))
err := model.MjBulkUpdate(taskIds, map[string]any{ err := model.MjBulkUpdate(taskIds, map[string]any{
"fail_reason": fmt.Sprintf("获取渠道信息失败,请联系管理员,渠道ID:%d", channelId), "fail_reason": fmt.Sprintf("获取渠道信息失败,请联系管理员,渠道ID:%d", channelId),
"status": "FAILURE", "status": "FAILURE",
"progress": "100%", "progress": "100%",
}) })
if err != nil { if err != nil {
common.LogInfo(ctx, fmt.Sprintf("UpdateMidjourneyTask error: %v", err)) logger.LogInfo(ctx, fmt.Sprintf("UpdateMidjourneyTask error: %v", err))
} }
continue continue
} }
...@@ -81,7 +82,7 @@ func UpdateMidjourneyTaskBulk() { ...@@ -81,7 +82,7 @@ func UpdateMidjourneyTaskBulk() {
}) })
req, err := http.NewRequest("POST", requestUrl, bytes.NewBuffer(body)) req, err := http.NewRequest("POST", requestUrl, bytes.NewBuffer(body))
if err != nil { if err != nil {
common.LogError(ctx, fmt.Sprintf("Get Task error: %v", err)) logger.LogError(ctx, fmt.Sprintf("Get Task error: %v", err))
continue continue
} }
// 设置超时时间 // 设置超时时间
...@@ -93,22 +94,22 @@ func UpdateMidjourneyTaskBulk() { ...@@ -93,22 +94,22 @@ func UpdateMidjourneyTaskBulk() {
req.Header.Set("mj-api-secret", midjourneyChannel.Key) req.Header.Set("mj-api-secret", midjourneyChannel.Key)
resp, err := service.GetHttpClient().Do(req) resp, err := service.GetHttpClient().Do(req)
if err != nil { if err != nil {
common.LogError(ctx, fmt.Sprintf("Get Task Do req error: %v", err)) logger.LogError(ctx, fmt.Sprintf("Get Task Do req error: %v", err))
continue continue
} }
if resp.StatusCode != http.StatusOK { if resp.StatusCode != http.StatusOK {
common.LogError(ctx, fmt.Sprintf("Get Task status code: %d", resp.StatusCode)) logger.LogError(ctx, fmt.Sprintf("Get Task status code: %d", resp.StatusCode))
continue continue
} }
responseBody, err := io.ReadAll(resp.Body) responseBody, err := io.ReadAll(resp.Body)
if err != nil { if err != nil {
common.LogError(ctx, fmt.Sprintf("Get Task parse body error: %v", err)) logger.LogError(ctx, fmt.Sprintf("Get Task parse body error: %v", err))
continue continue
} }
var responseItems []dto.MidjourneyDto var responseItems []dto.MidjourneyDto
err = json.Unmarshal(responseBody, &responseItems) err = json.Unmarshal(responseBody, &responseItems)
if err != nil { if err != nil {
common.LogError(ctx, fmt.Sprintf("Get Task parse body error2: %v, body: %s", err, string(responseBody))) logger.LogError(ctx, fmt.Sprintf("Get Task parse body error2: %v, body: %s", err, string(responseBody)))
continue continue
} }
resp.Body.Close() resp.Body.Close()
...@@ -147,12 +148,12 @@ func UpdateMidjourneyTaskBulk() { ...@@ -147,12 +148,12 @@ func UpdateMidjourneyTaskBulk() {
} }
// 映射 VideoUrl // 映射 VideoUrl
task.VideoUrl = responseItem.VideoUrl task.VideoUrl = responseItem.VideoUrl
// 映射 VideoUrls - 将数组序列化为 JSON 字符串 // 映射 VideoUrls - 将数组序列化为 JSON 字符串
if responseItem.VideoUrls != nil && len(responseItem.VideoUrls) > 0 { if responseItem.VideoUrls != nil && len(responseItem.VideoUrls) > 0 {
videoUrlsStr, err := json.Marshal(responseItem.VideoUrls) videoUrlsStr, err := json.Marshal(responseItem.VideoUrls)
if err != nil { if err != nil {
common.LogError(ctx, fmt.Sprintf("序列化 VideoUrls 失败: %v", err)) logger.LogError(ctx, fmt.Sprintf("序列化 VideoUrls 失败: %v", err))
task.VideoUrls = "[]" // 失败时设置为空数组 task.VideoUrls = "[]" // 失败时设置为空数组
} else { } else {
task.VideoUrls = string(videoUrlsStr) task.VideoUrls = string(videoUrlsStr)
...@@ -160,10 +161,10 @@ func UpdateMidjourneyTaskBulk() { ...@@ -160,10 +161,10 @@ func UpdateMidjourneyTaskBulk() {
} else { } else {
task.VideoUrls = "" // 空值时清空字段 task.VideoUrls = "" // 空值时清空字段
} }
shouldReturnQuota := false shouldReturnQuota := false
if (task.Progress != "100%" && responseItem.FailReason != "") || (task.Progress == "100%" && task.Status == "FAILURE") { if (task.Progress != "100%" && responseItem.FailReason != "") || (task.Progress == "100%" && task.Status == "FAILURE") {
common.LogInfo(ctx, task.MjId+" 构建失败,"+task.FailReason) logger.LogInfo(ctx, task.MjId+" 构建失败,"+task.FailReason)
task.Progress = "100%" task.Progress = "100%"
if task.Quota != 0 { if task.Quota != 0 {
shouldReturnQuota = true shouldReturnQuota = true
...@@ -171,14 +172,14 @@ func UpdateMidjourneyTaskBulk() { ...@@ -171,14 +172,14 @@ func UpdateMidjourneyTaskBulk() {
} }
err = task.Update() err = task.Update()
if err != nil { if err != nil {
common.LogError(ctx, "UpdateMidjourneyTask task error: "+err.Error()) logger.LogError(ctx, "UpdateMidjourneyTask task error: "+err.Error())
} else { } else {
if shouldReturnQuota { if shouldReturnQuota {
err = model.IncreaseUserQuota(task.UserId, task.Quota, false) err = model.IncreaseUserQuota(task.UserId, task.Quota, false)
if err != nil { if err != nil {
common.LogError(ctx, "fail to increase user quota: "+err.Error()) logger.LogError(ctx, "fail to increase user quota: "+err.Error())
} }
logContent := fmt.Sprintf("构图失败 %s,补偿 %s", task.MjId, common.LogQuota(task.Quota)) logContent := fmt.Sprintf("构图失败 %s,补偿 %s", task.MjId, logger.LogQuota(task.Quota))
model.RecordLog(task.UserId, model.LogTypeSystem, logContent) model.RecordLog(task.UserId, model.LogTypeSystem, logContent)
} }
} }
......
...@@ -7,6 +7,7 @@ import ( ...@@ -7,6 +7,7 @@ import (
"net/http" "net/http"
"net/url" "net/url"
"one-api/common" "one-api/common"
"one-api/logger"
"one-api/model" "one-api/model"
"one-api/setting" "one-api/setting"
"one-api/setting/system_setting" "one-api/setting/system_setting"
...@@ -58,7 +59,7 @@ func getOidcUserInfoByCode(code string) (*OidcUser, error) { ...@@ -58,7 +59,7 @@ func getOidcUserInfoByCode(code string) (*OidcUser, error) {
} }
res, err := client.Do(req) res, err := client.Do(req)
if err != nil { if err != nil {
common.SysLog(err.Error()) logger.SysLog(err.Error())
return nil, errors.New("无法连接至 OIDC 服务器,请稍后重试!") return nil, errors.New("无法连接至 OIDC 服务器,请稍后重试!")
} }
defer res.Body.Close() defer res.Body.Close()
...@@ -69,7 +70,7 @@ func getOidcUserInfoByCode(code string) (*OidcUser, error) { ...@@ -69,7 +70,7 @@ func getOidcUserInfoByCode(code string) (*OidcUser, error) {
} }
if oidcResponse.AccessToken == "" { if oidcResponse.AccessToken == "" {
common.SysError("OIDC 获取 Token 失败,请检查设置!") logger.SysError("OIDC 获取 Token 失败,请检查设置!")
return nil, errors.New("OIDC 获取 Token 失败,请检查设置!") return nil, errors.New("OIDC 获取 Token 失败,请检查设置!")
} }
...@@ -80,12 +81,12 @@ func getOidcUserInfoByCode(code string) (*OidcUser, error) { ...@@ -80,12 +81,12 @@ func getOidcUserInfoByCode(code string) (*OidcUser, error) {
req.Header.Set("Authorization", "Bearer "+oidcResponse.AccessToken) req.Header.Set("Authorization", "Bearer "+oidcResponse.AccessToken)
res2, err := client.Do(req) res2, err := client.Do(req)
if err != nil { if err != nil {
common.SysLog(err.Error()) logger.SysLog(err.Error())
return nil, errors.New("无法连接至 OIDC 服务器,请稍后重试!") return nil, errors.New("无法连接至 OIDC 服务器,请稍后重试!")
} }
defer res2.Body.Close() defer res2.Body.Close()
if res2.StatusCode != http.StatusOK { if res2.StatusCode != http.StatusOK {
common.SysError("OIDC 获取用户信息失败!请检查设置!") logger.SysError("OIDC 获取用户信息失败!请检查设置!")
return nil, errors.New("OIDC 获取用户信息失败!请检查设置!") return nil, errors.New("OIDC 获取用户信息失败!请检查设置!")
} }
...@@ -95,7 +96,7 @@ func getOidcUserInfoByCode(code string) (*OidcUser, error) { ...@@ -95,7 +96,7 @@ func getOidcUserInfoByCode(code string) (*OidcUser, error) {
return nil, err return nil, err
} }
if oidcUser.OpenID == "" || oidcUser.Email == "" { if oidcUser.OpenID == "" || oidcUser.Email == "" {
common.SysError("OIDC 获取用户信息为空!请检查设置!") logger.SysError("OIDC 获取用户信息为空!请检查设置!")
return nil, errors.New("OIDC 获取用户信息为空!请检查设置!") return nil, errors.New("OIDC 获取用户信息为空!请检查设置!")
} }
return &oidcUser, nil return &oidcUser, nil
......
...@@ -10,6 +10,7 @@ import ( ...@@ -10,6 +10,7 @@ import (
"one-api/common" "one-api/common"
"one-api/constant" "one-api/constant"
"one-api/dto" "one-api/dto"
"one-api/logger"
"one-api/model" "one-api/model"
"one-api/relay" "one-api/relay"
"sort" "sort"
...@@ -25,7 +26,7 @@ func UpdateTaskBulk() { ...@@ -25,7 +26,7 @@ func UpdateTaskBulk() {
//imageModel := "midjourney" //imageModel := "midjourney"
for { for {
time.Sleep(time.Duration(15) * time.Second) time.Sleep(time.Duration(15) * time.Second)
common.SysLog("任务进度轮询开始") logger.SysLog("任务进度轮询开始")
ctx := context.TODO() ctx := context.TODO()
allTasks := model.GetAllUnFinishSyncTasks(500) allTasks := model.GetAllUnFinishSyncTasks(500)
platformTask := make(map[constant.TaskPlatform][]*model.Task) platformTask := make(map[constant.TaskPlatform][]*model.Task)
...@@ -54,9 +55,9 @@ func UpdateTaskBulk() { ...@@ -54,9 +55,9 @@ func UpdateTaskBulk() {
"progress": "100%", "progress": "100%",
}) })
if err != nil { if err != nil {
common.LogError(ctx, fmt.Sprintf("Fix null task_id task error: %v", err)) logger.LogError(ctx, fmt.Sprintf("Fix null task_id task error: %v", err))
} else { } else {
common.LogInfo(ctx, fmt.Sprintf("Fix null task_id task success: %v", nullTaskIds)) logger.LogInfo(ctx, fmt.Sprintf("Fix null task_id task success: %v", nullTaskIds))
} }
} }
if len(taskChannelM) == 0 { if len(taskChannelM) == 0 {
...@@ -65,7 +66,7 @@ func UpdateTaskBulk() { ...@@ -65,7 +66,7 @@ func UpdateTaskBulk() {
UpdateTaskByPlatform(platform, taskChannelM, taskM) UpdateTaskByPlatform(platform, taskChannelM, taskM)
} }
common.SysLog("任务进度轮询完成") logger.SysLog("任务进度轮询完成")
} }
} }
...@@ -77,7 +78,7 @@ func UpdateTaskByPlatform(platform constant.TaskPlatform, taskChannelM map[int][ ...@@ -77,7 +78,7 @@ func UpdateTaskByPlatform(platform constant.TaskPlatform, taskChannelM map[int][
_ = UpdateSunoTaskAll(context.Background(), taskChannelM, taskM) _ = UpdateSunoTaskAll(context.Background(), taskChannelM, taskM)
default: default:
if err := UpdateVideoTaskAll(context.Background(), platform, taskChannelM, taskM); err != nil { if err := UpdateVideoTaskAll(context.Background(), platform, taskChannelM, taskM); err != nil {
common.SysLog(fmt.Sprintf("UpdateVideoTaskAll fail: %s", err)) logger.SysLog(fmt.Sprintf("UpdateVideoTaskAll fail: %s", err))
} }
} }
} }
...@@ -86,27 +87,27 @@ func UpdateSunoTaskAll(ctx context.Context, taskChannelM map[int][]string, taskM ...@@ -86,27 +87,27 @@ func UpdateSunoTaskAll(ctx context.Context, taskChannelM map[int][]string, taskM
for channelId, taskIds := range taskChannelM { for channelId, taskIds := range taskChannelM {
err := updateSunoTaskAll(ctx, channelId, taskIds, taskM) err := updateSunoTaskAll(ctx, channelId, taskIds, taskM)
if err != nil { if err != nil {
common.LogError(ctx, fmt.Sprintf("渠道 #%d 更新异步任务失败: %d", channelId, err.Error())) logger.LogError(ctx, fmt.Sprintf("渠道 #%d 更新异步任务失败: %d", channelId, err.Error()))
} }
} }
return nil return nil
} }
func updateSunoTaskAll(ctx context.Context, channelId int, taskIds []string, taskM map[string]*model.Task) error { func updateSunoTaskAll(ctx context.Context, channelId int, taskIds []string, taskM map[string]*model.Task) error {
common.LogInfo(ctx, fmt.Sprintf("渠道 #%d 未完成的任务有: %d", channelId, len(taskIds))) logger.LogInfo(ctx, fmt.Sprintf("渠道 #%d 未完成的任务有: %d", channelId, len(taskIds)))
if len(taskIds) == 0 { if len(taskIds) == 0 {
return nil return nil
} }
channel, err := model.CacheGetChannel(channelId) channel, err := model.CacheGetChannel(channelId)
if err != nil { if err != nil {
common.SysLog(fmt.Sprintf("CacheGetChannel: %v", err)) logger.SysLog(fmt.Sprintf("CacheGetChannel: %v", err))
err = model.TaskBulkUpdate(taskIds, map[string]any{ err = model.TaskBulkUpdate(taskIds, map[string]any{
"fail_reason": fmt.Sprintf("获取渠道信息失败,请联系管理员,渠道ID:%d", channelId), "fail_reason": fmt.Sprintf("获取渠道信息失败,请联系管理员,渠道ID:%d", channelId),
"status": "FAILURE", "status": "FAILURE",
"progress": "100%", "progress": "100%",
}) })
if err != nil { if err != nil {
common.SysError(fmt.Sprintf("UpdateMidjourneyTask error2: %v", err)) logger.SysError(fmt.Sprintf("UpdateMidjourneyTask error2: %v", err))
} }
return err return err
} }
...@@ -118,27 +119,27 @@ func updateSunoTaskAll(ctx context.Context, channelId int, taskIds []string, tas ...@@ -118,27 +119,27 @@ func updateSunoTaskAll(ctx context.Context, channelId int, taskIds []string, tas
"ids": taskIds, "ids": taskIds,
}) })
if err != nil { if err != nil {
common.SysError(fmt.Sprintf("Get Task Do req error: %v", err)) logger.SysError(fmt.Sprintf("Get Task Do req error: %v", err))
return err return err
} }
if resp.StatusCode != http.StatusOK { if resp.StatusCode != http.StatusOK {
common.LogError(ctx, fmt.Sprintf("Get Task status code: %d", resp.StatusCode)) logger.LogError(ctx, fmt.Sprintf("Get Task status code: %d", resp.StatusCode))
return errors.New(fmt.Sprintf("Get Task status code: %d", resp.StatusCode)) return errors.New(fmt.Sprintf("Get Task status code: %d", resp.StatusCode))
} }
defer resp.Body.Close() defer resp.Body.Close()
responseBody, err := io.ReadAll(resp.Body) responseBody, err := io.ReadAll(resp.Body)
if err != nil { if err != nil {
common.SysError(fmt.Sprintf("Get Task parse body error: %v", err)) logger.SysError(fmt.Sprintf("Get Task parse body error: %v", err))
return err return err
} }
var responseItems dto.TaskResponse[[]dto.SunoDataResponse] var responseItems dto.TaskResponse[[]dto.SunoDataResponse]
err = json.Unmarshal(responseBody, &responseItems) err = json.Unmarshal(responseBody, &responseItems)
if err != nil { if err != nil {
common.LogError(ctx, fmt.Sprintf("Get Task parse body error2: %v, body: %s", err, string(responseBody))) logger.LogError(ctx, fmt.Sprintf("Get Task parse body error2: %v, body: %s", err, string(responseBody)))
return err return err
} }
if !responseItems.IsSuccess() { if !responseItems.IsSuccess() {
common.SysLog(fmt.Sprintf("渠道 #%d 未完成的任务有: %d, 成功获取到任务数: %d", channelId, len(taskIds), string(responseBody))) logger.SysLog(fmt.Sprintf("渠道 #%d 未完成的任务有: %d, 成功获取到任务数: %d", channelId, len(taskIds), string(responseBody)))
return err return err
} }
...@@ -154,19 +155,19 @@ func updateSunoTaskAll(ctx context.Context, channelId int, taskIds []string, tas ...@@ -154,19 +155,19 @@ func updateSunoTaskAll(ctx context.Context, channelId int, taskIds []string, tas
task.StartTime = lo.If(responseItem.StartTime != 0, responseItem.StartTime).Else(task.StartTime) task.StartTime = lo.If(responseItem.StartTime != 0, responseItem.StartTime).Else(task.StartTime)
task.FinishTime = lo.If(responseItem.FinishTime != 0, responseItem.FinishTime).Else(task.FinishTime) task.FinishTime = lo.If(responseItem.FinishTime != 0, responseItem.FinishTime).Else(task.FinishTime)
if responseItem.FailReason != "" || task.Status == model.TaskStatusFailure { if responseItem.FailReason != "" || task.Status == model.TaskStatusFailure {
common.LogInfo(ctx, task.TaskID+" 构建失败,"+task.FailReason) logger.LogInfo(ctx, task.TaskID+" 构建失败,"+task.FailReason)
task.Progress = "100%" task.Progress = "100%"
//err = model.CacheUpdateUserQuota(task.UserId) ? //err = model.CacheUpdateUserQuota(task.UserId) ?
if err != nil { if err != nil {
common.LogError(ctx, "error update user quota cache: "+err.Error()) logger.LogError(ctx, "error update user quota cache: "+err.Error())
} else { } else {
quota := task.Quota quota := task.Quota
if quota != 0 { if quota != 0 {
err = model.IncreaseUserQuota(task.UserId, quota, false) err = model.IncreaseUserQuota(task.UserId, quota, false)
if err != nil { if err != nil {
common.LogError(ctx, "fail to increase user quota: "+err.Error()) logger.LogError(ctx, "fail to increase user quota: "+err.Error())
} }
logContent := fmt.Sprintf("异步任务执行失败 %s,补偿 %s", task.TaskID, common.LogQuota(quota)) logContent := fmt.Sprintf("异步任务执行失败 %s,补偿 %s", task.TaskID, logger.LogQuota(quota))
model.RecordLog(task.UserId, model.LogTypeSystem, logContent) model.RecordLog(task.UserId, model.LogTypeSystem, logContent)
} }
} }
...@@ -178,7 +179,7 @@ func updateSunoTaskAll(ctx context.Context, channelId int, taskIds []string, tas ...@@ -178,7 +179,7 @@ func updateSunoTaskAll(ctx context.Context, channelId int, taskIds []string, tas
err = task.Update() err = task.Update()
if err != nil { if err != nil {
common.SysError("UpdateMidjourneyTask task error: " + err.Error()) logger.SysError("UpdateMidjourneyTask task error: " + err.Error())
} }
} }
return nil return nil
......
...@@ -5,9 +5,9 @@ import ( ...@@ -5,9 +5,9 @@ import (
"encoding/json" "encoding/json"
"fmt" "fmt"
"io" "io"
"one-api/common"
"one-api/constant" "one-api/constant"
"one-api/dto" "one-api/dto"
"one-api/logger"
"one-api/model" "one-api/model"
"one-api/relay" "one-api/relay"
"one-api/relay/channel" "one-api/relay/channel"
...@@ -18,14 +18,14 @@ import ( ...@@ -18,14 +18,14 @@ import (
func UpdateVideoTaskAll(ctx context.Context, platform constant.TaskPlatform, taskChannelM map[int][]string, taskM map[string]*model.Task) error { func UpdateVideoTaskAll(ctx context.Context, platform constant.TaskPlatform, taskChannelM map[int][]string, taskM map[string]*model.Task) error {
for channelId, taskIds := range taskChannelM { for channelId, taskIds := range taskChannelM {
if err := updateVideoTaskAll(ctx, platform, channelId, taskIds, taskM); err != nil { if err := updateVideoTaskAll(ctx, platform, channelId, taskIds, taskM); err != nil {
common.LogError(ctx, fmt.Sprintf("Channel #%d failed to update video async tasks: %s", channelId, err.Error())) logger.LogError(ctx, fmt.Sprintf("Channel #%d failed to update video async tasks: %s", channelId, err.Error()))
} }
} }
return nil return nil
} }
func updateVideoTaskAll(ctx context.Context, platform constant.TaskPlatform, channelId int, taskIds []string, taskM map[string]*model.Task) error { func updateVideoTaskAll(ctx context.Context, platform constant.TaskPlatform, channelId int, taskIds []string, taskM map[string]*model.Task) error {
common.LogInfo(ctx, fmt.Sprintf("Channel #%d pending video tasks: %d", channelId, len(taskIds))) logger.LogInfo(ctx, fmt.Sprintf("Channel #%d pending video tasks: %d", channelId, len(taskIds)))
if len(taskIds) == 0 { if len(taskIds) == 0 {
return nil return nil
} }
...@@ -37,7 +37,7 @@ func updateVideoTaskAll(ctx context.Context, platform constant.TaskPlatform, cha ...@@ -37,7 +37,7 @@ func updateVideoTaskAll(ctx context.Context, platform constant.TaskPlatform, cha
"progress": "100%", "progress": "100%",
}) })
if errUpdate != nil { if errUpdate != nil {
common.SysError(fmt.Sprintf("UpdateVideoTask error: %v", errUpdate)) logger.SysError(fmt.Sprintf("UpdateVideoTask error: %v", errUpdate))
} }
return fmt.Errorf("CacheGetChannel failed: %w", err) return fmt.Errorf("CacheGetChannel failed: %w", err)
} }
...@@ -47,7 +47,7 @@ func updateVideoTaskAll(ctx context.Context, platform constant.TaskPlatform, cha ...@@ -47,7 +47,7 @@ func updateVideoTaskAll(ctx context.Context, platform constant.TaskPlatform, cha
} }
for _, taskId := range taskIds { for _, taskId := range taskIds {
if err := updateVideoSingleTask(ctx, adaptor, cacheGetChannel, taskId, taskM); err != nil { if err := updateVideoSingleTask(ctx, adaptor, cacheGetChannel, taskId, taskM); err != nil {
common.LogError(ctx, fmt.Sprintf("Failed to update video task %s: %s", taskId, err.Error())) logger.LogError(ctx, fmt.Sprintf("Failed to update video task %s: %s", taskId, err.Error()))
} }
} }
return nil return nil
...@@ -61,7 +61,7 @@ func updateVideoSingleTask(ctx context.Context, adaptor channel.TaskAdaptor, cha ...@@ -61,7 +61,7 @@ func updateVideoSingleTask(ctx context.Context, adaptor channel.TaskAdaptor, cha
task := taskM[taskId] task := taskM[taskId]
if task == nil { if task == nil {
common.LogError(ctx, fmt.Sprintf("Task %s not found in taskM", taskId)) logger.LogError(ctx, fmt.Sprintf("Task %s not found in taskM", taskId))
return fmt.Errorf("task %s not found", taskId) return fmt.Errorf("task %s not found", taskId)
} }
resp, err := adaptor.FetchTask(baseURL, channel.Key, map[string]any{ resp, err := adaptor.FetchTask(baseURL, channel.Key, map[string]any{
...@@ -124,13 +124,13 @@ func updateVideoSingleTask(ctx context.Context, adaptor channel.TaskAdaptor, cha ...@@ -124,13 +124,13 @@ func updateVideoSingleTask(ctx context.Context, adaptor channel.TaskAdaptor, cha
task.FinishTime = now task.FinishTime = now
} }
task.FailReason = taskResult.Reason task.FailReason = taskResult.Reason
common.LogInfo(ctx, fmt.Sprintf("Task %s failed: %s", task.TaskID, task.FailReason)) logger.LogInfo(ctx, fmt.Sprintf("Task %s failed: %s", task.TaskID, task.FailReason))
quota := task.Quota quota := task.Quota
if quota != 0 { if quota != 0 {
if err := model.IncreaseUserQuota(task.UserId, quota, false); err != nil { if err := model.IncreaseUserQuota(task.UserId, quota, false); err != nil {
common.LogError(ctx, "Failed to increase user quota: "+err.Error()) logger.LogError(ctx, "Failed to increase user quota: "+err.Error())
} }
logContent := fmt.Sprintf("Video async task failed %s, refund %s", task.TaskID, common.LogQuota(quota)) logContent := fmt.Sprintf("Video async task failed %s, refund %s", task.TaskID, logger.LogQuota(quota))
model.RecordLog(task.UserId, model.LogTypeSystem, logContent) model.RecordLog(task.UserId, model.LogTypeSystem, logContent)
} }
default: default:
...@@ -140,7 +140,7 @@ func updateVideoSingleTask(ctx context.Context, adaptor channel.TaskAdaptor, cha ...@@ -140,7 +140,7 @@ func updateVideoSingleTask(ctx context.Context, adaptor channel.TaskAdaptor, cha
task.Progress = taskResult.Progress task.Progress = taskResult.Progress
} }
if err := task.Update(); err != nil { if err := task.Update(); err != nil {
common.SysError("UpdateVideoTask task error: " + err.Error()) logger.SysError("UpdateVideoTask task error: " + err.Error())
} }
return nil return nil
......
...@@ -3,6 +3,7 @@ package controller ...@@ -3,6 +3,7 @@ package controller
import ( import (
"net/http" "net/http"
"one-api/common" "one-api/common"
"one-api/logger"
"one-api/model" "one-api/model"
"strconv" "strconv"
...@@ -102,7 +103,7 @@ func AddToken(c *gin.Context) { ...@@ -102,7 +103,7 @@ func AddToken(c *gin.Context) {
"success": false, "success": false,
"message": "生成令牌失败", "message": "生成令牌失败",
}) })
common.SysError("failed to generate token key: " + err.Error()) logger.SysError("failed to generate token key: " + err.Error())
return return
} }
cleanToken := model.Token{ cleanToken := model.Token{
......
...@@ -5,6 +5,7 @@ import ( ...@@ -5,6 +5,7 @@ import (
"log" "log"
"net/url" "net/url"
"one-api/common" "one-api/common"
"one-api/logger"
"one-api/model" "one-api/model"
"one-api/service" "one-api/service"
"one-api/setting" "one-api/setting"
...@@ -231,7 +232,7 @@ func EpayNotify(c *gin.Context) { ...@@ -231,7 +232,7 @@ func EpayNotify(c *gin.Context) {
return return
} }
log.Printf("易支付回调更新用户成功 %v", topUp) log.Printf("易支付回调更新用户成功 %v", topUp)
model.RecordLog(topUp.UserId, model.LogTypeTopup, fmt.Sprintf("使用在线充值成功,充值金额: %v,支付金额:%f", common.LogQuota(quotaToAdd), topUp.Money)) model.RecordLog(topUp.UserId, model.LogTypeTopup, fmt.Sprintf("使用在线充值成功,充值金额: %v,支付金额:%f", logger.LogQuota(quotaToAdd), topUp.Money))
} }
} else { } else {
log.Printf("易支付异常回调: %v", verifyInfo) log.Printf("易支付异常回调: %v", verifyInfo)
......
...@@ -5,6 +5,7 @@ import ( ...@@ -5,6 +5,7 @@ import (
"fmt" "fmt"
"net/http" "net/http"
"one-api/common" "one-api/common"
"one-api/logger"
"one-api/model" "one-api/model"
"strconv" "strconv"
...@@ -70,7 +71,7 @@ func Setup2FA(c *gin.Context) { ...@@ -70,7 +71,7 @@ func Setup2FA(c *gin.Context) {
"success": false, "success": false,
"message": "生成2FA密钥失败", "message": "生成2FA密钥失败",
}) })
common.SysError("生成TOTP密钥失败: " + err.Error()) logger.SysError("生成TOTP密钥失败: " + err.Error())
return return
} }
...@@ -81,7 +82,7 @@ func Setup2FA(c *gin.Context) { ...@@ -81,7 +82,7 @@ func Setup2FA(c *gin.Context) {
"success": false, "success": false,
"message": "生成备用码失败", "message": "生成备用码失败",
}) })
common.SysError("生成备用码失败: " + err.Error()) logger.SysError("生成备用码失败: " + err.Error())
return return
} }
...@@ -115,7 +116,7 @@ func Setup2FA(c *gin.Context) { ...@@ -115,7 +116,7 @@ func Setup2FA(c *gin.Context) {
"success": false, "success": false,
"message": "保存备用码失败", "message": "保存备用码失败",
}) })
common.SysError("保存备用码失败: " + err.Error()) logger.SysError("保存备用码失败: " + err.Error())
return return
} }
...@@ -294,7 +295,7 @@ func Get2FAStatus(c *gin.Context) { ...@@ -294,7 +295,7 @@ func Get2FAStatus(c *gin.Context) {
// 获取剩余备用码数量 // 获取剩余备用码数量
backupCount, err := model.GetUnusedBackupCodeCount(userId) backupCount, err := model.GetUnusedBackupCodeCount(userId)
if err != nil { if err != nil {
common.SysError("获取备用码数量失败: " + err.Error()) logger.SysError("获取备用码数量失败: " + err.Error())
} else { } else {
status["backup_codes_remaining"] = backupCount status["backup_codes_remaining"] = backupCount
} }
...@@ -368,7 +369,7 @@ func RegenerateBackupCodes(c *gin.Context) { ...@@ -368,7 +369,7 @@ func RegenerateBackupCodes(c *gin.Context) {
"success": false, "success": false,
"message": "生成备用码失败", "message": "生成备用码失败",
}) })
common.SysError("生成备用码失败: " + err.Error()) logger.SysError("生成备用码失败: " + err.Error())
return return
} }
...@@ -378,7 +379,7 @@ func RegenerateBackupCodes(c *gin.Context) { ...@@ -378,7 +379,7 @@ func RegenerateBackupCodes(c *gin.Context) {
"success": false, "success": false,
"message": "保存备用码失败", "message": "保存备用码失败",
}) })
common.SysError("保存备用码失败: " + err.Error()) logger.SysError("保存备用码失败: " + err.Error())
return return
} }
......
...@@ -7,6 +7,7 @@ import ( ...@@ -7,6 +7,7 @@ import (
"net/url" "net/url"
"one-api/common" "one-api/common"
"one-api/dto" "one-api/dto"
"one-api/logger"
"one-api/model" "one-api/model"
"one-api/setting" "one-api/setting"
"strconv" "strconv"
...@@ -192,7 +193,7 @@ func Register(c *gin.Context) { ...@@ -192,7 +193,7 @@ func Register(c *gin.Context) {
"success": false, "success": false,
"message": "数据库错误,请稍后重试", "message": "数据库错误,请稍后重试",
}) })
common.SysError(fmt.Sprintf("CheckUserExistOrDeleted error: %v", err)) logger.SysError(fmt.Sprintf("CheckUserExistOrDeleted error: %v", err))
return return
} }
if exist { if exist {
...@@ -235,7 +236,7 @@ func Register(c *gin.Context) { ...@@ -235,7 +236,7 @@ func Register(c *gin.Context) {
"success": false, "success": false,
"message": "生成默认令牌失败", "message": "生成默认令牌失败",
}) })
common.SysError("failed to generate token key: " + err.Error()) logger.SysError("failed to generate token key: " + err.Error())
return return
} }
// 生成默认令牌 // 生成默认令牌
...@@ -342,7 +343,7 @@ func GenerateAccessToken(c *gin.Context) { ...@@ -342,7 +343,7 @@ func GenerateAccessToken(c *gin.Context) {
"success": false, "success": false,
"message": "生成失败", "message": "生成失败",
}) })
common.SysError("failed to generate key: " + err.Error()) logger.SysError("failed to generate key: " + err.Error())
return return
} }
user.SetAccessToken(key) user.SetAccessToken(key)
...@@ -517,7 +518,7 @@ func UpdateUser(c *gin.Context) { ...@@ -517,7 +518,7 @@ func UpdateUser(c *gin.Context) {
return return
} }
if originUser.Quota != updatedUser.Quota { if originUser.Quota != updatedUser.Quota {
model.RecordLog(originUser.Id, model.LogTypeManage, fmt.Sprintf("管理员将用户额度从 %s修改为 %s", common.LogQuota(originUser.Quota), common.LogQuota(updatedUser.Quota))) model.RecordLog(originUser.Id, model.LogTypeManage, fmt.Sprintf("管理员将用户额度从 %s修改为 %s", logger.LogQuota(originUser.Quota), logger.LogQuota(updatedUser.Quota)))
} }
c.JSON(http.StatusOK, gin.H{ c.JSON(http.StatusOK, gin.H{
"success": true, "success": true,
......
package dto package dto
import (
"one-api/types"
"github.com/gin-gonic/gin"
)
type AudioRequest struct { type AudioRequest struct {
Model string `json:"model"` Model string `json:"model"`
Input string `json:"input"` Input string `json:"input"`
...@@ -8,6 +14,18 @@ type AudioRequest struct { ...@@ -8,6 +14,18 @@ type AudioRequest struct {
ResponseFormat string `json:"response_format,omitempty"` ResponseFormat string `json:"response_format,omitempty"`
} }
func (r *AudioRequest) GetTokenCountMeta() *types.TokenCountMeta {
meta := &types.TokenCountMeta{
CombineText: r.Input,
TokenType: types.TokenTypeTextNumber,
}
return meta
}
func (r *AudioRequest) IsStream(c *gin.Context) bool {
return false
}
type AudioResponse struct { type AudioResponse struct {
Text string `json:"text"` Text string `json:"text"`
} }
......
...@@ -5,6 +5,9 @@ import ( ...@@ -5,6 +5,9 @@ import (
"fmt" "fmt"
"one-api/common" "one-api/common"
"one-api/types" "one-api/types"
"strings"
"github.com/gin-gonic/gin"
) )
type ClaudeMetadata struct { type ClaudeMetadata struct {
...@@ -81,7 +84,7 @@ func (c *ClaudeMediaMessage) GetStringContent() string { ...@@ -81,7 +84,7 @@ func (c *ClaudeMediaMessage) GetStringContent() string {
} }
func (c *ClaudeMediaMessage) GetJsonRowString() string { func (c *ClaudeMediaMessage) GetJsonRowString() string {
jsonContent, _ := json.Marshal(c) jsonContent, _ := common.Marshal(c)
return string(jsonContent) return string(jsonContent)
} }
...@@ -199,6 +202,129 @@ type ClaudeRequest struct { ...@@ -199,6 +202,129 @@ type ClaudeRequest struct {
Thinking *Thinking `json:"thinking,omitempty"` Thinking *Thinking `json:"thinking,omitempty"`
} }
func (c *ClaudeRequest) GetTokenCountMeta() *types.TokenCountMeta {
var tokenCountMeta = types.TokenCountMeta{
TokenType: types.TokenTypeTextNumber,
MaxTokens: int(c.MaxTokens),
}
var texts = make([]string, 0)
var fileMeta = make([]*types.FileMeta, 0)
// system
if c.System != nil {
if c.IsStringSystem() {
sys := c.GetStringSystem()
if sys != "" {
texts = append(texts, sys)
}
} else {
systemMedia := c.ParseSystem()
for _, media := range systemMedia {
switch media.Type {
case "text":
texts = append(texts, media.GetText())
case "image":
if media.Source != nil {
data := media.Source.Url
if data == "" {
data = common.Interface2String(media.Source.Data)
}
if data != "" {
fileMeta = append(fileMeta, &types.FileMeta{FileType: types.FileTypeImage, Data: data})
}
}
}
}
}
}
// messages
for _, message := range c.Messages {
tokenCountMeta.MessagesCount++
texts = append(texts, message.Role)
if message.IsStringContent() {
content := message.GetStringContent()
if content != "" {
texts = append(texts, content)
}
continue
}
content, _ := message.ParseContent()
for _, media := range content {
switch media.Type {
case "text":
texts = append(texts, media.GetText())
case "image":
if media.Source != nil {
data := media.Source.Url
if data == "" {
data = common.Interface2String(media.Source.Data)
}
if data != "" {
fileMeta = append(fileMeta, &types.FileMeta{FileType: types.FileTypeImage, Data: data})
}
}
case "tool_use":
if media.Name != "" {
texts = append(texts, media.Name)
}
if media.Input != nil {
b, _ := common.Marshal(media.Input)
texts = append(texts, string(b))
}
case "tool_result":
if media.Content != nil {
b, _ := common.Marshal(media.Content)
texts = append(texts, string(b))
}
}
}
}
// tools
if c.Tools != nil {
tools := c.GetTools()
normalTools, webSearchTools := ProcessTools(tools)
if normalTools != nil {
for _, t := range normalTools {
tokenCountMeta.ToolsCount++
if t.Name != "" {
texts = append(texts, t.Name)
}
if t.Description != "" {
texts = append(texts, t.Description)
}
if t.InputSchema != nil {
b, _ := common.Marshal(t.InputSchema)
texts = append(texts, string(b))
}
}
}
if webSearchTools != nil {
for _, t := range webSearchTools {
tokenCountMeta.ToolsCount++
if t.Name != "" {
texts = append(texts, t.Name)
}
if t.UserLocation != nil {
b, _ := common.Marshal(t.UserLocation)
texts = append(texts, string(b))
}
}
}
}
tokenCountMeta.CombineText = strings.Join(texts, "\n")
tokenCountMeta.Files = fileMeta
return &tokenCountMeta
}
func (claudeRequest *ClaudeRequest) IsStream(c *gin.Context) bool {
return claudeRequest.Stream
}
func (c *ClaudeRequest) SearchToolNameByToolCallId(toolCallId string) string { func (c *ClaudeRequest) SearchToolNameByToolCallId(toolCallId string) string {
for _, message := range c.Messages { for _, message := range c.Messages {
content, _ := message.ParseContent() content, _ := message.ParseContent()
......
package dto package dto
import (
"one-api/types"
"strings"
"github.com/gin-gonic/gin"
)
type EmbeddingOptions struct { type EmbeddingOptions struct {
Seed int `json:"seed,omitempty"` Seed int `json:"seed,omitempty"`
Temperature *float64 `json:"temperature,omitempty"` Temperature *float64 `json:"temperature,omitempty"`
...@@ -24,9 +31,26 @@ type EmbeddingRequest struct { ...@@ -24,9 +31,26 @@ type EmbeddingRequest struct {
PresencePenalty float64 `json:"presence_penalty,omitempty"` PresencePenalty float64 `json:"presence_penalty,omitempty"`
} }
func (r EmbeddingRequest) ParseInput() []string { func (r *EmbeddingRequest) GetTokenCountMeta() *types.TokenCountMeta {
var texts = make([]string, 0)
inputs := r.ParseInput()
for _, input := range inputs {
texts = append(texts, input)
}
return &types.TokenCountMeta{
CombineText: strings.Join(texts, "\n"),
}
}
func (r *EmbeddingRequest) IsStream(c *gin.Context) bool {
return false
}
func (r *EmbeddingRequest) ParseInput() []string {
if r.Input == nil { if r.Input == nil {
return nil return make([]string, 0)
} }
var input []string var input []string
switch r.Input.(type) { switch r.Input.(type) {
......
...@@ -2,7 +2,10 @@ package dto ...@@ -2,7 +2,10 @@ package dto
import ( import (
"encoding/json" "encoding/json"
"github.com/gin-gonic/gin"
"one-api/common" "one-api/common"
"one-api/logger"
"one-api/types"
"strings" "strings"
) )
...@@ -14,19 +17,75 @@ type GeminiChatRequest struct { ...@@ -14,19 +17,75 @@ type GeminiChatRequest struct {
SystemInstructions *GeminiChatContent `json:"systemInstruction,omitempty"` SystemInstructions *GeminiChatContent `json:"systemInstruction,omitempty"`
} }
func (r *GeminiChatRequest) GetTokenCountMeta() *types.TokenCountMeta {
var files []*types.FileMeta = make([]*types.FileMeta, 0)
var maxTokens int
if r.GenerationConfig.MaxOutputTokens > 0 {
maxTokens = int(r.GenerationConfig.MaxOutputTokens)
}
var inputTexts []string
for _, content := range r.Contents {
for _, part := range content.Parts {
if part.Text != "" {
inputTexts = append(inputTexts, part.Text)
}
if part.InlineData != nil && part.InlineData.Data != "" {
if strings.HasPrefix(part.InlineData.MimeType, "image/") {
files = append(files, &types.FileMeta{
FileType: types.FileTypeImage,
Data: part.InlineData.Data,
})
} else if strings.HasPrefix(part.InlineData.MimeType, "audio/") {
files = append(files, &types.FileMeta{
FileType: types.FileTypeAudio,
Data: part.InlineData.Data,
})
} else if strings.HasPrefix(part.InlineData.MimeType, "video/") {
files = append(files, &types.FileMeta{
FileType: types.FileTypeVideo,
Data: part.InlineData.Data,
})
} else {
files = append(files, &types.FileMeta{
FileType: types.FileTypeFile,
Data: part.InlineData.Data,
})
}
}
}
}
inputText := strings.Join(inputTexts, "\n")
return &types.TokenCountMeta{
CombineText: inputText,
Files: files,
MaxTokens: maxTokens,
}
}
func (r *GeminiChatRequest) IsStream(c *gin.Context) bool {
if c.Query("alt") == "sse" {
return true
}
return false
}
func (r *GeminiChatRequest) GetTools() []GeminiChatTool { func (r *GeminiChatRequest) GetTools() []GeminiChatTool {
var tools []GeminiChatTool var tools []GeminiChatTool
if strings.HasSuffix(string(r.Tools), "[") { if strings.HasSuffix(string(r.Tools), "[") {
// is array // is array
if err := common.Unmarshal(r.Tools, &tools); err != nil { if err := common.Unmarshal(r.Tools, &tools); err != nil {
common.LogError(nil, "error_unmarshalling_tools: "+err.Error()) logger.LogError(nil, "error_unmarshalling_tools: "+err.Error())
return nil return nil
} }
} else if strings.HasPrefix(string(r.Tools), "{") { } else if strings.HasPrefix(string(r.Tools), "{") {
// is object // is object
singleTool := GeminiChatTool{} singleTool := GeminiChatTool{}
if err := common.Unmarshal(r.Tools, &singleTool); err != nil { if err := common.Unmarshal(r.Tools, &singleTool); err != nil {
common.LogError(nil, "error_unmarshalling_single_tool: "+err.Error()) logger.LogError(nil, "error_unmarshalling_single_tool: "+err.Error())
return nil return nil
} }
tools = []GeminiChatTool{singleTool} tools = []GeminiChatTool{singleTool}
...@@ -43,7 +102,7 @@ func (r *GeminiChatRequest) SetTools(tools []GeminiChatTool) { ...@@ -43,7 +102,7 @@ func (r *GeminiChatRequest) SetTools(tools []GeminiChatTool) {
// Marshal the tools to JSON // Marshal the tools to JSON
data, err := common.Marshal(tools) data, err := common.Marshal(tools)
if err != nil { if err != nil {
common.LogError(nil, "error_marshalling_tools: "+err.Error()) logger.LogError(nil, "error_marshalling_tools: "+err.Error())
return return
} }
r.Tools = data r.Tools = data
......
package dto package dto
import "encoding/json" import (
"encoding/json"
"one-api/types"
"strings"
"github.com/gin-gonic/gin"
)
type ImageRequest struct { type ImageRequest struct {
Model string `json:"model"` Model string `json:"model"`
Prompt string `json:"prompt" binding:"required"` Prompt string `json:"prompt" binding:"required"`
N int `json:"n,omitempty"` N uint `json:"n,omitempty"`
Size string `json:"size,omitempty"` Size string `json:"size,omitempty"`
Quality string `json:"quality,omitempty"` Quality string `json:"quality,omitempty"`
ResponseFormat string `json:"response_format,omitempty"` ResponseFormat string `json:"response_format,omitempty"`
...@@ -18,6 +24,42 @@ type ImageRequest struct { ...@@ -18,6 +24,42 @@ type ImageRequest struct {
Watermark *bool `json:"watermark,omitempty"` Watermark *bool `json:"watermark,omitempty"`
} }
func (i *ImageRequest) GetTokenCountMeta() *types.TokenCountMeta {
var sizeRatio = 1.0
var qualityRatio = 1.0
if strings.HasPrefix(i.Model, "dall-e") {
// Size
if i.Size == "256x256" {
sizeRatio = 0.4
} else if i.Size == "512x512" {
sizeRatio = 0.45
} else if i.Size == "1024x1024" {
sizeRatio = 1
} else if i.Size == "1024x1792" || i.Size == "1792x1024" {
sizeRatio = 2
}
if i.Model == "dall-e-3" && i.Quality == "hd" {
qualityRatio = 2.0
if i.Size == "1024x1792" || i.Size == "1792x1024" {
qualityRatio = 1.5
}
}
}
// not support token count for dalle
return &types.TokenCountMeta{
CombineText: i.Prompt,
MaxTokens: 1584,
ImagePriceRatio: sizeRatio * qualityRatio * float64(i.N),
}
}
func (i *ImageRequest) IsStream(c *gin.Context) bool {
return false
}
type ImageResponse struct { type ImageResponse struct {
Data []ImageData `json:"data"` Data []ImageData `json:"data"`
Created int64 `json:"created"` Created int64 `json:"created"`
......
...@@ -2,8 +2,12 @@ package dto ...@@ -2,8 +2,12 @@ package dto
import ( import (
"encoding/json" "encoding/json"
"fmt"
"one-api/common" "one-api/common"
"one-api/types"
"strings" "strings"
"github.com/gin-gonic/gin"
) )
type ResponseFormat struct { type ResponseFormat struct {
...@@ -67,6 +71,116 @@ type GeneralOpenAIRequest struct { ...@@ -67,6 +71,116 @@ type GeneralOpenAIRequest struct {
Extra map[string]json.RawMessage `json:"-"` Extra map[string]json.RawMessage `json:"-"`
} }
func (r *GeneralOpenAIRequest) GetTokenCountMeta() *types.TokenCountMeta {
var tokenCountMeta types.TokenCountMeta
var texts = make([]string, 0)
var fileMeta = make([]*types.FileMeta, 0)
if r.Prompt != nil {
switch v := r.Prompt.(type) {
case string:
texts = append(texts, v)
case []any:
for _, item := range v {
if str, ok := item.(string); ok {
texts = append(texts, str)
}
}
default:
texts = append(texts, fmt.Sprintf("%v", r.Prompt))
}
}
if r.Input != nil {
inputs := r.ParseInput()
texts = append(texts, inputs...)
}
if r.MaxCompletionTokens > r.MaxTokens {
tokenCountMeta.MaxTokens = int(r.MaxCompletionTokens)
} else {
tokenCountMeta.MaxTokens = int(r.MaxTokens)
}
for _, message := range r.Messages {
tokenCountMeta.MessagesCount++
texts = append(texts, message.Role)
if message.Content != nil {
if message.Name != nil {
tokenCountMeta.NameCount++
texts = append(texts, *message.Name)
}
arrayContent := message.ParseContent()
for _, m := range arrayContent {
if m.Type == ContentTypeImageURL {
imageUrl := m.GetImageMedia()
if imageUrl != nil {
meta := &types.FileMeta{
FileType: types.FileTypeImage,
}
meta.Data = imageUrl.Url
meta.Detail = imageUrl.Detail
fileMeta = append(fileMeta, meta)
}
} else if m.Type == ContentTypeInputAudio {
inputAudio := m.GetInputAudio()
if inputAudio != nil {
meta := &types.FileMeta{
FileType: types.FileTypeAudio,
}
meta.Data = inputAudio.Data
fileMeta = append(fileMeta, meta)
}
} else if m.Type == ContentTypeFile {
file := m.GetFile()
if file != nil {
meta := &types.FileMeta{
FileType: types.FileTypeFile,
}
meta.Data = file.FileData
fileMeta = append(fileMeta, meta)
}
} else if m.Type == ContentTypeVideoUrl {
videoUrl := m.GetVideoUrl()
if videoUrl != nil {
meta := &types.FileMeta{
FileType: types.FileTypeVideo,
}
meta.Data = videoUrl.Url
fileMeta = append(fileMeta, meta)
}
} else {
texts = append(texts, m.Text)
}
}
}
}
if r.Tools != nil {
openaiTools := r.Tools
for _, tool := range openaiTools {
tokenCountMeta.ToolsCount++
texts = append(texts, tool.Function.Name)
if tool.Function.Description != "" {
texts = append(texts, tool.Function.Description)
}
if tool.Function.Parameters != nil {
texts = append(texts, fmt.Sprintf("%v", tool.Function.Parameters))
}
}
//toolTokens := CountTokenInput(countStr, request.Model)
//tkm += 8
//tkm += toolTokens
}
tokenCountMeta.CombineText = strings.Join(texts, "\n")
tokenCountMeta.Files = fileMeta
return &tokenCountMeta
}
func (r *GeneralOpenAIRequest) IsStream(c *gin.Context) bool {
return r.Stream
}
func (r *GeneralOpenAIRequest) ToMap() map[string]any { func (r *GeneralOpenAIRequest) ToMap() map[string]any {
result := make(map[string]any) result := make(map[string]any)
data, _ := common.Marshal(r) data, _ := common.Marshal(r)
...@@ -202,10 +316,25 @@ func (m *MediaContent) GetFile() *MessageFile { ...@@ -202,10 +316,25 @@ func (m *MediaContent) GetFile() *MessageFile {
return nil return nil
} }
func (m *MediaContent) GetVideoUrl() *MessageVideoUrl {
if m.VideoUrl != nil {
if _, ok := m.VideoUrl.(*MessageVideoUrl); ok {
return m.VideoUrl.(*MessageVideoUrl)
}
if itemMap, ok := m.VideoUrl.(map[string]any); ok {
out := &MessageVideoUrl{
Url: common.Interface2String(itemMap["url"]),
}
return out
}
}
return nil
}
type MessageImageUrl struct { type MessageImageUrl struct {
Url string `json:"url"` Url string `json:"url"`
Detail string `json:"detail"` Detail string `json:"detail"`
MimeType string //MimeType string
} }
func (m *MessageImageUrl) IsRemoteImage() bool { func (m *MessageImageUrl) IsRemoteImage() bool {
...@@ -233,6 +362,7 @@ const ( ...@@ -233,6 +362,7 @@ const (
ContentTypeInputAudio = "input_audio" ContentTypeInputAudio = "input_audio"
ContentTypeFile = "file" ContentTypeFile = "file"
ContentTypeVideoUrl = "video_url" // 阿里百炼视频识别 ContentTypeVideoUrl = "video_url" // 阿里百炼视频识别
//ContentTypeAudioUrl = "audio_url"
) )
func (m *Message) GetPrefix() bool { func (m *Message) GetPrefix() bool {
...@@ -623,7 +753,7 @@ type WebSearchOptions struct { ...@@ -623,7 +753,7 @@ type WebSearchOptions struct {
// https://platform.openai.com/docs/api-reference/responses/create // https://platform.openai.com/docs/api-reference/responses/create
type OpenAIResponsesRequest struct { type OpenAIResponsesRequest struct {
Model string `json:"model"` Model string `json:"model"`
Input json.RawMessage `json:"input,omitempty"` Input any `json:"input,omitempty"`
Include json.RawMessage `json:"include,omitempty"` Include json.RawMessage `json:"include,omitempty"`
Instructions json.RawMessage `json:"instructions,omitempty"` Instructions json.RawMessage `json:"instructions,omitempty"`
MaxOutputTokens uint `json:"max_output_tokens,omitempty"` MaxOutputTokens uint `json:"max_output_tokens,omitempty"`
...@@ -645,28 +775,145 @@ type OpenAIResponsesRequest struct { ...@@ -645,28 +775,145 @@ type OpenAIResponsesRequest struct {
Prompt json.RawMessage `json:"prompt,omitempty"` Prompt json.RawMessage `json:"prompt,omitempty"`
} }
func (r *OpenAIResponsesRequest) GetTokenCountMeta() *types.TokenCountMeta {
var fileMeta = make([]*types.FileMeta, 0)
var texts = make([]string, 0)
if r.Input != nil {
inputs := r.ParseInput()
for _, input := range inputs {
if input.Type == "input_image" {
fileMeta = append(fileMeta, &types.FileMeta{
FileType: types.FileTypeImage,
Data: input.ImageUrl,
Detail: input.Detail,
})
} else if input.Type == "input_file" {
fileMeta = append(fileMeta, &types.FileMeta{
FileType: types.FileTypeFile,
Data: input.FileUrl,
})
} else {
texts = append(texts, input.Text)
}
}
}
if len(r.Instructions) > 0 {
texts = append(texts, string(r.Instructions))
}
if len(r.Metadata) > 0 {
texts = append(texts, string(r.Metadata))
}
if len(r.Text) > 0 {
texts = append(texts, string(r.Text))
}
if len(r.ToolChoice) > 0 {
texts = append(texts, string(r.ToolChoice))
}
if len(r.Prompt) > 0 {
texts = append(texts, string(r.Prompt))
}
if len(r.Tools) > 0 {
toolStr, _ := common.Marshal(r.Tools)
texts = append(texts, string(toolStr))
}
return &types.TokenCountMeta{
CombineText: strings.Join(texts, "\n"),
Files: fileMeta,
MaxTokens: int(r.MaxOutputTokens),
}
}
func (r *OpenAIResponsesRequest) IsStream(c *gin.Context) bool {
return r.Stream
}
type Reasoning struct { type Reasoning struct {
Effort string `json:"effort,omitempty"` Effort string `json:"effort,omitempty"`
Summary string `json:"summary,omitempty"` Summary string `json:"summary,omitempty"`
} }
//type ResponsesToolsCall struct { type MediaInput struct {
// Type string `json:"type"` Type string `json:"type"`
// // Web Search Text string `json:"text,omitempty"`
// UserLocation json.RawMessage `json:"user_location,omitempty"` FileUrl string `json:"file_url,omitempty"`
// SearchContextSize string `json:"search_context_size,omitempty"` ImageUrl string `json:"image_url,omitempty"`
// // File Search Detail string `json:"detail,omitempty"` // 仅 input_image 有效
// VectorStoreIds []string `json:"vector_store_ids,omitempty"` }
// MaxNumResults uint `json:"max_num_results,omitempty"`
// Filters json.RawMessage `json:"filters,omitempty"` // ParseInput parses the Responses API `input` field into a normalized slice of MediaInput.
// // Computer Use // Reference implementation mirrors Message.ParseContent:
// DisplayWidth uint `json:"display_width,omitempty"` // - input can be a string, treated as an input_text item
// DisplayHeight uint `json:"display_height,omitempty"` // - input can be an array of objects with a `type` field
// Environment string `json:"environment,omitempty"` // supported types: input_text, input_image, input_file
// // Function func (r *OpenAIResponsesRequest) ParseInput() []MediaInput {
// Name string `json:"name,omitempty"` if r.Input == nil {
// Description string `json:"description,omitempty"` return nil
// Parameters json.RawMessage `json:"parameters,omitempty"` }
// Function json.RawMessage `json:"function,omitempty"`
// Container json.RawMessage `json:"container,omitempty"` var inputs []MediaInput
//}
// Try string first
if str, ok := r.Input.(string); ok {
inputs = append(inputs, MediaInput{Type: "input_text", Text: str})
return inputs
}
// Try array of parts
if array, ok := r.Input.([]any); ok {
for _, itemAny := range array {
// Already parsed MediaInput
if media, ok := itemAny.(MediaInput); ok {
inputs = append(inputs, media)
continue
}
// Generic map
item, ok := itemAny.(map[string]any)
if !ok {
continue
}
typeVal, ok := item["type"].(string)
if !ok {
continue
}
switch typeVal {
case "input_text":
text, _ := item["text"].(string)
inputs = append(inputs, MediaInput{Type: "input_text", Text: text})
case "input_image":
// image_url may be string or object with url field
var imageUrl string
switch v := item["image_url"].(type) {
case string:
imageUrl = v
case map[string]any:
if url, ok := v["url"].(string); ok {
imageUrl = url
}
}
inputs = append(inputs, MediaInput{Type: "input_image", ImageUrl: imageUrl})
case "input_file":
// file_url may be string or object with url field
var fileUrl string
switch v := item["file_url"].(type) {
case string:
fileUrl = v
case map[string]any:
if url, ok := v["url"].(string); ok {
fileUrl = url
}
}
inputs = append(inputs, MediaInput{Type: "input_file", FileUrl: fileUrl})
}
}
}
return inputs
}
package dto
import (
"github.com/gin-gonic/gin"
"one-api/types"
)
type Request interface {
GetTokenCountMeta() *types.TokenCountMeta
IsStream(c *gin.Context) bool
}
package dto package dto
import (
"fmt"
"github.com/gin-gonic/gin"
"one-api/types"
"strings"
)
type RerankRequest struct { type RerankRequest struct {
Documents []any `json:"documents"` Documents []any `json:"documents"`
Query string `json:"query"` Query string `json:"query"`
...@@ -10,6 +17,26 @@ type RerankRequest struct { ...@@ -10,6 +17,26 @@ type RerankRequest struct {
OverLapTokens int `json:"overlap_tokens,omitempty"` OverLapTokens int `json:"overlap_tokens,omitempty"`
} }
func (r *RerankRequest) IsStream(c *gin.Context) bool {
return false
}
func (r *RerankRequest) GetTokenCountMeta() *types.TokenCountMeta {
var texts = make([]string, 0)
for _, document := range r.Documents {
texts = append(texts, fmt.Sprintf("%v", document))
}
if r.Query != "" {
texts = append(texts, r.Query)
}
return &types.TokenCountMeta{
CombineText: strings.Join(texts, "\n"),
}
}
func (r *RerankRequest) GetReturnDocuments() bool { func (r *RerankRequest) GetReturnDocuments() bool {
if r.ReturnDocuments == nil { if r.ReturnDocuments == nil {
return false return false
......
package logger
import (
"context"
"encoding/json"
"fmt"
"github.com/bytedance/gopkg/util/gopool"
"github.com/gin-gonic/gin"
"io"
"log"
"one-api/common"
"os"
"path/filepath"
"sync"
"time"
)
const (
loggerINFO = "INFO"
loggerWarn = "WARN"
loggerError = "ERR"
loggerDebug = "DEBUG"
)
const maxLogCount = 1000000
var logCount int
var setupLogLock sync.Mutex
var setupLogWorking bool
func SetupLogger() {
if *common.LogDir != "" {
ok := setupLogLock.TryLock()
if !ok {
log.Println("setup log is already working")
return
}
defer func() {
setupLogLock.Unlock()
setupLogWorking = false
}()
logPath := filepath.Join(*common.LogDir, fmt.Sprintf("oneapi-%s.log", time.Now().Format("20060102150405")))
fd, err := os.OpenFile(logPath, os.O_APPEND|os.O_CREATE|os.O_WRONLY, 0644)
if err != nil {
log.Fatal("failed to open log file")
}
gin.DefaultWriter = io.MultiWriter(os.Stdout, fd)
gin.DefaultErrorWriter = io.MultiWriter(os.Stderr, fd)
}
}
func LogInfo(ctx context.Context, msg string) {
logHelper(ctx, loggerINFO, msg)
}
func LogWarn(ctx context.Context, msg string) {
logHelper(ctx, loggerWarn, msg)
}
func LogError(ctx context.Context, msg string) {
logHelper(ctx, loggerError, msg)
}
func LogDebug(ctx context.Context, msg string) {
if common.DebugEnabled {
logHelper(ctx, loggerDebug, msg)
}
}
func logHelper(ctx context.Context, level string, msg string) {
writer := gin.DefaultErrorWriter
if level == loggerINFO {
writer = gin.DefaultWriter
}
id := ctx.Value(common.RequestIdKey)
if id == nil {
id = "SYSTEM"
}
now := time.Now()
_, _ = fmt.Fprintf(writer, "[%s] %v | %s | %s \n", level, now.Format("2006/01/02 - 15:04:05"), id, msg)
logCount++ // we don't need accurate count, so no lock here
if logCount > maxLogCount && !setupLogWorking {
logCount = 0
setupLogWorking = true
gopool.Go(func() {
SetupLogger()
})
}
}
func LogQuota(quota int) string {
if common.DisplayInCurrencyEnabled {
return fmt.Sprintf("$%.6f 额度", float64(quota)/common.QuotaPerUnit)
} else {
return fmt.Sprintf("%d 点额度", quota)
}
}
func FormatQuota(quota int) string {
if common.DisplayInCurrencyEnabled {
return fmt.Sprintf("$%.6f", float64(quota)/common.QuotaPerUnit)
} else {
return fmt.Sprintf("%d", quota)
}
}
// LogJson 仅供测试使用 only for test
func LogJson(ctx context.Context, msg string, obj any) {
jsonStr, err := json.Marshal(obj)
if err != nil {
LogError(ctx, fmt.Sprintf("json marshal failed: %s", err.Error()))
return
}
LogInfo(ctx, fmt.Sprintf("%s | %s", msg, string(jsonStr)))
}
...@@ -8,6 +8,7 @@ import ( ...@@ -8,6 +8,7 @@ import (
"one-api/common" "one-api/common"
"one-api/constant" "one-api/constant"
"one-api/controller" "one-api/controller"
"one-api/logger"
"one-api/middleware" "one-api/middleware"
"one-api/model" "one-api/model"
"one-api/router" "one-api/router"
...@@ -35,22 +36,22 @@ func main() { ...@@ -35,22 +36,22 @@ func main() {
err := InitResources() err := InitResources()
if err != nil { if err != nil {
common.FatalLog("failed to initialize resources: " + err.Error()) logger.FatalLog("failed to initialize resources: " + err.Error())
return return
} }
common.SysLog("New API " + common.Version + " started") logger.SysLog("New API " + common.Version + " started")
if os.Getenv("GIN_MODE") != "debug" { if os.Getenv("GIN_MODE") != "debug" {
gin.SetMode(gin.ReleaseMode) gin.SetMode(gin.ReleaseMode)
} }
if common.DebugEnabled { if common.DebugEnabled {
common.SysLog("running in debug mode") logger.SysLog("running in debug mode")
} }
defer func() { defer func() {
err := model.CloseDB() err := model.CloseDB()
if err != nil { if err != nil {
common.FatalLog("failed to close database: " + err.Error()) logger.FatalLog("failed to close database: " + err.Error())
} }
}() }()
...@@ -59,18 +60,18 @@ func main() { ...@@ -59,18 +60,18 @@ func main() {
common.MemoryCacheEnabled = true common.MemoryCacheEnabled = true
} }
if common.MemoryCacheEnabled { if common.MemoryCacheEnabled {
common.SysLog("memory cache enabled") logger.SysLog("memory cache enabled")
common.SysError(fmt.Sprintf("sync frequency: %d seconds", common.SyncFrequency)) logger.SysError(fmt.Sprintf("sync frequency: %d seconds", common.SyncFrequency))
// Add panic recovery and retry for InitChannelCache // Add panic recovery and retry for InitChannelCache
func() { func() {
defer func() { defer func() {
if r := recover(); r != nil { if r := recover(); r != nil {
common.SysError(fmt.Sprintf("InitChannelCache panic: %v, retrying once", r)) logger.SysError(fmt.Sprintf("InitChannelCache panic: %v, retrying once", r))
// Retry once // Retry once
_, _, fixErr := model.FixAbility() _, _, fixErr := model.FixAbility()
if fixErr != nil { if fixErr != nil {
common.FatalLog(fmt.Sprintf("InitChannelCache failed: %s", fixErr.Error())) logger.FatalLog(fmt.Sprintf("InitChannelCache failed: %s", fixErr.Error()))
} }
} }
}() }()
...@@ -89,14 +90,14 @@ func main() { ...@@ -89,14 +90,14 @@ func main() {
if os.Getenv("CHANNEL_UPDATE_FREQUENCY") != "" { if os.Getenv("CHANNEL_UPDATE_FREQUENCY") != "" {
frequency, err := strconv.Atoi(os.Getenv("CHANNEL_UPDATE_FREQUENCY")) frequency, err := strconv.Atoi(os.Getenv("CHANNEL_UPDATE_FREQUENCY"))
if err != nil { if err != nil {
common.FatalLog("failed to parse CHANNEL_UPDATE_FREQUENCY: " + err.Error()) logger.FatalLog("failed to parse CHANNEL_UPDATE_FREQUENCY: " + err.Error())
} }
go controller.AutomaticallyUpdateChannels(frequency) go controller.AutomaticallyUpdateChannels(frequency)
} }
if os.Getenv("CHANNEL_TEST_FREQUENCY") != "" { if os.Getenv("CHANNEL_TEST_FREQUENCY") != "" {
frequency, err := strconv.Atoi(os.Getenv("CHANNEL_TEST_FREQUENCY")) frequency, err := strconv.Atoi(os.Getenv("CHANNEL_TEST_FREQUENCY"))
if err != nil { if err != nil {
common.FatalLog("failed to parse CHANNEL_TEST_FREQUENCY: " + err.Error()) logger.FatalLog("failed to parse CHANNEL_TEST_FREQUENCY: " + err.Error())
} }
go controller.AutomaticallyTestChannels(frequency) go controller.AutomaticallyTestChannels(frequency)
} }
...@@ -110,7 +111,7 @@ func main() { ...@@ -110,7 +111,7 @@ func main() {
} }
if os.Getenv("BATCH_UPDATE_ENABLED") == "true" { if os.Getenv("BATCH_UPDATE_ENABLED") == "true" {
common.BatchUpdateEnabled = true common.BatchUpdateEnabled = true
common.SysLog("batch update enabled with interval " + strconv.Itoa(common.BatchUpdateInterval) + "s") logger.SysLog("batch update enabled with interval " + strconv.Itoa(common.BatchUpdateInterval) + "s")
model.InitBatchUpdater() model.InitBatchUpdater()
} }
...@@ -119,13 +120,13 @@ func main() { ...@@ -119,13 +120,13 @@ func main() {
log.Println(http.ListenAndServe("0.0.0.0:8005", nil)) log.Println(http.ListenAndServe("0.0.0.0:8005", nil))
}) })
go common.Monitor() go common.Monitor()
common.SysLog("pprof enabled") logger.SysLog("pprof enabled")
} }
// Initialize HTTP server // Initialize HTTP server
server := gin.New() server := gin.New()
server.Use(gin.CustomRecovery(func(c *gin.Context, err any) { server.Use(gin.CustomRecovery(func(c *gin.Context, err any) {
common.SysError(fmt.Sprintf("panic detected: %v", err)) logger.SysError(fmt.Sprintf("panic detected: %v", err))
c.JSON(http.StatusInternalServerError, gin.H{ c.JSON(http.StatusInternalServerError, gin.H{
"error": gin.H{ "error": gin.H{
"message": fmt.Sprintf("Panic detected, error: %v. Please submit a issue here: https://github.com/Calcium-Ion/new-api", err), "message": fmt.Sprintf("Panic detected, error: %v. Please submit a issue here: https://github.com/Calcium-Ion/new-api", err),
...@@ -155,7 +156,7 @@ func main() { ...@@ -155,7 +156,7 @@ func main() {
} }
err = server.Run(":" + port) err = server.Run(":" + port)
if err != nil { if err != nil {
common.FatalLog("failed to start HTTP server: " + err.Error()) logger.FatalLog("failed to start HTTP server: " + err.Error())
} }
} }
...@@ -164,14 +165,14 @@ func InitResources() error { ...@@ -164,14 +165,14 @@ func InitResources() error {
// This is a placeholder function for future resource initialization // This is a placeholder function for future resource initialization
err := godotenv.Load(".env") err := godotenv.Load(".env")
if err != nil { if err != nil {
common.SysLog("未找到 .env 文件,使用默认环境变量,如果需要,请创建 .env 文件并设置相关变量") logger.SysLog("未找到 .env 文件,使用默认环境变量,如果需要,请创建 .env 文件并设置相关变量")
common.SysLog("No .env file found, using default environment variables. If needed, please create a .env file and set the relevant variables.") logger.SysLog("No .env file found, using default environment variables. If needed, please create a .env file and set the relevant variables.")
} }
// 加载环境变量 // 加载环境变量
common.InitEnv() common.InitEnv()
common.SetupLogger() logger.SetupLogger()
// Initialize model settings // Initialize model settings
ratio_setting.InitRatioSettings() ratio_setting.InitRatioSettings()
...@@ -183,7 +184,7 @@ func InitResources() error { ...@@ -183,7 +184,7 @@ func InitResources() error {
// Initialize SQL Database // Initialize SQL Database
err = model.InitDB() err = model.InitDB()
if err != nil { if err != nil {
common.FatalLog("failed to initialize database: " + err.Error()) logger.FatalLog("failed to initialize database: " + err.Error())
return err return err
} }
......
...@@ -4,7 +4,7 @@ import ( ...@@ -4,7 +4,7 @@ import (
"fmt" "fmt"
"github.com/gin-gonic/gin" "github.com/gin-gonic/gin"
"net/http" "net/http"
"one-api/common" "one-api/logger"
"runtime/debug" "runtime/debug"
) )
...@@ -12,8 +12,8 @@ func RelayPanicRecover() gin.HandlerFunc { ...@@ -12,8 +12,8 @@ func RelayPanicRecover() gin.HandlerFunc {
return func(c *gin.Context) { return func(c *gin.Context) {
defer func() { defer func() {
if err := recover(); err != nil { if err := recover(); err != nil {
common.SysError(fmt.Sprintf("panic detected: %v", err)) logger.SysError(fmt.Sprintf("panic detected: %v", err))
common.SysError(fmt.Sprintf("stacktrace from panic: %s", string(debug.Stack()))) logger.SysError(fmt.Sprintf("stacktrace from panic: %s", string(debug.Stack())))
c.JSON(http.StatusInternalServerError, gin.H{ c.JSON(http.StatusInternalServerError, gin.H{
"error": gin.H{ "error": gin.H{
"message": fmt.Sprintf("Panic detected, error: %v. Please submit a issue here: https://github.com/Calcium-Ion/new-api", err), "message": fmt.Sprintf("Panic detected, error: %v. Please submit a issue here: https://github.com/Calcium-Ion/new-api", err),
......
...@@ -7,6 +7,7 @@ import ( ...@@ -7,6 +7,7 @@ import (
"net/http" "net/http"
"net/url" "net/url"
"one-api/common" "one-api/common"
"one-api/logger"
) )
type turnstileCheckResponse struct { type turnstileCheckResponse struct {
...@@ -37,7 +38,7 @@ func TurnstileCheck() gin.HandlerFunc { ...@@ -37,7 +38,7 @@ func TurnstileCheck() gin.HandlerFunc {
"remoteip": {c.ClientIP()}, "remoteip": {c.ClientIP()},
}) })
if err != nil { if err != nil {
common.SysError(err.Error()) logger.SysError(err.Error())
c.JSON(http.StatusOK, gin.H{ c.JSON(http.StatusOK, gin.H{
"success": false, "success": false,
"message": err.Error(), "message": err.Error(),
...@@ -49,7 +50,7 @@ func TurnstileCheck() gin.HandlerFunc { ...@@ -49,7 +50,7 @@ func TurnstileCheck() gin.HandlerFunc {
var res turnstileCheckResponse var res turnstileCheckResponse
err = json.NewDecoder(rawRes.Body).Decode(&res) err = json.NewDecoder(rawRes.Body).Decode(&res)
if err != nil { if err != nil {
common.SysError(err.Error()) logger.SysError(err.Error())
c.JSON(http.StatusOK, gin.H{ c.JSON(http.StatusOK, gin.H{
"success": false, "success": false,
"message": err.Error(), "message": err.Error(),
......
...@@ -4,6 +4,7 @@ import ( ...@@ -4,6 +4,7 @@ import (
"fmt" "fmt"
"github.com/gin-gonic/gin" "github.com/gin-gonic/gin"
"one-api/common" "one-api/common"
"one-api/logger"
) )
func abortWithOpenAiMessage(c *gin.Context, statusCode int, message string) { func abortWithOpenAiMessage(c *gin.Context, statusCode int, message string) {
...@@ -15,7 +16,7 @@ func abortWithOpenAiMessage(c *gin.Context, statusCode int, message string) { ...@@ -15,7 +16,7 @@ func abortWithOpenAiMessage(c *gin.Context, statusCode int, message string) {
}, },
}) })
c.Abort() c.Abort()
common.LogError(c.Request.Context(), fmt.Sprintf("user %d | %s", userId, message)) logger.LogError(c.Request.Context(), fmt.Sprintf("user %d | %s", userId, message))
} }
func abortWithMidjourneyMessage(c *gin.Context, statusCode int, code int, description string) { func abortWithMidjourneyMessage(c *gin.Context, statusCode int, code int, description string) {
...@@ -25,5 +26,5 @@ func abortWithMidjourneyMessage(c *gin.Context, statusCode int, code int, descri ...@@ -25,5 +26,5 @@ func abortWithMidjourneyMessage(c *gin.Context, statusCode int, code int, descri
"code": code, "code": code,
}) })
c.Abort() c.Abort()
common.LogError(c.Request.Context(), description) logger.LogError(c.Request.Context(), description)
} }
...@@ -4,6 +4,7 @@ import ( ...@@ -4,6 +4,7 @@ import (
"errors" "errors"
"fmt" "fmt"
"one-api/common" "one-api/common"
"one-api/logger"
"strings" "strings"
"sync" "sync"
...@@ -294,13 +295,13 @@ func FixAbility() (int, int, error) { ...@@ -294,13 +295,13 @@ func FixAbility() (int, int, error) {
if common.UsingSQLite { if common.UsingSQLite {
err := DB.Exec("DELETE FROM abilities").Error err := DB.Exec("DELETE FROM abilities").Error
if err != nil { if err != nil {
common.SysError(fmt.Sprintf("Delete abilities failed: %s", err.Error())) logger.SysError(fmt.Sprintf("Delete abilities failed: %s", err.Error()))
return 0, 0, err return 0, 0, err
} }
} else { } else {
err := DB.Exec("TRUNCATE TABLE abilities").Error err := DB.Exec("TRUNCATE TABLE abilities").Error
if err != nil { if err != nil {
common.SysError(fmt.Sprintf("Truncate abilities failed: %s", err.Error())) logger.SysError(fmt.Sprintf("Truncate abilities failed: %s", err.Error()))
return 0, 0, err return 0, 0, err
} }
} }
...@@ -320,7 +321,7 @@ func FixAbility() (int, int, error) { ...@@ -320,7 +321,7 @@ func FixAbility() (int, int, error) {
// Delete all abilities of this channel // Delete all abilities of this channel
err = DB.Where("channel_id IN ?", ids).Delete(&Ability{}).Error err = DB.Where("channel_id IN ?", ids).Delete(&Ability{}).Error
if err != nil { if err != nil {
common.SysError(fmt.Sprintf("Delete abilities failed: %s", err.Error())) logger.SysError(fmt.Sprintf("Delete abilities failed: %s", err.Error()))
failCount += len(chunk) failCount += len(chunk)
continue continue
} }
...@@ -328,7 +329,7 @@ func FixAbility() (int, int, error) { ...@@ -328,7 +329,7 @@ func FixAbility() (int, int, error) {
for _, channel := range chunk { for _, channel := range chunk {
err = channel.AddAbilities(nil) err = channel.AddAbilities(nil)
if err != nil { if err != nil {
common.SysError(fmt.Sprintf("Add abilities for channel %d failed: %s", channel.Id, err.Error())) logger.SysError(fmt.Sprintf("Add abilities for channel %d failed: %s", channel.Id, err.Error()))
failCount++ failCount++
} else { } else {
successCount++ successCount++
......
...@@ -9,6 +9,7 @@ import ( ...@@ -9,6 +9,7 @@ import (
"one-api/common" "one-api/common"
"one-api/constant" "one-api/constant"
"one-api/dto" "one-api/dto"
"one-api/logger"
"one-api/types" "one-api/types"
"strings" "strings"
"sync" "sync"
...@@ -209,7 +210,7 @@ func (channel *Channel) GetOtherInfo() map[string]interface{} { ...@@ -209,7 +210,7 @@ func (channel *Channel) GetOtherInfo() map[string]interface{} {
if channel.OtherInfo != "" { if channel.OtherInfo != "" {
err := common.Unmarshal([]byte(channel.OtherInfo), &otherInfo) err := common.Unmarshal([]byte(channel.OtherInfo), &otherInfo)
if err != nil { if err != nil {
common.SysError("failed to unmarshal other info: " + err.Error()) logger.SysError("failed to unmarshal other info: " + err.Error())
} }
} }
return otherInfo return otherInfo
...@@ -218,7 +219,7 @@ func (channel *Channel) GetOtherInfo() map[string]interface{} { ...@@ -218,7 +219,7 @@ func (channel *Channel) GetOtherInfo() map[string]interface{} {
func (channel *Channel) SetOtherInfo(otherInfo map[string]interface{}) { func (channel *Channel) SetOtherInfo(otherInfo map[string]interface{}) {
otherInfoBytes, err := json.Marshal(otherInfo) otherInfoBytes, err := json.Marshal(otherInfo)
if err != nil { if err != nil {
common.SysError("failed to marshal other info: " + err.Error()) logger.SysError("failed to marshal other info: " + err.Error())
return return
} }
channel.OtherInfo = string(otherInfoBytes) channel.OtherInfo = string(otherInfoBytes)
...@@ -488,7 +489,7 @@ func (channel *Channel) UpdateResponseTime(responseTime int64) { ...@@ -488,7 +489,7 @@ func (channel *Channel) UpdateResponseTime(responseTime int64) {
ResponseTime: int(responseTime), ResponseTime: int(responseTime),
}).Error }).Error
if err != nil { if err != nil {
common.SysError("failed to update response time: " + err.Error()) logger.SysError("failed to update response time: " + err.Error())
} }
} }
...@@ -498,7 +499,7 @@ func (channel *Channel) UpdateBalance(balance float64) { ...@@ -498,7 +499,7 @@ func (channel *Channel) UpdateBalance(balance float64) {
Balance: balance, Balance: balance,
}).Error }).Error
if err != nil { if err != nil {
common.SysError("failed to update balance: " + err.Error()) logger.SysError("failed to update balance: " + err.Error())
} }
} }
...@@ -614,7 +615,7 @@ func UpdateChannelStatus(channelId int, usingKey string, status int, reason stri ...@@ -614,7 +615,7 @@ func UpdateChannelStatus(channelId int, usingKey string, status int, reason stri
if shouldUpdateAbilities { if shouldUpdateAbilities {
err := UpdateAbilityStatus(channelId, status == common.ChannelStatusEnabled) err := UpdateAbilityStatus(channelId, status == common.ChannelStatusEnabled)
if err != nil { if err != nil {
common.SysError("failed to update ability status: " + err.Error()) logger.SysError("failed to update ability status: " + err.Error())
} }
} }
}() }()
...@@ -642,7 +643,7 @@ func UpdateChannelStatus(channelId int, usingKey string, status int, reason stri ...@@ -642,7 +643,7 @@ func UpdateChannelStatus(channelId int, usingKey string, status int, reason stri
} }
err = channel.Save() err = channel.Save()
if err != nil { if err != nil {
common.SysError("failed to update channel status: " + err.Error()) logger.SysError("failed to update channel status: " + err.Error())
return false return false
} }
} }
...@@ -704,7 +705,7 @@ func EditChannelByTag(tag string, newTag *string, modelMapping *string, models * ...@@ -704,7 +705,7 @@ func EditChannelByTag(tag string, newTag *string, modelMapping *string, models *
for _, channel := range channels { for _, channel := range channels {
err = channel.UpdateAbilities(nil) err = channel.UpdateAbilities(nil)
if err != nil { if err != nil {
common.SysError("failed to update abilities: " + err.Error()) logger.SysError("failed to update abilities: " + err.Error())
} }
} }
} }
...@@ -728,7 +729,7 @@ func UpdateChannelUsedQuota(id int, quota int) { ...@@ -728,7 +729,7 @@ func UpdateChannelUsedQuota(id int, quota int) {
func updateChannelUsedQuota(id int, quota int) { func updateChannelUsedQuota(id int, quota int) {
err := DB.Model(&Channel{}).Where("id = ?", id).Update("used_quota", gorm.Expr("used_quota + ?", quota)).Error err := DB.Model(&Channel{}).Where("id = ?", id).Update("used_quota", gorm.Expr("used_quota + ?", quota)).Error
if err != nil { if err != nil {
common.SysError("failed to update channel used quota: " + err.Error()) logger.SysError("failed to update channel used quota: " + err.Error())
} }
} }
...@@ -821,7 +822,7 @@ func (channel *Channel) GetSetting() dto.ChannelSettings { ...@@ -821,7 +822,7 @@ func (channel *Channel) GetSetting() dto.ChannelSettings {
if channel.Setting != nil && *channel.Setting != "" { if channel.Setting != nil && *channel.Setting != "" {
err := common.Unmarshal([]byte(*channel.Setting), &setting) err := common.Unmarshal([]byte(*channel.Setting), &setting)
if err != nil { if err != nil {
common.SysError("failed to unmarshal setting: " + err.Error()) logger.SysError("failed to unmarshal setting: " + err.Error())
channel.Setting = nil // 清空设置以避免后续错误 channel.Setting = nil // 清空设置以避免后续错误
_ = channel.Save() // 保存修改 _ = channel.Save() // 保存修改
} }
...@@ -832,7 +833,7 @@ func (channel *Channel) GetSetting() dto.ChannelSettings { ...@@ -832,7 +833,7 @@ func (channel *Channel) GetSetting() dto.ChannelSettings {
func (channel *Channel) SetSetting(setting dto.ChannelSettings) { func (channel *Channel) SetSetting(setting dto.ChannelSettings) {
settingBytes, err := common.Marshal(setting) settingBytes, err := common.Marshal(setting)
if err != nil { if err != nil {
common.SysError("failed to marshal setting: " + err.Error()) logger.SysError("failed to marshal setting: " + err.Error())
return return
} }
channel.Setting = common.GetPointer[string](string(settingBytes)) channel.Setting = common.GetPointer[string](string(settingBytes))
...@@ -843,7 +844,7 @@ func (channel *Channel) GetOtherSettings() dto.ChannelOtherSettings { ...@@ -843,7 +844,7 @@ func (channel *Channel) GetOtherSettings() dto.ChannelOtherSettings {
if channel.OtherSettings != "" { if channel.OtherSettings != "" {
err := common.UnmarshalJsonStr(channel.OtherSettings, &setting) err := common.UnmarshalJsonStr(channel.OtherSettings, &setting)
if err != nil { if err != nil {
common.SysError("failed to unmarshal setting: " + err.Error()) logger.SysError("failed to unmarshal setting: " + err.Error())
channel.OtherSettings = "{}" // 清空设置以避免后续错误 channel.OtherSettings = "{}" // 清空设置以避免后续错误
_ = channel.Save() // 保存修改 _ = channel.Save() // 保存修改
} }
...@@ -854,7 +855,7 @@ func (channel *Channel) GetOtherSettings() dto.ChannelOtherSettings { ...@@ -854,7 +855,7 @@ func (channel *Channel) GetOtherSettings() dto.ChannelOtherSettings {
func (channel *Channel) SetOtherSettings(setting dto.ChannelOtherSettings) { func (channel *Channel) SetOtherSettings(setting dto.ChannelOtherSettings) {
settingBytes, err := common.Marshal(setting) settingBytes, err := common.Marshal(setting)
if err != nil { if err != nil {
common.SysError("failed to marshal setting: " + err.Error()) logger.SysError("failed to marshal setting: " + err.Error())
return return
} }
channel.OtherSettings = string(settingBytes) channel.OtherSettings = string(settingBytes)
...@@ -865,7 +866,7 @@ func (channel *Channel) GetParamOverride() map[string]interface{} { ...@@ -865,7 +866,7 @@ func (channel *Channel) GetParamOverride() map[string]interface{} {
if channel.ParamOverride != nil && *channel.ParamOverride != "" { if channel.ParamOverride != nil && *channel.ParamOverride != "" {
err := common.Unmarshal([]byte(*channel.ParamOverride), &paramOverride) err := common.Unmarshal([]byte(*channel.ParamOverride), &paramOverride)
if err != nil { if err != nil {
common.SysError("failed to unmarshal param override: " + err.Error()) logger.SysError("failed to unmarshal param override: " + err.Error())
} }
} }
return paramOverride return paramOverride
......
...@@ -6,6 +6,7 @@ import ( ...@@ -6,6 +6,7 @@ import (
"math/rand" "math/rand"
"one-api/common" "one-api/common"
"one-api/constant" "one-api/constant"
"one-api/logger"
"one-api/setting" "one-api/setting"
"one-api/setting/ratio_setting" "one-api/setting/ratio_setting"
"sort" "sort"
...@@ -84,13 +85,13 @@ func InitChannelCache() { ...@@ -84,13 +85,13 @@ func InitChannelCache() {
} }
channelsIDM = newChannelId2channel channelsIDM = newChannelId2channel
channelSyncLock.Unlock() channelSyncLock.Unlock()
common.SysLog("channels synced from database") logger.SysLog("channels synced from database")
} }
func SyncChannelCache(frequency int) { func SyncChannelCache(frequency int) {
for { for {
time.Sleep(time.Duration(frequency) * time.Second) time.Sleep(time.Duration(frequency) * time.Second)
common.SysLog("syncing channels from database") logger.SysLog("syncing channels from database")
InitChannelCache() InitChannelCache()
} }
} }
......
...@@ -4,6 +4,7 @@ import ( ...@@ -4,6 +4,7 @@ import (
"context" "context"
"fmt" "fmt"
"one-api/common" "one-api/common"
"one-api/logger"
"os" "os"
"strings" "strings"
"time" "time"
...@@ -87,13 +88,13 @@ func RecordLog(userId int, logType int, content string) { ...@@ -87,13 +88,13 @@ func RecordLog(userId int, logType int, content string) {
} }
err := LOG_DB.Create(log).Error err := LOG_DB.Create(log).Error
if err != nil { if err != nil {
common.SysError("failed to record log: " + err.Error()) logger.SysError("failed to record log: " + err.Error())
} }
} }
func RecordErrorLog(c *gin.Context, userId int, channelId int, modelName string, tokenName string, content string, tokenId int, useTimeSeconds int, func RecordErrorLog(c *gin.Context, userId int, channelId int, modelName string, tokenName string, content string, tokenId int, useTimeSeconds int,
isStream bool, group string, other map[string]interface{}) { isStream bool, group string, other map[string]interface{}) {
common.LogInfo(c, fmt.Sprintf("record error log: userId=%d, channelId=%d, modelName=%s, tokenName=%s, content=%s", userId, channelId, modelName, tokenName, content)) logger.LogInfo(c, fmt.Sprintf("record error log: userId=%d, channelId=%d, modelName=%s, tokenName=%s, content=%s", userId, channelId, modelName, tokenName, content))
username := c.GetString("username") username := c.GetString("username")
otherStr := common.MapToJsonStr(other) otherStr := common.MapToJsonStr(other)
// 判断是否需要记录 IP // 判断是否需要记录 IP
...@@ -129,7 +130,7 @@ func RecordErrorLog(c *gin.Context, userId int, channelId int, modelName string, ...@@ -129,7 +130,7 @@ func RecordErrorLog(c *gin.Context, userId int, channelId int, modelName string,
} }
err := LOG_DB.Create(log).Error err := LOG_DB.Create(log).Error
if err != nil { if err != nil {
common.LogError(c, "failed to record log: "+err.Error()) logger.LogError(c, "failed to record log: "+err.Error())
} }
} }
...@@ -142,7 +143,6 @@ type RecordConsumeLogParams struct { ...@@ -142,7 +143,6 @@ type RecordConsumeLogParams struct {
Quota int `json:"quota"` Quota int `json:"quota"`
Content string `json:"content"` Content string `json:"content"`
TokenId int `json:"token_id"` TokenId int `json:"token_id"`
UserQuota int `json:"user_quota"`
UseTimeSeconds int `json:"use_time_seconds"` UseTimeSeconds int `json:"use_time_seconds"`
IsStream bool `json:"is_stream"` IsStream bool `json:"is_stream"`
Group string `json:"group"` Group string `json:"group"`
...@@ -150,7 +150,7 @@ type RecordConsumeLogParams struct { ...@@ -150,7 +150,7 @@ type RecordConsumeLogParams struct {
} }
func RecordConsumeLog(c *gin.Context, userId int, params RecordConsumeLogParams) { func RecordConsumeLog(c *gin.Context, userId int, params RecordConsumeLogParams) {
common.LogInfo(c, fmt.Sprintf("record consume log: userId=%d, params=%s", userId, common.GetJsonString(params))) logger.LogInfo(c, fmt.Sprintf("record consume log: userId=%d, params=%s", userId, common.GetJsonString(params)))
if !common.LogConsumeEnabled { if !common.LogConsumeEnabled {
return return
} }
...@@ -189,7 +189,7 @@ func RecordConsumeLog(c *gin.Context, userId int, params RecordConsumeLogParams) ...@@ -189,7 +189,7 @@ func RecordConsumeLog(c *gin.Context, userId int, params RecordConsumeLogParams)
} }
err := LOG_DB.Create(log).Error err := LOG_DB.Create(log).Error
if err != nil { if err != nil {
common.LogError(c, "failed to record log: "+err.Error()) logger.LogError(c, "failed to record log: "+err.Error())
} }
if common.DataExportEnabled { if common.DataExportEnabled {
gopool.Go(func() { gopool.Go(func() {
......
...@@ -5,6 +5,7 @@ import ( ...@@ -5,6 +5,7 @@ import (
"log" "log"
"one-api/common" "one-api/common"
"one-api/constant" "one-api/constant"
"one-api/logger"
"os" "os"
"strings" "strings"
"sync" "sync"
...@@ -84,7 +85,7 @@ func createRootAccountIfNeed() error { ...@@ -84,7 +85,7 @@ func createRootAccountIfNeed() error {
var user User var user User
//if user.Status != common.UserStatusEnabled { //if user.Status != common.UserStatusEnabled {
if err := DB.First(&user).Error; err != nil { if err := DB.First(&user).Error; err != nil {
common.SysLog("no user exists, create a root user for you: username is root, password is 123456") logger.SysLog("no user exists, create a root user for you: username is root, password is 123456")
hashedPassword, err := common.Password2Hash("123456") hashedPassword, err := common.Password2Hash("123456")
if err != nil { if err != nil {
return err return err
...@@ -108,7 +109,7 @@ func CheckSetup() { ...@@ -108,7 +109,7 @@ func CheckSetup() {
if setup == nil { if setup == nil {
// No setup record exists, check if we have a root user // No setup record exists, check if we have a root user
if RootUserExists() { if RootUserExists() {
common.SysLog("system is not initialized, but root user exists") logger.SysLog("system is not initialized, but root user exists")
// Create setup record // Create setup record
newSetup := Setup{ newSetup := Setup{
Version: common.Version, Version: common.Version,
...@@ -116,16 +117,16 @@ func CheckSetup() { ...@@ -116,16 +117,16 @@ func CheckSetup() {
} }
err := DB.Create(&newSetup).Error err := DB.Create(&newSetup).Error
if err != nil { if err != nil {
common.SysLog("failed to create setup record: " + err.Error()) logger.SysLog("failed to create setup record: " + err.Error())
} }
constant.Setup = true constant.Setup = true
} else { } else {
common.SysLog("system is not initialized and no root user exists") logger.SysLog("system is not initialized and no root user exists")
constant.Setup = false constant.Setup = false
} }
} else { } else {
// Setup record exists, system is initialized // Setup record exists, system is initialized
common.SysLog("system is already initialized at: " + time.Unix(setup.InitializedAt, 0).String()) logger.SysLog("system is already initialized at: " + time.Unix(setup.InitializedAt, 0).String())
constant.Setup = true constant.Setup = true
} }
} }
...@@ -138,7 +139,7 @@ func chooseDB(envName string, isLog bool) (*gorm.DB, error) { ...@@ -138,7 +139,7 @@ func chooseDB(envName string, isLog bool) (*gorm.DB, error) {
if dsn != "" { if dsn != "" {
if strings.HasPrefix(dsn, "postgres://") || strings.HasPrefix(dsn, "postgresql://") { if strings.HasPrefix(dsn, "postgres://") || strings.HasPrefix(dsn, "postgresql://") {
// Use PostgreSQL // Use PostgreSQL
common.SysLog("using PostgreSQL as database") logger.SysLog("using PostgreSQL as database")
if !isLog { if !isLog {
common.UsingPostgreSQL = true common.UsingPostgreSQL = true
} else { } else {
...@@ -152,7 +153,7 @@ func chooseDB(envName string, isLog bool) (*gorm.DB, error) { ...@@ -152,7 +153,7 @@ func chooseDB(envName string, isLog bool) (*gorm.DB, error) {
}) })
} }
if strings.HasPrefix(dsn, "local") { if strings.HasPrefix(dsn, "local") {
common.SysLog("SQL_DSN not set, using SQLite as database") logger.SysLog("SQL_DSN not set, using SQLite as database")
if !isLog { if !isLog {
common.UsingSQLite = true common.UsingSQLite = true
} else { } else {
...@@ -163,7 +164,7 @@ func chooseDB(envName string, isLog bool) (*gorm.DB, error) { ...@@ -163,7 +164,7 @@ func chooseDB(envName string, isLog bool) (*gorm.DB, error) {
}) })
} }
// Use MySQL // Use MySQL
common.SysLog("using MySQL as database") logger.SysLog("using MySQL as database")
// check parseTime // check parseTime
if !strings.Contains(dsn, "parseTime") { if !strings.Contains(dsn, "parseTime") {
if strings.Contains(dsn, "?") { if strings.Contains(dsn, "?") {
...@@ -182,7 +183,7 @@ func chooseDB(envName string, isLog bool) (*gorm.DB, error) { ...@@ -182,7 +183,7 @@ func chooseDB(envName string, isLog bool) (*gorm.DB, error) {
}) })
} }
// Use SQLite // Use SQLite
common.SysLog("SQL_DSN not set, using SQLite as database") logger.SysLog("SQL_DSN not set, using SQLite as database")
common.UsingSQLite = true common.UsingSQLite = true
return gorm.Open(sqlite.Open(common.SQLitePath), &gorm.Config{ return gorm.Open(sqlite.Open(common.SQLitePath), &gorm.Config{
PrepareStmt: true, // precompile SQL PrepareStmt: true, // precompile SQL
...@@ -216,11 +217,11 @@ func InitDB() (err error) { ...@@ -216,11 +217,11 @@ func InitDB() (err error) {
if common.UsingMySQL { if common.UsingMySQL {
//_, _ = sqlDB.Exec("ALTER TABLE channels MODIFY model_mapping TEXT;") // TODO: delete this line when most users have upgraded //_, _ = sqlDB.Exec("ALTER TABLE channels MODIFY model_mapping TEXT;") // TODO: delete this line when most users have upgraded
} }
common.SysLog("database migration started") logger.SysLog("database migration started")
err = migrateDB() err = migrateDB()
return err return err
} else { } else {
common.FatalLog(err) logger.FatalLog(err)
} }
return err return err
} }
...@@ -253,11 +254,11 @@ func InitLogDB() (err error) { ...@@ -253,11 +254,11 @@ func InitLogDB() (err error) {
if !common.IsMasterNode { if !common.IsMasterNode {
return nil return nil
} }
common.SysLog("database migration started") logger.SysLog("database migration started")
err = migrateLOGDB() err = migrateLOGDB()
return err return err
} else { } else {
common.FatalLog(err) logger.FatalLog(err)
} }
return err return err
} }
...@@ -354,7 +355,7 @@ func migrateDBFast() error { ...@@ -354,7 +355,7 @@ func migrateDBFast() error {
return err return err
} }
} }
common.SysLog("database migrated") logger.SysLog("database migrated")
return nil return nil
} }
...@@ -503,6 +504,6 @@ func PingDB() error { ...@@ -503,6 +504,6 @@ func PingDB() error {
} }
lastPingTime = time.Now() lastPingTime = time.Now()
common.SysLog("Database pinged successfully") logger.SysLog("Database pinged successfully")
return nil return nil
} }
...@@ -2,6 +2,7 @@ package model ...@@ -2,6 +2,7 @@ package model
import ( import (
"one-api/common" "one-api/common"
"one-api/logger"
"one-api/setting" "one-api/setting"
"one-api/setting/config" "one-api/setting/config"
"one-api/setting/operation_setting" "one-api/setting/operation_setting"
...@@ -150,7 +151,7 @@ func loadOptionsFromDatabase() { ...@@ -150,7 +151,7 @@ func loadOptionsFromDatabase() {
for _, option := range options { for _, option := range options {
err := updateOptionMap(option.Key, option.Value) err := updateOptionMap(option.Key, option.Value)
if err != nil { if err != nil {
common.SysError("failed to update option map: " + err.Error()) logger.SysError("failed to update option map: " + err.Error())
} }
} }
} }
...@@ -158,7 +159,7 @@ func loadOptionsFromDatabase() { ...@@ -158,7 +159,7 @@ func loadOptionsFromDatabase() {
func SyncOptions(frequency int) { func SyncOptions(frequency int) {
for { for {
time.Sleep(time.Duration(frequency) * time.Second) time.Sleep(time.Duration(frequency) * time.Second)
common.SysLog("syncing options from database") logger.SysLog("syncing options from database")
loadOptionsFromDatabase() loadOptionsFromDatabase()
} }
} }
......
...@@ -3,6 +3,7 @@ package model ...@@ -3,6 +3,7 @@ package model
import ( import (
"encoding/json" "encoding/json"
"fmt" "fmt"
"one-api/logger"
"strings" "strings"
"one-api/common" "one-api/common"
...@@ -92,7 +93,7 @@ func updatePricing() { ...@@ -92,7 +93,7 @@ func updatePricing() {
//modelRatios := common.GetModelRatios() //modelRatios := common.GetModelRatios()
enableAbilities, err := GetAllEnableAbilityWithChannels() enableAbilities, err := GetAllEnableAbilityWithChannels()
if err != nil { if err != nil {
common.SysError(fmt.Sprintf("GetAllEnableAbilityWithChannels error: %v", err)) logger.SysError(fmt.Sprintf("GetAllEnableAbilityWithChannels error: %v", err))
return return
} }
// 预加载模型元数据与供应商一次,避免循环查询 // 预加载模型元数据与供应商一次,避免循环查询
......
...@@ -4,6 +4,7 @@ import ( ...@@ -4,6 +4,7 @@ import (
"errors" "errors"
"fmt" "fmt"
"one-api/common" "one-api/common"
"one-api/logger"
"strconv" "strconv"
"gorm.io/gorm" "gorm.io/gorm"
...@@ -148,7 +149,7 @@ func Redeem(key string, userId int) (quota int, err error) { ...@@ -148,7 +149,7 @@ func Redeem(key string, userId int) (quota int, err error) {
if err != nil { if err != nil {
return 0, errors.New("兑换失败," + err.Error()) return 0, errors.New("兑换失败," + err.Error())
} }
RecordLog(userId, LogTypeTopup, fmt.Sprintf("通过兑换码充值 %s,兑换码ID %d", common.LogQuota(redemption.Quota), redemption.Id)) RecordLog(userId, LogTypeTopup, fmt.Sprintf("通过兑换码充值 %s,兑换码ID %d", logger.LogQuota(redemption.Quota), redemption.Id))
return redemption.Quota, nil return redemption.Quota, nil
} }
......
...@@ -4,6 +4,7 @@ import ( ...@@ -4,6 +4,7 @@ import (
"errors" "errors"
"fmt" "fmt"
"one-api/common" "one-api/common"
"one-api/logger"
"strings" "strings"
"github.com/bytedance/gopkg/util/gopool" "github.com/bytedance/gopkg/util/gopool"
...@@ -91,7 +92,7 @@ func ValidateUserToken(key string) (token *Token, err error) { ...@@ -91,7 +92,7 @@ func ValidateUserToken(key string) (token *Token, err error) {
token.Status = common.TokenStatusExpired token.Status = common.TokenStatusExpired
err := token.SelectUpdate() err := token.SelectUpdate()
if err != nil { if err != nil {
common.SysError("failed to update token status" + err.Error()) logger.SysError("failed to update token status" + err.Error())
} }
} }
return token, errors.New("该令牌已过期") return token, errors.New("该令牌已过期")
...@@ -102,7 +103,7 @@ func ValidateUserToken(key string) (token *Token, err error) { ...@@ -102,7 +103,7 @@ func ValidateUserToken(key string) (token *Token, err error) {
token.Status = common.TokenStatusExhausted token.Status = common.TokenStatusExhausted
err := token.SelectUpdate() err := token.SelectUpdate()
if err != nil { if err != nil {
common.SysError("failed to update token status" + err.Error()) logger.SysError("failed to update token status" + err.Error())
} }
} }
keyPrefix := key[:3] keyPrefix := key[:3]
...@@ -134,7 +135,7 @@ func GetTokenById(id int) (*Token, error) { ...@@ -134,7 +135,7 @@ func GetTokenById(id int) (*Token, error) {
if shouldUpdateRedis(true, err) { if shouldUpdateRedis(true, err) {
gopool.Go(func() { gopool.Go(func() {
if err := cacheSetToken(token); err != nil { if err := cacheSetToken(token); err != nil {
common.SysError("failed to update user status cache: " + err.Error()) logger.SysError("failed to update user status cache: " + err.Error())
} }
}) })
} }
...@@ -147,7 +148,7 @@ func GetTokenByKey(key string, fromDB bool) (token *Token, err error) { ...@@ -147,7 +148,7 @@ func GetTokenByKey(key string, fromDB bool) (token *Token, err error) {
if shouldUpdateRedis(fromDB, err) && token != nil { if shouldUpdateRedis(fromDB, err) && token != nil {
gopool.Go(func() { gopool.Go(func() {
if err := cacheSetToken(*token); err != nil { if err := cacheSetToken(*token); err != nil {
common.SysError("failed to update user status cache: " + err.Error()) logger.SysError("failed to update user status cache: " + err.Error())
} }
}) })
} }
...@@ -178,7 +179,7 @@ func (token *Token) Update() (err error) { ...@@ -178,7 +179,7 @@ func (token *Token) Update() (err error) {
gopool.Go(func() { gopool.Go(func() {
err := cacheSetToken(*token) err := cacheSetToken(*token)
if err != nil { if err != nil {
common.SysError("failed to update token cache: " + err.Error()) logger.SysError("failed to update token cache: " + err.Error())
} }
}) })
} }
...@@ -194,7 +195,7 @@ func (token *Token) SelectUpdate() (err error) { ...@@ -194,7 +195,7 @@ func (token *Token) SelectUpdate() (err error) {
gopool.Go(func() { gopool.Go(func() {
err := cacheSetToken(*token) err := cacheSetToken(*token)
if err != nil { if err != nil {
common.SysError("failed to update token cache: " + err.Error()) logger.SysError("failed to update token cache: " + err.Error())
} }
}) })
} }
...@@ -209,7 +210,7 @@ func (token *Token) Delete() (err error) { ...@@ -209,7 +210,7 @@ func (token *Token) Delete() (err error) {
gopool.Go(func() { gopool.Go(func() {
err := cacheDeleteToken(token.Key) err := cacheDeleteToken(token.Key)
if err != nil { if err != nil {
common.SysError("failed to delete token cache: " + err.Error()) logger.SysError("failed to delete token cache: " + err.Error())
} }
}) })
} }
...@@ -269,7 +270,7 @@ func IncreaseTokenQuota(id int, key string, quota int) (err error) { ...@@ -269,7 +270,7 @@ func IncreaseTokenQuota(id int, key string, quota int) (err error) {
gopool.Go(func() { gopool.Go(func() {
err := cacheIncrTokenQuota(key, int64(quota)) err := cacheIncrTokenQuota(key, int64(quota))
if err != nil { if err != nil {
common.SysError("failed to increase token quota: " + err.Error()) logger.SysError("failed to increase token quota: " + err.Error())
} }
}) })
} }
...@@ -299,7 +300,7 @@ func DecreaseTokenQuota(id int, key string, quota int) (err error) { ...@@ -299,7 +300,7 @@ func DecreaseTokenQuota(id int, key string, quota int) (err error) {
gopool.Go(func() { gopool.Go(func() {
err := cacheDecrTokenQuota(key, int64(quota)) err := cacheDecrTokenQuota(key, int64(quota))
if err != nil { if err != nil {
common.SysError("failed to decrease token quota: " + err.Error()) logger.SysError("failed to decrease token quota: " + err.Error())
} }
}) })
} }
......
...@@ -4,6 +4,7 @@ import ( ...@@ -4,6 +4,7 @@ import (
"errors" "errors"
"fmt" "fmt"
"one-api/common" "one-api/common"
"one-api/logger"
"gorm.io/gorm" "gorm.io/gorm"
) )
...@@ -94,7 +95,7 @@ func Recharge(referenceId string, customerId string) (err error) { ...@@ -94,7 +95,7 @@ func Recharge(referenceId string, customerId string) (err error) {
return errors.New("充值失败," + err.Error()) return errors.New("充值失败," + err.Error())
} }
RecordLog(topUp.UserId, LogTypeTopup, fmt.Sprintf("使用在线充值成功,充值金额: %v,支付金额:%d", common.FormatQuota(int(quota)), topUp.Amount)) RecordLog(topUp.UserId, LogTypeTopup, fmt.Sprintf("使用在线充值成功,充值金额: %v,支付金额:%d", logger.FormatQuota(int(quota)), topUp.Amount))
return nil return nil
} }
...@@ -4,6 +4,7 @@ import ( ...@@ -4,6 +4,7 @@ import (
"errors" "errors"
"fmt" "fmt"
"one-api/common" "one-api/common"
"one-api/logger"
"time" "time"
"gorm.io/gorm" "gorm.io/gorm"
...@@ -243,7 +244,7 @@ func (t *TwoFA) ValidateTOTPAndUpdateUsage(code string) (bool, error) { ...@@ -243,7 +244,7 @@ func (t *TwoFA) ValidateTOTPAndUpdateUsage(code string) (bool, error) {
if !common.ValidateTOTPCode(t.Secret, code) { if !common.ValidateTOTPCode(t.Secret, code) {
// 增加失败次数 // 增加失败次数
if err := t.IncrementFailedAttempts(); err != nil { if err := t.IncrementFailedAttempts(); err != nil {
common.SysError("更新2FA失败次数失败: " + err.Error()) logger.SysError("更新2FA失败次数失败: " + err.Error())
} }
return false, nil return false, nil
} }
...@@ -255,7 +256,7 @@ func (t *TwoFA) ValidateTOTPAndUpdateUsage(code string) (bool, error) { ...@@ -255,7 +256,7 @@ func (t *TwoFA) ValidateTOTPAndUpdateUsage(code string) (bool, error) {
t.LastUsedAt = &now t.LastUsedAt = &now
if err := t.Update(); err != nil { if err := t.Update(); err != nil {
common.SysError("更新2FA使用记录失败: " + err.Error()) logger.SysError("更新2FA使用记录失败: " + err.Error())
} }
return true, nil return true, nil
...@@ -277,7 +278,7 @@ func (t *TwoFA) ValidateBackupCodeAndUpdateUsage(code string) (bool, error) { ...@@ -277,7 +278,7 @@ func (t *TwoFA) ValidateBackupCodeAndUpdateUsage(code string) (bool, error) {
if !valid { if !valid {
// 增加失败次数 // 增加失败次数
if err := t.IncrementFailedAttempts(); err != nil { if err := t.IncrementFailedAttempts(); err != nil {
common.SysError("更新2FA失败次数失败: " + err.Error()) logger.SysError("更新2FA失败次数失败: " + err.Error())
} }
return false, nil return false, nil
} }
...@@ -289,7 +290,7 @@ func (t *TwoFA) ValidateBackupCodeAndUpdateUsage(code string) (bool, error) { ...@@ -289,7 +290,7 @@ func (t *TwoFA) ValidateBackupCodeAndUpdateUsage(code string) (bool, error) {
t.LastUsedAt = &now t.LastUsedAt = &now
if err := t.Update(); err != nil { if err := t.Update(); err != nil {
common.SysError("更新2FA使用记录失败: " + err.Error()) logger.SysError("更新2FA使用记录失败: " + err.Error())
} }
return true, nil return true, nil
......
...@@ -4,6 +4,7 @@ import ( ...@@ -4,6 +4,7 @@ import (
"fmt" "fmt"
"gorm.io/gorm" "gorm.io/gorm"
"one-api/common" "one-api/common"
"one-api/logger"
"sync" "sync"
"time" "time"
) )
...@@ -24,12 +25,12 @@ func UpdateQuotaData() { ...@@ -24,12 +25,12 @@ func UpdateQuotaData() {
// recover // recover
defer func() { defer func() {
if r := recover(); r != nil { if r := recover(); r != nil {
common.SysLog(fmt.Sprintf("UpdateQuotaData panic: %s", r)) logger.SysLog(fmt.Sprintf("UpdateQuotaData panic: %s", r))
} }
}() }()
for { for {
if common.DataExportEnabled { if common.DataExportEnabled {
common.SysLog("正在更新数据看板数据...") logger.SysLog("正在更新数据看板数据...")
SaveQuotaDataCache() SaveQuotaDataCache()
} }
time.Sleep(time.Duration(common.DataExportInterval) * time.Minute) time.Sleep(time.Duration(common.DataExportInterval) * time.Minute)
...@@ -91,7 +92,7 @@ func SaveQuotaDataCache() { ...@@ -91,7 +92,7 @@ func SaveQuotaDataCache() {
} }
} }
CacheQuotaData = make(map[string]*QuotaData) CacheQuotaData = make(map[string]*QuotaData)
common.SysLog(fmt.Sprintf("保存数据看板数据成功,共保存%d条数据", size)) logger.SysLog(fmt.Sprintf("保存数据看板数据成功,共保存%d条数据", size))
} }
func increaseQuotaData(userId int, username string, modelName string, count int, quota int, createdAt int64, tokenUsed int) { func increaseQuotaData(userId int, username string, modelName string, count int, quota int, createdAt int64, tokenUsed int) {
...@@ -102,7 +103,7 @@ func increaseQuotaData(userId int, username string, modelName string, count int, ...@@ -102,7 +103,7 @@ func increaseQuotaData(userId int, username string, modelName string, count int,
"token_used": gorm.Expr("token_used + ?", tokenUsed), "token_used": gorm.Expr("token_used + ?", tokenUsed),
}).Error }).Error
if err != nil { if err != nil {
common.SysLog(fmt.Sprintf("increaseQuotaData error: %s", err)) logger.SysLog(fmt.Sprintf("increaseQuotaData error: %s", err))
} }
} }
......
...@@ -6,6 +6,7 @@ import ( ...@@ -6,6 +6,7 @@ import (
"fmt" "fmt"
"one-api/common" "one-api/common"
"one-api/dto" "one-api/dto"
"one-api/logger"
"strconv" "strconv"
"strings" "strings"
...@@ -75,7 +76,7 @@ func (user *User) GetSetting() dto.UserSetting { ...@@ -75,7 +76,7 @@ func (user *User) GetSetting() dto.UserSetting {
if user.Setting != "" { if user.Setting != "" {
err := json.Unmarshal([]byte(user.Setting), &setting) err := json.Unmarshal([]byte(user.Setting), &setting)
if err != nil { if err != nil {
common.SysError("failed to unmarshal setting: " + err.Error()) logger.SysError("failed to unmarshal setting: " + err.Error())
} }
} }
return setting return setting
...@@ -84,7 +85,7 @@ func (user *User) GetSetting() dto.UserSetting { ...@@ -84,7 +85,7 @@ func (user *User) GetSetting() dto.UserSetting {
func (user *User) SetSetting(setting dto.UserSetting) { func (user *User) SetSetting(setting dto.UserSetting) {
settingBytes, err := json.Marshal(setting) settingBytes, err := json.Marshal(setting)
if err != nil { if err != nil {
common.SysError("failed to marshal setting: " + err.Error()) logger.SysError("failed to marshal setting: " + err.Error())
return return
} }
user.Setting = string(settingBytes) user.Setting = string(settingBytes)
...@@ -274,7 +275,7 @@ func inviteUser(inviterId int) (err error) { ...@@ -274,7 +275,7 @@ func inviteUser(inviterId int) (err error) {
func (user *User) TransferAffQuotaToQuota(quota int) error { func (user *User) TransferAffQuotaToQuota(quota int) error {
// 检查quota是否小于最小额度 // 检查quota是否小于最小额度
if float64(quota) < common.QuotaPerUnit { if float64(quota) < common.QuotaPerUnit {
return fmt.Errorf("转移额度最小为%s!", common.LogQuota(int(common.QuotaPerUnit))) return fmt.Errorf("转移额度最小为%s!", logger.LogQuota(int(common.QuotaPerUnit)))
} }
// 开始数据库事务 // 开始数据库事务
...@@ -324,16 +325,16 @@ func (user *User) Insert(inviterId int) error { ...@@ -324,16 +325,16 @@ func (user *User) Insert(inviterId int) error {
return result.Error return result.Error
} }
if common.QuotaForNewUser > 0 { if common.QuotaForNewUser > 0 {
RecordLog(user.Id, LogTypeSystem, fmt.Sprintf("新用户注册赠送 %s", common.LogQuota(common.QuotaForNewUser))) RecordLog(user.Id, LogTypeSystem, fmt.Sprintf("新用户注册赠送 %s", logger.LogQuota(common.QuotaForNewUser)))
} }
if inviterId != 0 { if inviterId != 0 {
if common.QuotaForInvitee > 0 { if common.QuotaForInvitee > 0 {
_ = IncreaseUserQuota(user.Id, common.QuotaForInvitee, true) _ = IncreaseUserQuota(user.Id, common.QuotaForInvitee, true)
RecordLog(user.Id, LogTypeSystem, fmt.Sprintf("使用邀请码赠送 %s", common.LogQuota(common.QuotaForInvitee))) RecordLog(user.Id, LogTypeSystem, fmt.Sprintf("使用邀请码赠送 %s", logger.LogQuota(common.QuotaForInvitee)))
} }
if common.QuotaForInviter > 0 { if common.QuotaForInviter > 0 {
//_ = IncreaseUserQuota(inviterId, common.QuotaForInviter) //_ = IncreaseUserQuota(inviterId, common.QuotaForInviter)
RecordLog(inviterId, LogTypeSystem, fmt.Sprintf("邀请用户赠送 %s", common.LogQuota(common.QuotaForInviter))) RecordLog(inviterId, LogTypeSystem, fmt.Sprintf("邀请用户赠送 %s", logger.LogQuota(common.QuotaForInviter)))
_ = inviteUser(inviterId) _ = inviteUser(inviterId)
} }
} }
...@@ -517,7 +518,7 @@ func IsAdmin(userId int) bool { ...@@ -517,7 +518,7 @@ func IsAdmin(userId int) bool {
var user User var user User
err := DB.Where("id = ?", userId).Select("role").Find(&user).Error err := DB.Where("id = ?", userId).Select("role").Find(&user).Error
if err != nil { if err != nil {
common.SysError("no such user " + err.Error()) logger.SysError("no such user " + err.Error())
return false return false
} }
return user.Role >= common.RoleAdminUser return user.Role >= common.RoleAdminUser
...@@ -572,7 +573,7 @@ func GetUserQuota(id int, fromDB bool) (quota int, err error) { ...@@ -572,7 +573,7 @@ func GetUserQuota(id int, fromDB bool) (quota int, err error) {
if shouldUpdateRedis(fromDB, err) { if shouldUpdateRedis(fromDB, err) {
gopool.Go(func() { gopool.Go(func() {
if err := updateUserQuotaCache(id, quota); err != nil { if err := updateUserQuotaCache(id, quota); err != nil {
common.SysError("failed to update user quota cache: " + err.Error()) logger.SysError("failed to update user quota cache: " + err.Error())
} }
}) })
} }
...@@ -610,7 +611,7 @@ func GetUserGroup(id int, fromDB bool) (group string, err error) { ...@@ -610,7 +611,7 @@ func GetUserGroup(id int, fromDB bool) (group string, err error) {
if shouldUpdateRedis(fromDB, err) { if shouldUpdateRedis(fromDB, err) {
gopool.Go(func() { gopool.Go(func() {
if err := updateUserGroupCache(id, group); err != nil { if err := updateUserGroupCache(id, group); err != nil {
common.SysError("failed to update user group cache: " + err.Error()) logger.SysError("failed to update user group cache: " + err.Error())
} }
}) })
} }
...@@ -639,7 +640,7 @@ func GetUserSetting(id int, fromDB bool) (settingMap dto.UserSetting, err error) ...@@ -639,7 +640,7 @@ func GetUserSetting(id int, fromDB bool) (settingMap dto.UserSetting, err error)
if shouldUpdateRedis(fromDB, err) { if shouldUpdateRedis(fromDB, err) {
gopool.Go(func() { gopool.Go(func() {
if err := updateUserSettingCache(id, setting); err != nil { if err := updateUserSettingCache(id, setting); err != nil {
common.SysError("failed to update user setting cache: " + err.Error()) logger.SysError("failed to update user setting cache: " + err.Error())
} }
}) })
} }
...@@ -669,7 +670,7 @@ func IncreaseUserQuota(id int, quota int, db bool) (err error) { ...@@ -669,7 +670,7 @@ func IncreaseUserQuota(id int, quota int, db bool) (err error) {
gopool.Go(func() { gopool.Go(func() {
err := cacheIncrUserQuota(id, int64(quota)) err := cacheIncrUserQuota(id, int64(quota))
if err != nil { if err != nil {
common.SysError("failed to increase user quota: " + err.Error()) logger.SysError("failed to increase user quota: " + err.Error())
} }
}) })
if !db && common.BatchUpdateEnabled { if !db && common.BatchUpdateEnabled {
...@@ -694,7 +695,7 @@ func DecreaseUserQuota(id int, quota int) (err error) { ...@@ -694,7 +695,7 @@ func DecreaseUserQuota(id int, quota int) (err error) {
gopool.Go(func() { gopool.Go(func() {
err := cacheDecrUserQuota(id, int64(quota)) err := cacheDecrUserQuota(id, int64(quota))
if err != nil { if err != nil {
common.SysError("failed to decrease user quota: " + err.Error()) logger.SysError("failed to decrease user quota: " + err.Error())
} }
}) })
if common.BatchUpdateEnabled { if common.BatchUpdateEnabled {
...@@ -750,7 +751,7 @@ func updateUserUsedQuotaAndRequestCount(id int, quota int, count int) { ...@@ -750,7 +751,7 @@ func updateUserUsedQuotaAndRequestCount(id int, quota int, count int) {
}, },
).Error ).Error
if err != nil { if err != nil {
common.SysError("failed to update user used quota and request count: " + err.Error()) logger.SysError("failed to update user used quota and request count: " + err.Error())
return return
} }
...@@ -767,14 +768,14 @@ func updateUserUsedQuota(id int, quota int) { ...@@ -767,14 +768,14 @@ func updateUserUsedQuota(id int, quota int) {
}, },
).Error ).Error
if err != nil { if err != nil {
common.SysError("failed to update user used quota: " + err.Error()) logger.SysError("failed to update user used quota: " + err.Error())
} }
} }
func updateUserRequestCount(id int, count int) { func updateUserRequestCount(id int, count int) {
err := DB.Model(&User{}).Where("id = ?", id).Update("request_count", gorm.Expr("request_count + ?", count)).Error err := DB.Model(&User{}).Where("id = ?", id).Update("request_count", gorm.Expr("request_count + ?", count)).Error
if err != nil { if err != nil {
common.SysError("failed to update user request count: " + err.Error()) logger.SysError("failed to update user request count: " + err.Error())
} }
} }
...@@ -785,7 +786,7 @@ func GetUsernameById(id int, fromDB bool) (username string, err error) { ...@@ -785,7 +786,7 @@ func GetUsernameById(id int, fromDB bool) (username string, err error) {
if shouldUpdateRedis(fromDB, err) { if shouldUpdateRedis(fromDB, err) {
gopool.Go(func() { gopool.Go(func() {
if err := updateUserNameCache(id, username); err != nil { if err := updateUserNameCache(id, username); err != nil {
common.SysError("failed to update user name cache: " + err.Error()) logger.SysError("failed to update user name cache: " + err.Error())
} }
}) })
} }
......
...@@ -5,6 +5,7 @@ import ( ...@@ -5,6 +5,7 @@ import (
"one-api/common" "one-api/common"
"one-api/constant" "one-api/constant"
"one-api/dto" "one-api/dto"
"one-api/logger"
"time" "time"
"github.com/gin-gonic/gin" "github.com/gin-gonic/gin"
...@@ -37,7 +38,7 @@ func (user *UserBase) GetSetting() dto.UserSetting { ...@@ -37,7 +38,7 @@ func (user *UserBase) GetSetting() dto.UserSetting {
if user.Setting != "" { if user.Setting != "" {
err := common.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()) logger.SysError("failed to unmarshal setting: " + err.Error())
} }
} }
return setting return setting
...@@ -78,7 +79,7 @@ func GetUserCache(userId int) (userCache *UserBase, err error) { ...@@ -78,7 +79,7 @@ func GetUserCache(userId int) (userCache *UserBase, err error) {
if shouldUpdateRedis(fromDB, err) && user != nil { if shouldUpdateRedis(fromDB, err) && user != nil {
gopool.Go(func() { gopool.Go(func() {
if err := updateUserCache(*user); err != nil { if err := updateUserCache(*user); err != nil {
common.SysError("failed to update user status cache: " + err.Error()) logger.SysError("failed to update user status cache: " + err.Error())
} }
}) })
} }
......
...@@ -3,6 +3,7 @@ package model ...@@ -3,6 +3,7 @@ package model
import ( import (
"errors" "errors"
"one-api/common" "one-api/common"
"one-api/logger"
"sync" "sync"
"time" "time"
...@@ -65,7 +66,7 @@ func batchUpdate() { ...@@ -65,7 +66,7 @@ func batchUpdate() {
return return
} }
common.SysLog("batch update started") logger.SysLog("batch update started")
for i := 0; i < BatchUpdateTypeCount; i++ { for i := 0; i < BatchUpdateTypeCount; i++ {
batchUpdateLocks[i].Lock() batchUpdateLocks[i].Lock()
store := batchUpdateStores[i] store := batchUpdateStores[i]
...@@ -77,12 +78,12 @@ func batchUpdate() { ...@@ -77,12 +78,12 @@ func batchUpdate() {
case BatchUpdateTypeUserQuota: case BatchUpdateTypeUserQuota:
err := increaseUserQuota(key, value) err := increaseUserQuota(key, value)
if err != nil { if err != nil {
common.SysError("failed to batch update user quota: " + err.Error()) logger.SysError("failed to batch update user quota: " + err.Error())
} }
case BatchUpdateTypeTokenQuota: case BatchUpdateTypeTokenQuota:
err := increaseTokenQuota(key, value) err := increaseTokenQuota(key, value)
if err != nil { if err != nil {
common.SysError("failed to batch update token quota: " + err.Error()) logger.SysError("failed to batch update token quota: " + err.Error())
} }
case BatchUpdateTypeUsedQuota: case BatchUpdateTypeUsedQuota:
updateUserUsedQuota(key, value) updateUserUsedQuota(key, value)
...@@ -93,7 +94,7 @@ func batchUpdate() { ...@@ -93,7 +94,7 @@ func batchUpdate() {
} }
} }
} }
common.SysLog("batch update finished") logger.SysLog("batch update finished")
} }
func RecordExist(err error) (bool, error) { func RecordExist(err error) (bool, error) {
......
...@@ -4,107 +4,40 @@ import ( ...@@ -4,107 +4,40 @@ import (
"errors" "errors"
"fmt" "fmt"
"net/http" "net/http"
"one-api/common"
"one-api/dto" "one-api/dto"
relaycommon "one-api/relay/common" relaycommon "one-api/relay/common"
relayconstant "one-api/relay/constant"
"one-api/relay/helper" "one-api/relay/helper"
"one-api/service" "one-api/service"
"one-api/setting"
"one-api/types" "one-api/types"
"strings"
"github.com/gin-gonic/gin" "github.com/gin-gonic/gin"
) )
func getAndValidAudioRequest(c *gin.Context, info *relaycommon.RelayInfo) (*dto.AudioRequest, error) { func AudioHelper(c *gin.Context, info *relaycommon.RelayInfo) (newAPIError *types.NewAPIError) {
audioRequest := &dto.AudioRequest{} info.InitChannelMeta(c)
err := common.UnmarshalBodyReusable(c, audioRequest)
if err != nil {
return nil, err
}
switch info.RelayMode {
case relayconstant.RelayModeAudioSpeech:
if audioRequest.Model == "" {
return nil, errors.New("model is required")
}
if setting.ShouldCheckPromptSensitive() {
words, err := service.CheckSensitiveInput(audioRequest.Input)
if err != nil {
common.LogWarn(c, fmt.Sprintf("user sensitive words detected: %s", strings.Join(words, ",")))
return nil, err
}
}
default:
err = c.Request.ParseForm()
if err != nil {
return nil, err
}
formData := c.Request.PostForm
if audioRequest.Model == "" {
audioRequest.Model = formData.Get("model")
}
if audioRequest.Model == "" { audioRequest, ok := info.Request.(*dto.AudioRequest)
return nil, errors.New("model is required") if !ok {
} return types.NewError(errors.New("invalid request type"), types.ErrorCodeInvalidRequest, types.ErrOptionWithSkipRetry())
audioRequest.ResponseFormat = formData.Get("response_format")
if audioRequest.ResponseFormat == "" {
audioRequest.ResponseFormat = "json"
}
} }
return audioRequest, nil
}
func AudioHelper(c *gin.Context) (newAPIError *types.NewAPIError) {
relayInfo := relaycommon.GenRelayInfoOpenAIAudio(c)
audioRequest, err := getAndValidAudioRequest(c, relayInfo)
if err != nil {
common.LogError(c, fmt.Sprintf("getAndValidAudioRequest failed: %s", err.Error()))
return types.NewError(err, types.ErrorCodeInvalidRequest, types.ErrOptionWithSkipRetry())
}
promptTokens := 0
preConsumedTokens := common.PreConsumedQuota
if relayInfo.RelayMode == relayconstant.RelayModeAudioSpeech {
promptTokens = service.CountTTSToken(audioRequest.Input, audioRequest.Model)
preConsumedTokens = promptTokens
relayInfo.PromptTokens = promptTokens
}
priceData, err := helper.ModelPriceHelper(c, relayInfo, preConsumedTokens, 0)
if err != nil {
return types.NewError(err, types.ErrorCodeModelPriceError, types.ErrOptionWithSkipRetry())
}
preConsumedQuota, userQuota, openaiErr := preConsumeQuota(c, priceData.ShouldPreConsumedQuota, relayInfo)
if openaiErr != nil {
return openaiErr
}
defer func() {
if openaiErr != nil {
returnPreConsumedQuota(c, relayInfo, userQuota, preConsumedQuota)
}
}()
err = helper.ModelMappedHelper(c, relayInfo, audioRequest) err := helper.ModelMappedHelper(c, info, audioRequest)
if err != nil { if err != nil {
return types.NewError(err, types.ErrorCodeChannelModelMappedError, types.ErrOptionWithSkipRetry()) return types.NewError(err, types.ErrorCodeChannelModelMappedError, types.ErrOptionWithSkipRetry())
} }
adaptor := GetAdaptor(relayInfo.ApiType) adaptor := GetAdaptor(info.ApiType)
if adaptor == nil { if adaptor == nil {
return types.NewError(fmt.Errorf("invalid api type: %d", relayInfo.ApiType), types.ErrorCodeInvalidApiType, types.ErrOptionWithSkipRetry()) return types.NewError(fmt.Errorf("invalid api type: %d", info.ApiType), types.ErrorCodeInvalidApiType, types.ErrOptionWithSkipRetry())
} }
adaptor.Init(relayInfo) adaptor.Init(info)
ioReader, err := adaptor.ConvertAudioRequest(c, relayInfo, *audioRequest) ioReader, err := adaptor.ConvertAudioRequest(c, info, *audioRequest)
if err != nil { if err != nil {
return types.NewError(err, types.ErrorCodeConvertRequestFailed, types.ErrOptionWithSkipRetry()) return types.NewError(err, types.ErrorCodeConvertRequestFailed, types.ErrOptionWithSkipRetry())
} }
resp, err := adaptor.DoRequest(c, relayInfo, ioReader) resp, err := adaptor.DoRequest(c, info, ioReader)
if err != nil { if err != nil {
return types.NewError(err, types.ErrorCodeDoRequestFailed) return types.NewError(err, types.ErrorCodeDoRequestFailed)
} }
...@@ -121,14 +54,14 @@ func AudioHelper(c *gin.Context) (newAPIError *types.NewAPIError) { ...@@ -121,14 +54,14 @@ func AudioHelper(c *gin.Context) (newAPIError *types.NewAPIError) {
} }
} }
usage, newAPIError := adaptor.DoResponse(c, httpResp, relayInfo) usage, newAPIError := adaptor.DoResponse(c, httpResp, info)
if newAPIError != nil { if newAPIError != nil {
// reset status code 重置状态码 // reset status code 重置状态码
service.ResetStatusCode(newAPIError, statusCodeMappingStr) service.ResetStatusCode(newAPIError, statusCodeMappingStr)
return newAPIError return newAPIError
} }
postConsumeQuota(c, relayInfo, usage.(*dto.Usage), preConsumedQuota, userQuota, priceData, "") postConsumeQuota(c, info, usage.(*dto.Usage), "")
return nil return nil
} }
...@@ -6,8 +6,8 @@ import ( ...@@ -6,8 +6,8 @@ import (
"fmt" "fmt"
"io" "io"
"net/http" "net/http"
"one-api/common"
"one-api/dto" "one-api/dto"
"one-api/logger"
relaycommon "one-api/relay/common" relaycommon "one-api/relay/common"
"one-api/service" "one-api/service"
"one-api/types" "one-api/types"
...@@ -43,7 +43,7 @@ func updateTask(info *relaycommon.RelayInfo, taskID string) (*AliResponse, error ...@@ -43,7 +43,7 @@ func updateTask(info *relaycommon.RelayInfo, taskID string) (*AliResponse, error
client := &http.Client{} client := &http.Client{}
resp, err := client.Do(req) resp, err := client.Do(req)
if err != nil { if err != nil {
common.SysError("updateTask client.Do err: " + err.Error()) logger.SysError("updateTask client.Do err: " + err.Error())
return &aliResponse, err, nil return &aliResponse, err, nil
} }
defer resp.Body.Close() defer resp.Body.Close()
...@@ -53,7 +53,7 @@ func updateTask(info *relaycommon.RelayInfo, taskID string) (*AliResponse, error ...@@ -53,7 +53,7 @@ func updateTask(info *relaycommon.RelayInfo, taskID string) (*AliResponse, error
var response AliResponse var response AliResponse
err = json.Unmarshal(responseBody, &response) err = json.Unmarshal(responseBody, &response)
if err != nil { if err != nil {
common.SysError("updateTask NewDecoder err: " + err.Error()) logger.SysError("updateTask NewDecoder err: " + err.Error())
return &aliResponse, err, nil return &aliResponse, err, nil
} }
...@@ -109,7 +109,7 @@ func responseAli2OpenAIImage(c *gin.Context, response *AliResponse, info *relayc ...@@ -109,7 +109,7 @@ func responseAli2OpenAIImage(c *gin.Context, response *AliResponse, info *relayc
if responseFormat == "b64_json" { if responseFormat == "b64_json" {
_, b64, err := service.GetImageFromUrl(data.Url) _, b64, err := service.GetImageFromUrl(data.Url)
if err != nil { if err != nil {
common.LogError(c, "get_image_data_failed: "+err.Error()) logger.LogError(c, "get_image_data_failed: "+err.Error())
continue continue
} }
b64Json = b64 b64Json = b64
...@@ -134,14 +134,14 @@ func aliImageHandler(c *gin.Context, resp *http.Response, info *relaycommon.Rela ...@@ -134,14 +134,14 @@ func aliImageHandler(c *gin.Context, resp *http.Response, info *relaycommon.Rela
if err != nil { if err != nil {
return types.NewOpenAIError(err, types.ErrorCodeReadResponseBodyFailed, http.StatusInternalServerError), nil return types.NewOpenAIError(err, types.ErrorCodeReadResponseBodyFailed, http.StatusInternalServerError), nil
} }
common.CloseResponseBodyGracefully(resp) service.CloseResponseBodyGracefully(resp)
err = json.Unmarshal(responseBody, &aliTaskResponse) err = json.Unmarshal(responseBody, &aliTaskResponse)
if err != nil { if err != nil {
return types.NewOpenAIError(err, types.ErrorCodeBadResponseBody, http.StatusInternalServerError), nil return types.NewOpenAIError(err, types.ErrorCodeBadResponseBody, http.StatusInternalServerError), nil
} }
if aliTaskResponse.Message != "" { if aliTaskResponse.Message != "" {
common.LogError(c, "ali_async_task_failed: "+aliTaskResponse.Message) logger.LogError(c, "ali_async_task_failed: "+aliTaskResponse.Message)
return types.NewError(errors.New(aliTaskResponse.Message), types.ErrorCodeBadResponse), nil return types.NewError(errors.New(aliTaskResponse.Message), types.ErrorCodeBadResponse), nil
} }
......
...@@ -4,9 +4,9 @@ import ( ...@@ -4,9 +4,9 @@ import (
"encoding/json" "encoding/json"
"io" "io"
"net/http" "net/http"
"one-api/common"
"one-api/dto" "one-api/dto"
relaycommon "one-api/relay/common" relaycommon "one-api/relay/common"
"one-api/service"
"one-api/types" "one-api/types"
"github.com/gin-gonic/gin" "github.com/gin-gonic/gin"
...@@ -36,7 +36,7 @@ func RerankHandler(c *gin.Context, resp *http.Response, info *relaycommon.RelayI ...@@ -36,7 +36,7 @@ func RerankHandler(c *gin.Context, resp *http.Response, info *relaycommon.RelayI
if err != nil { if err != nil {
return types.NewOpenAIError(err, types.ErrorCodeReadResponseBodyFailed, http.StatusInternalServerError), nil return types.NewOpenAIError(err, types.ErrorCodeReadResponseBodyFailed, http.StatusInternalServerError), nil
} }
common.CloseResponseBodyGracefully(resp) service.CloseResponseBodyGracefully(resp)
var aliResponse AliRerankResponse var aliResponse AliRerankResponse
err = json.Unmarshal(responseBody, &aliResponse) err = json.Unmarshal(responseBody, &aliResponse)
......
...@@ -7,7 +7,9 @@ import ( ...@@ -7,7 +7,9 @@ import (
"net/http" "net/http"
"one-api/common" "one-api/common"
"one-api/dto" "one-api/dto"
"one-api/logger"
"one-api/relay/helper" "one-api/relay/helper"
"one-api/service"
"strings" "strings"
"one-api/types" "one-api/types"
...@@ -46,7 +48,7 @@ func aliEmbeddingHandler(c *gin.Context, resp *http.Response) (*types.NewAPIErro ...@@ -46,7 +48,7 @@ func aliEmbeddingHandler(c *gin.Context, resp *http.Response) (*types.NewAPIErro
return types.NewOpenAIError(err, types.ErrorCodeBadResponseBody, http.StatusInternalServerError), nil return types.NewOpenAIError(err, types.ErrorCodeBadResponseBody, http.StatusInternalServerError), nil
} }
common.CloseResponseBodyGracefully(resp) service.CloseResponseBodyGracefully(resp)
model := c.GetString("model") model := c.GetString("model")
if model == "" { if model == "" {
...@@ -148,7 +150,7 @@ func aliStreamHandler(c *gin.Context, resp *http.Response) (*types.NewAPIError, ...@@ -148,7 +150,7 @@ func aliStreamHandler(c *gin.Context, resp *http.Response) (*types.NewAPIError,
var aliResponse AliResponse var aliResponse AliResponse
err := json.Unmarshal([]byte(data), &aliResponse) err := json.Unmarshal([]byte(data), &aliResponse)
if err != nil { if err != nil {
common.SysError("error unmarshalling stream response: " + err.Error()) logger.SysError("error unmarshalling stream response: " + err.Error())
return true return true
} }
if aliResponse.Usage.OutputTokens != 0 { if aliResponse.Usage.OutputTokens != 0 {
...@@ -161,7 +163,7 @@ func aliStreamHandler(c *gin.Context, resp *http.Response) (*types.NewAPIError, ...@@ -161,7 +163,7 @@ func aliStreamHandler(c *gin.Context, resp *http.Response) (*types.NewAPIError,
lastResponseText = aliResponse.Output.Text lastResponseText = aliResponse.Output.Text
jsonResponse, err := json.Marshal(response) jsonResponse, err := json.Marshal(response)
if err != nil { if err != nil {
common.SysError("error marshalling stream response: " + err.Error()) logger.SysError("error marshalling stream response: " + err.Error())
return true return true
} }
c.Render(-1, common.CustomEvent{Data: "data: " + string(jsonResponse)}) c.Render(-1, common.CustomEvent{Data: "data: " + string(jsonResponse)})
...@@ -171,7 +173,7 @@ func aliStreamHandler(c *gin.Context, resp *http.Response) (*types.NewAPIError, ...@@ -171,7 +173,7 @@ func aliStreamHandler(c *gin.Context, resp *http.Response) (*types.NewAPIError,
return false return false
} }
}) })
common.CloseResponseBodyGracefully(resp) service.CloseResponseBodyGracefully(resp)
return nil, &usage return nil, &usage
} }
...@@ -181,7 +183,7 @@ func aliHandler(c *gin.Context, resp *http.Response) (*types.NewAPIError, *dto.U ...@@ -181,7 +183,7 @@ func aliHandler(c *gin.Context, resp *http.Response) (*types.NewAPIError, *dto.U
if err != nil { if err != nil {
return types.NewOpenAIError(err, types.ErrorCodeReadResponseBodyFailed, http.StatusInternalServerError), nil return types.NewOpenAIError(err, types.ErrorCodeReadResponseBodyFailed, http.StatusInternalServerError), nil
} }
common.CloseResponseBodyGracefully(resp) service.CloseResponseBodyGracefully(resp)
err = json.Unmarshal(responseBody, &aliResponse) err = json.Unmarshal(responseBody, &aliResponse)
if err != nil { if err != nil {
return types.NewOpenAIError(err, types.ErrorCodeBadResponseBody, http.StatusInternalServerError), nil return types.NewOpenAIError(err, types.ErrorCodeBadResponseBody, http.StatusInternalServerError), nil
......
...@@ -7,6 +7,7 @@ import ( ...@@ -7,6 +7,7 @@ import (
"io" "io"
"net/http" "net/http"
common2 "one-api/common" common2 "one-api/common"
"one-api/logger"
"one-api/relay/common" "one-api/relay/common"
"one-api/relay/constant" "one-api/relay/constant"
"one-api/relay/helper" "one-api/relay/helper"
...@@ -181,7 +182,7 @@ func sendPingData(c *gin.Context, mutex *sync.Mutex) error { ...@@ -181,7 +182,7 @@ func sendPingData(c *gin.Context, mutex *sync.Mutex) error {
err := helper.PingData(c) err := helper.PingData(c)
if err != nil { if err != nil {
common2.LogError(c, "SSE ping error: "+err.Error()) logger.LogError(c, "SSE ping error: "+err.Error())
done <- err done <- err
return return
} }
......
...@@ -9,6 +9,7 @@ import ( ...@@ -9,6 +9,7 @@ import (
"one-api/common" "one-api/common"
"one-api/constant" "one-api/constant"
"one-api/dto" "one-api/dto"
"one-api/logger"
relaycommon "one-api/relay/common" relaycommon "one-api/relay/common"
"one-api/relay/helper" "one-api/relay/helper"
"one-api/service" "one-api/service"
...@@ -118,7 +119,7 @@ func baiduStreamHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http. ...@@ -118,7 +119,7 @@ func baiduStreamHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.
var baiduResponse BaiduChatStreamResponse var baiduResponse BaiduChatStreamResponse
err := common.Unmarshal([]byte(data), &baiduResponse) err := common.Unmarshal([]byte(data), &baiduResponse)
if err != nil { if err != nil {
common.SysError("error unmarshalling stream response: " + err.Error()) logger.SysError("error unmarshalling stream response: " + err.Error())
return true return true
} }
if baiduResponse.Usage.TotalTokens != 0 { if baiduResponse.Usage.TotalTokens != 0 {
...@@ -129,11 +130,11 @@ func baiduStreamHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http. ...@@ -129,11 +130,11 @@ func baiduStreamHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.
response := streamResponseBaidu2OpenAI(&baiduResponse) response := streamResponseBaidu2OpenAI(&baiduResponse)
err = helper.ObjectData(c, response) err = helper.ObjectData(c, response)
if err != nil { if err != nil {
common.SysError("error sending stream response: " + err.Error()) logger.SysError("error sending stream response: " + err.Error())
} }
return true return true
}) })
common.CloseResponseBodyGracefully(resp) service.CloseResponseBodyGracefully(resp)
return nil, usage return nil, usage
} }
...@@ -143,7 +144,7 @@ func baiduHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Respon ...@@ -143,7 +144,7 @@ func baiduHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Respon
if err != nil { if err != nil {
return types.NewError(err, types.ErrorCodeBadResponseBody), nil return types.NewError(err, types.ErrorCodeBadResponseBody), nil
} }
common.CloseResponseBodyGracefully(resp) service.CloseResponseBodyGracefully(resp)
err = json.Unmarshal(responseBody, &baiduResponse) err = json.Unmarshal(responseBody, &baiduResponse)
if err != nil { if err != nil {
return types.NewError(err, types.ErrorCodeBadResponseBody), nil return types.NewError(err, types.ErrorCodeBadResponseBody), nil
...@@ -168,7 +169,7 @@ func baiduEmbeddingHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *ht ...@@ -168,7 +169,7 @@ func baiduEmbeddingHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *ht
if err != nil { if err != nil {
return types.NewError(err, types.ErrorCodeBadResponseBody), nil return types.NewError(err, types.ErrorCodeBadResponseBody), nil
} }
common.CloseResponseBodyGracefully(resp) service.CloseResponseBodyGracefully(resp)
err = json.Unmarshal(responseBody, &baiduResponse) err = json.Unmarshal(responseBody, &baiduResponse)
if err != nil { if err != nil {
return types.NewError(err, types.ErrorCodeBadResponseBody), nil return types.NewError(err, types.ErrorCodeBadResponseBody), nil
......
...@@ -7,6 +7,7 @@ import ( ...@@ -7,6 +7,7 @@ import (
"net/http" "net/http"
"one-api/common" "one-api/common"
"one-api/dto" "one-api/dto"
"one-api/logger"
"one-api/relay/channel/openrouter" "one-api/relay/channel/openrouter"
relaycommon "one-api/relay/common" relaycommon "one-api/relay/common"
"one-api/relay/helper" "one-api/relay/helper"
...@@ -375,7 +376,7 @@ func RequestOpenAI2ClaudeMessage(textRequest dto.GeneralOpenAIRequest) (*dto.Cla ...@@ -375,7 +376,7 @@ func RequestOpenAI2ClaudeMessage(textRequest dto.GeneralOpenAIRequest) (*dto.Cla
for _, toolCall := range message.ParseToolCalls() { for _, toolCall := range message.ParseToolCalls() {
inputObj := make(map[string]any) inputObj := make(map[string]any)
if err := json.Unmarshal([]byte(toolCall.Function.Arguments), &inputObj); err != nil { if err := json.Unmarshal([]byte(toolCall.Function.Arguments), &inputObj); err != nil {
common.SysError("tool call function arguments is not a map[string]any: " + fmt.Sprintf("%v", toolCall.Function.Arguments)) logger.SysError("tool call function arguments is not a map[string]any: " + fmt.Sprintf("%v", toolCall.Function.Arguments))
continue continue
} }
claudeMediaMessages = append(claudeMediaMessages, dto.ClaudeMediaMessage{ claudeMediaMessages = append(claudeMediaMessages, dto.ClaudeMediaMessage{
...@@ -609,7 +610,7 @@ func HandleStreamResponseData(c *gin.Context, info *relaycommon.RelayInfo, claud ...@@ -609,7 +610,7 @@ func HandleStreamResponseData(c *gin.Context, info *relaycommon.RelayInfo, claud
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()) logger.SysError("error unmarshalling stream response: " + err.Error())
return types.NewError(err, types.ErrorCodeBadResponseBody) return types.NewError(err, types.ErrorCodeBadResponseBody)
} }
if claudeError := claudeResponse.GetClaudeError(); claudeError != nil && claudeError.Type != "" { if claudeError := claudeResponse.GetClaudeError(); claudeError != nil && claudeError.Type != "" {
...@@ -637,7 +638,7 @@ func HandleStreamResponseData(c *gin.Context, info *relaycommon.RelayInfo, claud ...@@ -637,7 +638,7 @@ func HandleStreamResponseData(c *gin.Context, info *relaycommon.RelayInfo, claud
err = helper.ObjectData(c, response) err = helper.ObjectData(c, response)
if err != nil { if err != nil {
common.LogError(c, "send_stream_response_failed: "+err.Error()) logger.LogError(c, "send_stream_response_failed: "+err.Error())
} }
} }
return nil return nil
...@@ -653,7 +654,7 @@ func HandleStreamFinalResponse(c *gin.Context, info *relaycommon.RelayInfo, clau ...@@ -653,7 +654,7 @@ func HandleStreamFinalResponse(c *gin.Context, info *relaycommon.RelayInfo, clau
} }
if claudeInfo.Usage.CompletionTokens == 0 || !claudeInfo.Done { if claudeInfo.Usage.CompletionTokens == 0 || !claudeInfo.Done {
if common.DebugEnabled { if common.DebugEnabled {
common.SysError("claude response usage is not complete, maybe upstream error") logger.SysError("claude response usage is not complete, maybe upstream error")
} }
claudeInfo.Usage = service.ResponseText2Usage(claudeInfo.ResponseText.String(), info.UpstreamModelName, claudeInfo.Usage.PromptTokens) claudeInfo.Usage = service.ResponseText2Usage(claudeInfo.ResponseText.String(), info.UpstreamModelName, claudeInfo.Usage.PromptTokens)
} }
...@@ -667,7 +668,7 @@ func HandleStreamFinalResponse(c *gin.Context, info *relaycommon.RelayInfo, clau ...@@ -667,7 +668,7 @@ func HandleStreamFinalResponse(c *gin.Context, info *relaycommon.RelayInfo, clau
response := helper.GenerateFinalUsageResponse(claudeInfo.ResponseId, claudeInfo.Created, info.UpstreamModelName, *claudeInfo.Usage) response := helper.GenerateFinalUsageResponse(claudeInfo.ResponseId, claudeInfo.Created, info.UpstreamModelName, *claudeInfo.Usage)
err := helper.ObjectData(c, response) err := helper.ObjectData(c, response)
if err != nil { if err != nil {
common.SysError("send final response failed: " + err.Error()) logger.SysError("send final response failed: " + err.Error())
} }
} }
helper.Done(c) helper.Done(c)
...@@ -736,12 +737,12 @@ func HandleClaudeResponseData(c *gin.Context, info *relaycommon.RelayInfo, claud ...@@ -736,12 +737,12 @@ func HandleClaudeResponseData(c *gin.Context, info *relaycommon.RelayInfo, claud
c.Set("claude_web_search_requests", claudeResponse.Usage.ServerToolUse.WebSearchRequests) c.Set("claude_web_search_requests", claudeResponse.Usage.ServerToolUse.WebSearchRequests)
} }
common.IOCopyBytesGracefully(c, nil, responseData) service.IOCopyBytesGracefully(c, nil, responseData)
return nil return nil
} }
func ClaudeHandler(c *gin.Context, resp *http.Response, info *relaycommon.RelayInfo, requestMode int) (*types.NewAPIError, *dto.Usage) { func ClaudeHandler(c *gin.Context, resp *http.Response, info *relaycommon.RelayInfo, requestMode int) (*types.NewAPIError, *dto.Usage) {
defer common.CloseResponseBodyGracefully(resp) defer service.CloseResponseBodyGracefully(resp)
claudeInfo := &ClaudeResponseInfo{ claudeInfo := &ClaudeResponseInfo{
ResponseId: helper.GetResponseID(c), ResponseId: helper.GetResponseID(c),
......
...@@ -5,8 +5,8 @@ import ( ...@@ -5,8 +5,8 @@ import (
"encoding/json" "encoding/json"
"io" "io"
"net/http" "net/http"
"one-api/common"
"one-api/dto" "one-api/dto"
"one-api/logger"
relaycommon "one-api/relay/common" relaycommon "one-api/relay/common"
"one-api/relay/helper" "one-api/relay/helper"
"one-api/service" "one-api/service"
...@@ -51,7 +51,7 @@ func cfStreamHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Res ...@@ -51,7 +51,7 @@ func cfStreamHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Res
var response dto.ChatCompletionsStreamResponse var response dto.ChatCompletionsStreamResponse
err := json.Unmarshal([]byte(data), &response) err := json.Unmarshal([]byte(data), &response)
if err != nil { if err != nil {
common.LogError(c, "error_unmarshalling_stream_response: "+err.Error()) logger.LogError(c, "error_unmarshalling_stream_response: "+err.Error())
continue continue
} }
for _, choice := range response.Choices { for _, choice := range response.Choices {
...@@ -66,24 +66,24 @@ func cfStreamHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Res ...@@ -66,24 +66,24 @@ func cfStreamHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Res
info.FirstResponseTime = time.Now() info.FirstResponseTime = time.Now()
} }
if err != nil { if err != nil {
common.LogError(c, "error_rendering_stream_response: "+err.Error()) logger.LogError(c, "error_rendering_stream_response: "+err.Error())
} }
} }
if err := scanner.Err(); err != nil { if err := scanner.Err(); err != nil {
common.LogError(c, "error_scanning_stream_response: "+err.Error()) logger.LogError(c, "error_scanning_stream_response: "+err.Error())
} }
usage := service.ResponseText2Usage(responseText, info.UpstreamModelName, info.PromptTokens) usage := service.ResponseText2Usage(responseText, info.UpstreamModelName, info.PromptTokens)
if info.ShouldIncludeUsage { if info.ShouldIncludeUsage {
response := helper.GenerateFinalUsageResponse(id, info.StartTime.Unix(), info.UpstreamModelName, *usage) response := helper.GenerateFinalUsageResponse(id, info.StartTime.Unix(), info.UpstreamModelName, *usage)
err := helper.ObjectData(c, response) err := helper.ObjectData(c, response)
if err != nil { if err != nil {
common.LogError(c, "error_rendering_final_usage_response: "+err.Error()) logger.LogError(c, "error_rendering_final_usage_response: "+err.Error())
} }
} }
helper.Done(c) helper.Done(c)
common.CloseResponseBodyGracefully(resp) service.CloseResponseBodyGracefully(resp)
return nil, usage return nil, usage
} }
...@@ -93,7 +93,7 @@ func cfHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Response) ...@@ -93,7 +93,7 @@ func cfHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Response)
if err != nil { if err != nil {
return types.NewError(err, types.ErrorCodeBadResponseBody), nil return types.NewError(err, types.ErrorCodeBadResponseBody), nil
} }
common.CloseResponseBodyGracefully(resp) service.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 {
...@@ -123,7 +123,7 @@ func cfSTTHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Respon ...@@ -123,7 +123,7 @@ func cfSTTHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Respon
if err != nil { if err != nil {
return types.NewError(err, types.ErrorCodeBadResponseBody), nil return types.NewError(err, types.ErrorCodeBadResponseBody), nil
} }
common.CloseResponseBodyGracefully(resp) service.CloseResponseBodyGracefully(resp)
err = json.Unmarshal(responseBody, &cfResp) err = json.Unmarshal(responseBody, &cfResp)
if err != nil { if err != nil {
return types.NewError(err, types.ErrorCodeBadResponseBody), nil return types.NewError(err, types.ErrorCodeBadResponseBody), nil
......
...@@ -7,6 +7,7 @@ import ( ...@@ -7,6 +7,7 @@ import (
"net/http" "net/http"
"one-api/common" "one-api/common"
"one-api/dto" "one-api/dto"
"one-api/logger"
relaycommon "one-api/relay/common" relaycommon "one-api/relay/common"
"one-api/relay/helper" "one-api/relay/helper"
"one-api/service" "one-api/service"
...@@ -118,7 +119,7 @@ func cohereStreamHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http ...@@ -118,7 +119,7 @@ func cohereStreamHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http
var cohereResp CohereResponse var cohereResp CohereResponse
err := json.Unmarshal([]byte(data), &cohereResp) err := json.Unmarshal([]byte(data), &cohereResp)
if err != nil { if err != nil {
common.SysError("error unmarshalling stream response: " + err.Error()) logger.SysError("error unmarshalling stream response: " + err.Error())
return true return true
} }
var openaiResp dto.ChatCompletionsStreamResponse var openaiResp dto.ChatCompletionsStreamResponse
...@@ -153,7 +154,7 @@ func cohereStreamHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http ...@@ -153,7 +154,7 @@ func cohereStreamHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http
} }
jsonStr, err := json.Marshal(openaiResp) jsonStr, err := json.Marshal(openaiResp)
if err != nil { if err != nil {
common.SysError("error marshalling stream response: " + err.Error()) logger.SysError("error marshalling stream response: " + err.Error())
return true return true
} }
c.Render(-1, common.CustomEvent{Data: "data: " + string(jsonStr)}) c.Render(-1, common.CustomEvent{Data: "data: " + string(jsonStr)})
...@@ -175,7 +176,7 @@ func cohereHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Respo ...@@ -175,7 +176,7 @@ func cohereHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Respo
if err != nil { if err != nil {
return nil, types.NewError(err, types.ErrorCodeBadResponseBody) return nil, types.NewError(err, types.ErrorCodeBadResponseBody)
} }
common.CloseResponseBodyGracefully(resp) service.CloseResponseBodyGracefully(resp)
var cohereResp CohereResponseResult var cohereResp CohereResponseResult
err = json.Unmarshal(responseBody, &cohereResp) err = json.Unmarshal(responseBody, &cohereResp)
if err != nil { if err != nil {
...@@ -216,7 +217,7 @@ func cohereRerankHandler(c *gin.Context, resp *http.Response, info *relaycommon. ...@@ -216,7 +217,7 @@ func cohereRerankHandler(c *gin.Context, resp *http.Response, info *relaycommon.
if err != nil { if err != nil {
return nil, types.NewError(err, types.ErrorCodeBadResponseBody) return nil, types.NewError(err, types.ErrorCodeBadResponseBody)
} }
common.CloseResponseBodyGracefully(resp) service.CloseResponseBodyGracefully(resp)
var cohereResp CohereRerankResponseResult var cohereResp CohereRerankResponseResult
err = json.Unmarshal(responseBody, &cohereResp) err = json.Unmarshal(responseBody, &cohereResp)
if err != nil { if err != nil {
......
...@@ -9,6 +9,7 @@ import ( ...@@ -9,6 +9,7 @@ import (
"net/http" "net/http"
"one-api/common" "one-api/common"
"one-api/dto" "one-api/dto"
"one-api/logger"
relaycommon "one-api/relay/common" relaycommon "one-api/relay/common"
"one-api/relay/helper" "one-api/relay/helper"
"one-api/service" "one-api/service"
...@@ -49,7 +50,7 @@ func cozeChatHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Res ...@@ -49,7 +50,7 @@ func cozeChatHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Res
if err != nil { if err != nil {
return nil, types.NewError(err, types.ErrorCodeBadResponseBody) return nil, types.NewError(err, types.ErrorCodeBadResponseBody)
} }
common.CloseResponseBodyGracefully(resp) service.CloseResponseBodyGracefully(resp)
// convert coze response to openai response // convert coze response to openai response
var response dto.TextResponse var response dto.TextResponse
var cozeResponse CozeChatDetailResponse var cozeResponse CozeChatDetailResponse
...@@ -154,7 +155,7 @@ func handleCozeEvent(c *gin.Context, event string, data string, responseText *st ...@@ -154,7 +155,7 @@ func handleCozeEvent(c *gin.Context, event string, data string, responseText *st
var chatData CozeChatResponseData var chatData CozeChatResponseData
err := json.Unmarshal([]byte(data), &chatData) err := json.Unmarshal([]byte(data), &chatData)
if err != nil { if err != nil {
common.SysError("error_unmarshalling_stream_response: " + err.Error()) logger.SysError("error_unmarshalling_stream_response: " + err.Error())
return return
} }
...@@ -171,14 +172,14 @@ func handleCozeEvent(c *gin.Context, event string, data string, responseText *st ...@@ -171,14 +172,14 @@ func handleCozeEvent(c *gin.Context, event string, data string, responseText *st
var messageData CozeChatV3MessageDetail var messageData CozeChatV3MessageDetail
err := json.Unmarshal([]byte(data), &messageData) err := json.Unmarshal([]byte(data), &messageData)
if err != nil { if err != nil {
common.SysError("error_unmarshalling_stream_response: " + err.Error()) logger.SysError("error_unmarshalling_stream_response: " + err.Error())
return return
} }
var content string var content string
err = json.Unmarshal(messageData.Content, &content) err = json.Unmarshal(messageData.Content, &content)
if err != nil { if err != nil {
common.SysError("error_unmarshalling_stream_response: " + err.Error()) logger.SysError("error_unmarshalling_stream_response: " + err.Error())
return return
} }
...@@ -203,11 +204,11 @@ func handleCozeEvent(c *gin.Context, event string, data string, responseText *st ...@@ -203,11 +204,11 @@ func handleCozeEvent(c *gin.Context, event string, data string, responseText *st
var errorData CozeError var errorData CozeError
err := json.Unmarshal([]byte(data), &errorData) err := json.Unmarshal([]byte(data), &errorData)
if err != nil { if err != nil {
common.SysError("error_unmarshalling_stream_response: " + err.Error()) logger.SysError("error_unmarshalling_stream_response: " + err.Error())
return return
} }
common.SysError(fmt.Sprintf("stream event error: ", errorData.Code, errorData.Message)) logger.SysError(fmt.Sprintf("stream event error: ", errorData.Code, errorData.Message))
} }
} }
......
...@@ -11,6 +11,7 @@ import ( ...@@ -11,6 +11,7 @@ import (
"one-api/common" "one-api/common"
"one-api/constant" "one-api/constant"
"one-api/dto" "one-api/dto"
"one-api/logger"
relaycommon "one-api/relay/common" relaycommon "one-api/relay/common"
"one-api/relay/helper" "one-api/relay/helper"
"one-api/service" "one-api/service"
...@@ -36,14 +37,14 @@ func uploadDifyFile(c *gin.Context, info *relaycommon.RelayInfo, user string, me ...@@ -36,14 +37,14 @@ func uploadDifyFile(c *gin.Context, info *relaycommon.RelayInfo, user string, me
// Decode base64 string // Decode base64 string
decodedData, err := base64.StdEncoding.DecodeString(base64Data) decodedData, err := base64.StdEncoding.DecodeString(base64Data)
if err != nil { if err != nil {
common.SysError("failed to decode base64: " + err.Error()) logger.SysError("failed to decode base64: " + err.Error())
return nil return nil
} }
// Create temporary file // Create temporary file
tempFile, err := os.CreateTemp("", "dify-upload-*") tempFile, err := os.CreateTemp("", "dify-upload-*")
if err != nil { if err != nil {
common.SysError("failed to create temp file: " + err.Error()) logger.SysError("failed to create temp file: " + err.Error())
return nil return nil
} }
defer tempFile.Close() defer tempFile.Close()
...@@ -51,7 +52,7 @@ func uploadDifyFile(c *gin.Context, info *relaycommon.RelayInfo, user string, me ...@@ -51,7 +52,7 @@ func uploadDifyFile(c *gin.Context, info *relaycommon.RelayInfo, user string, me
// Write decoded data to temp file // Write decoded data to temp file
if _, err := tempFile.Write(decodedData); err != nil { if _, err := tempFile.Write(decodedData); err != nil {
common.SysError("failed to write to temp file: " + err.Error()) logger.SysError("failed to write to temp file: " + err.Error())
return nil return nil
} }
...@@ -61,7 +62,7 @@ func uploadDifyFile(c *gin.Context, info *relaycommon.RelayInfo, user string, me ...@@ -61,7 +62,7 @@ func uploadDifyFile(c *gin.Context, info *relaycommon.RelayInfo, user string, me
// Add user field // Add user field
if err := writer.WriteField("user", user); err != nil { if err := writer.WriteField("user", user); err != nil {
common.SysError("failed to add user field: " + err.Error()) logger.SysError("failed to add user field: " + err.Error())
return nil return nil
} }
...@@ -74,13 +75,13 @@ func uploadDifyFile(c *gin.Context, info *relaycommon.RelayInfo, user string, me ...@@ -74,13 +75,13 @@ func uploadDifyFile(c *gin.Context, info *relaycommon.RelayInfo, user string, me
// Create form file // Create form file
part, err := writer.CreateFormFile("file", fmt.Sprintf("image.%s", strings.TrimPrefix(mimeType, "image/"))) part, err := writer.CreateFormFile("file", fmt.Sprintf("image.%s", strings.TrimPrefix(mimeType, "image/")))
if err != nil { if err != nil {
common.SysError("failed to create form file: " + err.Error()) logger.SysError("failed to create form file: " + err.Error())
return nil return nil
} }
// Copy file content to form // Copy file content to form
if _, err = io.Copy(part, bytes.NewReader(decodedData)); err != nil { if _, err = io.Copy(part, bytes.NewReader(decodedData)); err != nil {
common.SysError("failed to copy file content: " + err.Error()) logger.SysError("failed to copy file content: " + err.Error())
return nil return nil
} }
writer.Close() writer.Close()
...@@ -88,7 +89,7 @@ func uploadDifyFile(c *gin.Context, info *relaycommon.RelayInfo, user string, me ...@@ -88,7 +89,7 @@ func uploadDifyFile(c *gin.Context, info *relaycommon.RelayInfo, user string, me
// Create HTTP request // Create HTTP request
req, err := http.NewRequest("POST", uploadUrl, body) req, err := http.NewRequest("POST", uploadUrl, body)
if err != nil { if err != nil {
common.SysError("failed to create request: " + err.Error()) logger.SysError("failed to create request: " + err.Error())
return nil return nil
} }
...@@ -99,7 +100,7 @@ func uploadDifyFile(c *gin.Context, info *relaycommon.RelayInfo, user string, me ...@@ -99,7 +100,7 @@ func uploadDifyFile(c *gin.Context, info *relaycommon.RelayInfo, user string, me
client := service.GetHttpClient() client := service.GetHttpClient()
resp, err := client.Do(req) resp, err := client.Do(req)
if err != nil { if err != nil {
common.SysError("failed to send request: " + err.Error()) logger.SysError("failed to send request: " + err.Error())
return nil return nil
} }
defer resp.Body.Close() defer resp.Body.Close()
...@@ -109,7 +110,7 @@ func uploadDifyFile(c *gin.Context, info *relaycommon.RelayInfo, user string, me ...@@ -109,7 +110,7 @@ func uploadDifyFile(c *gin.Context, info *relaycommon.RelayInfo, user string, me
Id string `json:"id"` Id string `json:"id"`
} }
if err := json.NewDecoder(resp.Body).Decode(&result); err != nil { if err := json.NewDecoder(resp.Body).Decode(&result); err != nil {
common.SysError("failed to decode response: " + err.Error()) logger.SysError("failed to decode response: " + err.Error())
return nil return nil
} }
...@@ -219,7 +220,7 @@ func difyStreamHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.R ...@@ -219,7 +220,7 @@ func difyStreamHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.R
var difyResponse DifyChunkChatCompletionResponse var difyResponse DifyChunkChatCompletionResponse
err := json.Unmarshal([]byte(data), &difyResponse) err := json.Unmarshal([]byte(data), &difyResponse)
if err != nil { if err != nil {
common.SysError("error unmarshalling stream response: " + err.Error()) logger.SysError("error unmarshalling stream response: " + err.Error())
return true return true
} }
var openaiResponse dto.ChatCompletionsStreamResponse var openaiResponse dto.ChatCompletionsStreamResponse
...@@ -239,7 +240,7 @@ func difyStreamHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.R ...@@ -239,7 +240,7 @@ func difyStreamHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.R
} }
err = helper.ObjectData(c, openaiResponse) err = helper.ObjectData(c, openaiResponse)
if err != nil { if err != nil {
common.SysError(err.Error()) logger.SysError(err.Error())
} }
return true return true
}) })
...@@ -258,7 +259,7 @@ func difyHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Respons ...@@ -258,7 +259,7 @@ func difyHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Respons
if err != nil { if err != nil {
return nil, types.NewError(err, types.ErrorCodeBadResponseBody) return nil, types.NewError(err, types.ErrorCodeBadResponseBody)
} }
common.CloseResponseBodyGracefully(resp) service.CloseResponseBodyGracefully(resp)
err = json.Unmarshal(responseBody, &difyResponse) err = json.Unmarshal(responseBody, &difyResponse)
if err != nil { if err != nil {
return nil, types.NewError(err, types.ErrorCodeBadResponseBody) return nil, types.NewError(err, types.ErrorCodeBadResponseBody)
......
...@@ -78,7 +78,7 @@ func (a *Adaptor) ConvertImageRequest(c *gin.Context, info *relaycommon.RelayInf ...@@ -78,7 +78,7 @@ func (a *Adaptor) ConvertImageRequest(c *gin.Context, info *relaycommon.RelayInf
}, },
}, },
Parameters: dto.GeminiImageParameters{ Parameters: dto.GeminiImageParameters{
SampleCount: request.N, SampleCount: int(request.N),
AspectRatio: aspectRatio, AspectRatio: aspectRatio,
PersonGeneration: "allow_adult", // default allow adult PersonGeneration: "allow_adult", // default allow adult
}, },
......
...@@ -5,6 +5,7 @@ import ( ...@@ -5,6 +5,7 @@ import (
"net/http" "net/http"
"one-api/common" "one-api/common"
"one-api/dto" "one-api/dto"
"one-api/logger"
relaycommon "one-api/relay/common" relaycommon "one-api/relay/common"
"one-api/relay/helper" "one-api/relay/helper"
"one-api/service" "one-api/service"
...@@ -17,7 +18,7 @@ import ( ...@@ -17,7 +18,7 @@ import (
) )
func GeminiTextGenerationHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Response) (*dto.Usage, *types.NewAPIError) { func GeminiTextGenerationHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Response) (*dto.Usage, *types.NewAPIError) {
defer common.CloseResponseBodyGracefully(resp) defer service.CloseResponseBodyGracefully(resp)
// 读取响应体 // 读取响应体
responseBody, err := io.ReadAll(resp.Body) responseBody, err := io.ReadAll(resp.Body)
...@@ -53,13 +54,13 @@ func GeminiTextGenerationHandler(c *gin.Context, info *relaycommon.RelayInfo, re ...@@ -53,13 +54,13 @@ func GeminiTextGenerationHandler(c *gin.Context, info *relaycommon.RelayInfo, re
} }
} }
common.IOCopyBytesGracefully(c, resp, responseBody) service.IOCopyBytesGracefully(c, resp, responseBody)
return &usage, nil return &usage, nil
} }
func NativeGeminiEmbeddingHandler(c *gin.Context, resp *http.Response, info *relaycommon.RelayInfo) (*dto.Usage, *types.NewAPIError) { func NativeGeminiEmbeddingHandler(c *gin.Context, resp *http.Response, info *relaycommon.RelayInfo) (*dto.Usage, *types.NewAPIError) {
defer common.CloseResponseBodyGracefully(resp) defer service.CloseResponseBodyGracefully(resp)
responseBody, err := io.ReadAll(resp.Body) responseBody, err := io.ReadAll(resp.Body)
if err != nil { if err != nil {
...@@ -89,7 +90,7 @@ func NativeGeminiEmbeddingHandler(c *gin.Context, resp *http.Response, info *rel ...@@ -89,7 +90,7 @@ func NativeGeminiEmbeddingHandler(c *gin.Context, resp *http.Response, info *rel
} }
} }
common.IOCopyBytesGracefully(c, resp, responseBody) service.IOCopyBytesGracefully(c, resp, responseBody)
return usage, nil return usage, nil
} }
...@@ -106,7 +107,7 @@ func GeminiTextGenerationStreamHandler(c *gin.Context, info *relaycommon.RelayIn ...@@ -106,7 +107,7 @@ func GeminiTextGenerationStreamHandler(c *gin.Context, info *relaycommon.RelayIn
var geminiResponse dto.GeminiChatResponse var geminiResponse dto.GeminiChatResponse
err := common.UnmarshalJsonStr(data, &geminiResponse) err := common.UnmarshalJsonStr(data, &geminiResponse)
if err != nil { if err != nil {
common.LogError(c, "error unmarshalling stream response: "+err.Error()) logger.LogError(c, "error unmarshalling stream response: "+err.Error())
return false return false
} }
...@@ -140,7 +141,7 @@ func GeminiTextGenerationStreamHandler(c *gin.Context, info *relaycommon.RelayIn ...@@ -140,7 +141,7 @@ func GeminiTextGenerationStreamHandler(c *gin.Context, info *relaycommon.RelayIn
// 直接发送 GeminiChatResponse 响应 // 直接发送 GeminiChatResponse 响应
err = helper.StringData(c, data) err = helper.StringData(c, data)
if err != nil { if err != nil {
common.LogError(c, err.Error()) logger.LogError(c, err.Error())
} }
info.SendResponseCount++ info.SendResponseCount++
return true return true
......
...@@ -9,6 +9,7 @@ import ( ...@@ -9,6 +9,7 @@ import (
"one-api/common" "one-api/common"
"one-api/constant" "one-api/constant"
"one-api/dto" "one-api/dto"
"one-api/logger"
"one-api/relay/channel/openai" "one-api/relay/channel/openai"
relaycommon "one-api/relay/common" relaycommon "one-api/relay/common"
"one-api/relay/helper" "one-api/relay/helper"
...@@ -901,7 +902,7 @@ func GeminiChatStreamHandler(c *gin.Context, info *relaycommon.RelayInfo, resp * ...@@ -901,7 +902,7 @@ func GeminiChatStreamHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *
var geminiResponse dto.GeminiChatResponse var geminiResponse dto.GeminiChatResponse
err := common.UnmarshalJsonStr(data, &geminiResponse) err := common.UnmarshalJsonStr(data, &geminiResponse)
if err != nil { if err != nil {
common.LogError(c, "error unmarshalling stream response: "+err.Error()) logger.LogError(c, "error unmarshalling stream response: "+err.Error())
return false return false
} }
...@@ -945,7 +946,7 @@ func GeminiChatStreamHandler(c *gin.Context, info *relaycommon.RelayInfo, resp * ...@@ -945,7 +946,7 @@ func GeminiChatStreamHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *
finishReason = constant.FinishReasonToolCalls finishReason = constant.FinishReasonToolCalls
err = handleStream(c, info, emptyResponse) err = handleStream(c, info, emptyResponse)
if err != nil { if err != nil {
common.LogError(c, err.Error()) logger.LogError(c, err.Error())
} }
response.ClearToolCalls() response.ClearToolCalls()
...@@ -957,7 +958,7 @@ func GeminiChatStreamHandler(c *gin.Context, info *relaycommon.RelayInfo, resp * ...@@ -957,7 +958,7 @@ func GeminiChatStreamHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *
err = handleStream(c, info, response) err = handleStream(c, info, response)
if err != nil { if err != nil {
common.LogError(c, err.Error()) logger.LogError(c, err.Error())
} }
if isStop { if isStop {
_ = handleStream(c, info, helper.GenerateStopResponse(id, createAt, info.UpstreamModelName, finishReason)) _ = handleStream(c, info, helper.GenerateStopResponse(id, createAt, info.UpstreamModelName, finishReason))
...@@ -993,7 +994,7 @@ func GeminiChatStreamHandler(c *gin.Context, info *relaycommon.RelayInfo, resp * ...@@ -993,7 +994,7 @@ func GeminiChatStreamHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *
response := helper.GenerateFinalUsageResponse(id, createAt, info.UpstreamModelName, *usage) response := helper.GenerateFinalUsageResponse(id, createAt, info.UpstreamModelName, *usage)
err := handleFinalStream(c, info, response) err := handleFinalStream(c, info, response)
if err != nil { if err != nil {
common.SysError("send final response failed: " + err.Error()) logger.SysError("send final response failed: " + err.Error())
} }
//if info.RelayFormat == relaycommon.RelayFormatOpenAI { //if info.RelayFormat == relaycommon.RelayFormatOpenAI {
// helper.Done(c) // helper.Done(c)
...@@ -1007,7 +1008,7 @@ func GeminiChatHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.R ...@@ -1007,7 +1008,7 @@ func GeminiChatHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.R
if err != nil { if err != nil {
return nil, types.NewOpenAIError(err, types.ErrorCodeBadResponseBody, http.StatusInternalServerError) return nil, types.NewOpenAIError(err, types.ErrorCodeBadResponseBody, http.StatusInternalServerError)
} }
common.CloseResponseBodyGracefully(resp) service.CloseResponseBodyGracefully(resp)
if common.DebugEnabled { if common.DebugEnabled {
println(string(responseBody)) println(string(responseBody))
} }
...@@ -1057,13 +1058,13 @@ func GeminiChatHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.R ...@@ -1057,13 +1058,13 @@ func GeminiChatHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.R
break break
} }
common.IOCopyBytesGracefully(c, resp, responseBody) service.IOCopyBytesGracefully(c, resp, responseBody)
return &usage, nil return &usage, nil
} }
func GeminiEmbeddingHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Response) (*dto.Usage, *types.NewAPIError) { func GeminiEmbeddingHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Response) (*dto.Usage, *types.NewAPIError) {
defer common.CloseResponseBodyGracefully(resp) defer service.CloseResponseBodyGracefully(resp)
responseBody, readErr := io.ReadAll(resp.Body) responseBody, readErr := io.ReadAll(resp.Body)
if readErr != nil { if readErr != nil {
...@@ -1107,7 +1108,7 @@ func GeminiEmbeddingHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *h ...@@ -1107,7 +1108,7 @@ func GeminiEmbeddingHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *h
return nil, types.NewOpenAIError(jsonErr, types.ErrorCodeBadResponseBody, http.StatusInternalServerError) return nil, types.NewOpenAIError(jsonErr, types.ErrorCodeBadResponseBody, http.StatusInternalServerError)
} }
common.IOCopyBytesGracefully(c, resp, jsonResponse) service.IOCopyBytesGracefully(c, resp, jsonResponse)
return usage, nil return usage, nil
} }
......
...@@ -5,9 +5,9 @@ import ( ...@@ -5,9 +5,9 @@ import (
"fmt" "fmt"
"io" "io"
"net/http" "net/http"
"one-api/common"
"one-api/dto" "one-api/dto"
relaycommon "one-api/relay/common" relaycommon "one-api/relay/common"
"one-api/service"
"one-api/types" "one-api/types"
"github.com/gin-gonic/gin" "github.com/gin-gonic/gin"
...@@ -54,7 +54,7 @@ func jimengImageHandler(c *gin.Context, resp *http.Response, info *relaycommon.R ...@@ -54,7 +54,7 @@ func jimengImageHandler(c *gin.Context, resp *http.Response, info *relaycommon.R
if err != nil { if err != nil {
return nil, types.NewOpenAIError(err, types.ErrorCodeReadResponseBodyFailed, http.StatusInternalServerError) return nil, types.NewOpenAIError(err, types.ErrorCodeReadResponseBodyFailed, http.StatusInternalServerError)
} }
common.CloseResponseBodyGracefully(resp) service.CloseResponseBodyGracefully(resp)
err = json.Unmarshal(responseBody, &jimengResponse) err = json.Unmarshal(responseBody, &jimengResponse)
if err != nil { if err != nil {
......
...@@ -12,7 +12,7 @@ import ( ...@@ -12,7 +12,7 @@ import (
"io" "io"
"net/http" "net/http"
"net/url" "net/url"
"one-api/common" "one-api/logger"
"sort" "sort"
"strings" "strings"
"time" "time"
...@@ -44,7 +44,7 @@ func SetPayloadHash(c *gin.Context, req any) error { ...@@ -44,7 +44,7 @@ func SetPayloadHash(c *gin.Context, req any) error {
if err != nil { if err != nil {
return err return err
} }
common.LogInfo(c, fmt.Sprintf("SetPayloadHash body: %s", body)) logger.LogInfo(c, fmt.Sprintf("SetPayloadHash body: %s", body))
payloadHash := sha256.Sum256(body) payloadHash := sha256.Sum256(body)
hexPayloadHash := hex.EncodeToString(payloadHash[:]) hexPayloadHash := hex.EncodeToString(payloadHash[:])
c.Set(HexPayloadHashKey, hexPayloadHash) c.Set(HexPayloadHashKey, hexPayloadHash)
......
...@@ -7,6 +7,7 @@ import ( ...@@ -7,6 +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" "one-api/types"
"github.com/gin-gonic/gin" "github.com/gin-gonic/gin"
...@@ -56,7 +57,7 @@ func mokaEmbeddingHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *htt ...@@ -56,7 +57,7 @@ func mokaEmbeddingHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *htt
if err != nil { if err != nil {
return nil, types.NewError(err, types.ErrorCodeBadResponseBody) return nil, types.NewError(err, types.ErrorCodeBadResponseBody)
} }
common.CloseResponseBodyGracefully(resp) service.CloseResponseBodyGracefully(resp)
err = json.Unmarshal(responseBody, &baiduResponse) err = json.Unmarshal(responseBody, &baiduResponse)
if err != nil { if err != nil {
return nil, types.NewError(err, types.ErrorCodeBadResponseBody) return nil, types.NewError(err, types.ErrorCodeBadResponseBody)
...@@ -77,6 +78,6 @@ func mokaEmbeddingHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *htt ...@@ -77,6 +78,6 @@ func mokaEmbeddingHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *htt
} }
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)
common.IOCopyBytesGracefully(c, resp, jsonResponse) service.IOCopyBytesGracefully(c, resp, jsonResponse)
return &fullTextResponse.Usage, nil return &fullTextResponse.Usage, nil
} }
...@@ -94,7 +94,7 @@ func ollamaEmbeddingHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *h ...@@ -94,7 +94,7 @@ func ollamaEmbeddingHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *h
if err != nil { if err != nil {
return nil, types.NewOpenAIError(err, types.ErrorCodeBadResponseBody, http.StatusInternalServerError) return nil, types.NewOpenAIError(err, types.ErrorCodeBadResponseBody, http.StatusInternalServerError)
} }
common.CloseResponseBodyGracefully(resp) service.CloseResponseBodyGracefully(resp)
err = common.Unmarshal(responseBody, &ollamaEmbeddingResponse) err = common.Unmarshal(responseBody, &ollamaEmbeddingResponse)
if err != nil { if err != nil {
return nil, types.NewOpenAIError(err, types.ErrorCodeBadResponseBody, http.StatusInternalServerError) return nil, types.NewOpenAIError(err, types.ErrorCodeBadResponseBody, http.StatusInternalServerError)
...@@ -123,7 +123,7 @@ func ollamaEmbeddingHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *h ...@@ -123,7 +123,7 @@ func ollamaEmbeddingHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *h
if err != nil { if err != nil {
return nil, types.NewOpenAIError(err, types.ErrorCodeBadResponseBody, http.StatusInternalServerError) return nil, types.NewOpenAIError(err, types.ErrorCodeBadResponseBody, http.StatusInternalServerError)
} }
common.IOCopyBytesGracefully(c, resp, doResponseBody) service.IOCopyBytesGracefully(c, resp, doResponseBody)
return usage, nil return usage, nil
} }
......
...@@ -7,6 +7,7 @@ import ( ...@@ -7,6 +7,7 @@ import (
"net/http" "net/http"
"one-api/common" "one-api/common"
"one-api/dto" "one-api/dto"
"one-api/logger"
relaycommon "one-api/relay/common" relaycommon "one-api/relay/common"
relayconstant "one-api/relay/constant" relayconstant "one-api/relay/constant"
"one-api/relay/helper" "one-api/relay/helper"
...@@ -50,7 +51,7 @@ func handleClaudeFormat(c *gin.Context, data string, info *relaycommon.RelayInfo ...@@ -50,7 +51,7 @@ func handleClaudeFormat(c *gin.Context, data string, info *relaycommon.RelayInfo
func handleGeminiFormat(c *gin.Context, data string, info *relaycommon.RelayInfo) error { func handleGeminiFormat(c *gin.Context, data string, info *relaycommon.RelayInfo) error {
var streamResponse dto.ChatCompletionsStreamResponse var streamResponse dto.ChatCompletionsStreamResponse
if err := common.Unmarshal(common.StringToByteSlice(data), &streamResponse); err != nil { if err := common.Unmarshal(common.StringToByteSlice(data), &streamResponse); err != nil {
common.LogError(c, "failed to unmarshal stream response: "+err.Error()) logger.LogError(c, "failed to unmarshal stream response: "+err.Error())
return err return err
} }
...@@ -63,7 +64,7 @@ func handleGeminiFormat(c *gin.Context, data string, info *relaycommon.RelayInfo ...@@ -63,7 +64,7 @@ func handleGeminiFormat(c *gin.Context, data string, info *relaycommon.RelayInfo
geminiResponseStr, err := common.Marshal(geminiResponse) geminiResponseStr, err := common.Marshal(geminiResponse)
if err != nil { if err != nil {
common.LogError(c, "failed to marshal gemini response: "+err.Error()) logger.LogError(c, "failed to marshal gemini response: "+err.Error())
return err return err
} }
...@@ -110,14 +111,14 @@ func processChatCompletions(streamResp string, streamItems []string, responseTex ...@@ -110,14 +111,14 @@ func processChatCompletions(streamResp string, streamItems []string, responseTex
var streamResponses []dto.ChatCompletionsStreamResponse var streamResponses []dto.ChatCompletionsStreamResponse
if err := json.Unmarshal(common.StringToByteSlice(streamResp), &streamResponses); err != nil { if err := json.Unmarshal(common.StringToByteSlice(streamResp), &streamResponses); err != nil {
// 一次性解析失败,逐个解析 // 一次性解析失败,逐个解析
common.SysError("error unmarshalling stream response: " + err.Error()) logger.SysError("error unmarshalling stream response: " + err.Error())
for _, item := range streamItems { for _, item := range streamItems {
var streamResponse dto.ChatCompletionsStreamResponse var streamResponse dto.ChatCompletionsStreamResponse
if err := json.Unmarshal(common.StringToByteSlice(item), &streamResponse); err != nil { if err := json.Unmarshal(common.StringToByteSlice(item), &streamResponse); err != nil {
return err return err
} }
if err := ProcessStreamResponse(streamResponse, responseTextBuilder, toolCount); err != nil { if err := ProcessStreamResponse(streamResponse, responseTextBuilder, toolCount); err != nil {
common.SysError("error processing stream response: " + err.Error()) logger.SysError("error processing stream response: " + err.Error())
} }
} }
return nil return nil
...@@ -146,7 +147,7 @@ func processCompletions(streamResp string, streamItems []string, responseTextBui ...@@ -146,7 +147,7 @@ func processCompletions(streamResp string, streamItems []string, responseTextBui
var streamResponses []dto.CompletionsStreamResponse var streamResponses []dto.CompletionsStreamResponse
if err := json.Unmarshal(common.StringToByteSlice(streamResp), &streamResponses); err != nil { if err := json.Unmarshal(common.StringToByteSlice(streamResp), &streamResponses); err != nil {
// 一次性解析失败,逐个解析 // 一次性解析失败,逐个解析
common.SysError("error unmarshalling stream response: " + err.Error()) logger.SysError("error unmarshalling stream response: " + err.Error())
for _, item := range streamItems { for _, item := range streamItems {
var streamResponse dto.CompletionsStreamResponse var streamResponse dto.CompletionsStreamResponse
if err := json.Unmarshal(common.StringToByteSlice(item), &streamResponse); err != nil { if err := json.Unmarshal(common.StringToByteSlice(item), &streamResponse); err != nil {
...@@ -213,7 +214,7 @@ func HandleFinalResponse(c *gin.Context, info *relaycommon.RelayInfo, lastStream ...@@ -213,7 +214,7 @@ func HandleFinalResponse(c *gin.Context, info *relaycommon.RelayInfo, lastStream
info.ClaudeConvertInfo.Done = true info.ClaudeConvertInfo.Done = true
var streamResponse dto.ChatCompletionsStreamResponse var streamResponse dto.ChatCompletionsStreamResponse
if err := common.Unmarshal(common.StringToByteSlice(lastStreamData), &streamResponse); err != nil { if err := common.Unmarshal(common.StringToByteSlice(lastStreamData), &streamResponse); err != nil {
common.SysError("error unmarshalling stream response: " + err.Error()) logger.SysError("error unmarshalling stream response: " + err.Error())
return return
} }
...@@ -227,7 +228,7 @@ func HandleFinalResponse(c *gin.Context, info *relaycommon.RelayInfo, lastStream ...@@ -227,7 +228,7 @@ func HandleFinalResponse(c *gin.Context, info *relaycommon.RelayInfo, lastStream
case relaycommon.RelayFormatGemini: case relaycommon.RelayFormatGemini:
var streamResponse dto.ChatCompletionsStreamResponse var streamResponse dto.ChatCompletionsStreamResponse
if err := common.Unmarshal(common.StringToByteSlice(lastStreamData), &streamResponse); err != nil { if err := common.Unmarshal(common.StringToByteSlice(lastStreamData), &streamResponse); err != nil {
common.SysError("error unmarshalling stream response: " + err.Error()) logger.SysError("error unmarshalling stream response: " + err.Error())
return return
} }
...@@ -245,7 +246,7 @@ func HandleFinalResponse(c *gin.Context, info *relaycommon.RelayInfo, lastStream ...@@ -245,7 +246,7 @@ func HandleFinalResponse(c *gin.Context, info *relaycommon.RelayInfo, lastStream
geminiResponseStr, err := common.Marshal(geminiResponse) geminiResponseStr, err := common.Marshal(geminiResponse)
if err != nil { if err != nil {
common.SysError("error marshalling gemini response: " + err.Error()) logger.SysError("error marshalling gemini response: " + err.Error())
return return
} }
......
...@@ -10,6 +10,7 @@ import ( ...@@ -10,6 +10,7 @@ import (
"one-api/common" "one-api/common"
"one-api/constant" "one-api/constant"
"one-api/dto" "one-api/dto"
"one-api/logger"
relaycommon "one-api/relay/common" relaycommon "one-api/relay/common"
"one-api/relay/helper" "one-api/relay/helper"
"one-api/service" "one-api/service"
...@@ -108,11 +109,11 @@ func sendStreamData(c *gin.Context, info *relaycommon.RelayInfo, data string, fo ...@@ -108,11 +109,11 @@ func sendStreamData(c *gin.Context, info *relaycommon.RelayInfo, data string, fo
func OaiStreamHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Response) (*dto.Usage, *types.NewAPIError) { 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") logger.LogError(c, "invalid response or response body")
return nil, types.NewOpenAIError(fmt.Errorf("invalid response"), types.ErrorCodeBadResponse, http.StatusInternalServerError) return nil, types.NewOpenAIError(fmt.Errorf("invalid response"), types.ErrorCodeBadResponse, http.StatusInternalServerError)
} }
defer common.CloseResponseBodyGracefully(resp) defer service.CloseResponseBodyGracefully(resp)
model := info.UpstreamModelName model := info.UpstreamModelName
var responseId string var responseId string
...@@ -129,7 +130,7 @@ func OaiStreamHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Re ...@@ -129,7 +130,7 @@ func OaiStreamHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Re
if lastStreamData != "" { if lastStreamData != "" {
err := HandleStreamFormat(c, info, lastStreamData, info.ChannelSetting.ForceFormat, info.ChannelSetting.ThinkingToContent) err := HandleStreamFormat(c, info, lastStreamData, info.ChannelSetting.ForceFormat, info.ChannelSetting.ThinkingToContent)
if err != nil { if err != nil {
common.SysError("error handling stream format: " + err.Error()) logger.SysError("error handling stream format: " + err.Error())
} }
} }
if len(data) > 0 { if len(data) > 0 {
...@@ -143,7 +144,7 @@ func OaiStreamHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Re ...@@ -143,7 +144,7 @@ func OaiStreamHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Re
shouldSendLastResp := true shouldSendLastResp := true
if err := handleLastResponse(lastStreamData, &responseId, &createAt, &systemFingerprint, &model, &usage, if err := handleLastResponse(lastStreamData, &responseId, &createAt, &systemFingerprint, &model, &usage,
&containStreamUsage, info, &shouldSendLastResp); err != nil { &containStreamUsage, info, &shouldSendLastResp); err != nil {
common.LogError(c, fmt.Sprintf("error handling last response: %s, lastStreamData: [%s]", err.Error(), lastStreamData)) logger.LogError(c, fmt.Sprintf("error handling last response: %s, lastStreamData: [%s]", err.Error(), lastStreamData))
} }
if info.RelayFormat == relaycommon.RelayFormatOpenAI { if info.RelayFormat == relaycommon.RelayFormatOpenAI {
...@@ -154,7 +155,7 @@ func OaiStreamHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Re ...@@ -154,7 +155,7 @@ func OaiStreamHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Re
// 处理token计算 // 处理token计算
if err := processTokens(info.RelayMode, streamItems, &responseTextBuilder, &toolCount); err != nil { if err := processTokens(info.RelayMode, streamItems, &responseTextBuilder, &toolCount); err != nil {
common.LogError(c, "error processing tokens: "+err.Error()) logger.LogError(c, "error processing tokens: "+err.Error())
} }
if !containStreamUsage { if !containStreamUsage {
...@@ -173,7 +174,7 @@ func OaiStreamHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Re ...@@ -173,7 +174,7 @@ func OaiStreamHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Re
} }
func OpenaiHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Response) (*dto.Usage, *types.NewAPIError) { func OpenaiHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Response) (*dto.Usage, *types.NewAPIError) {
defer common.CloseResponseBodyGracefully(resp) defer service.CloseResponseBodyGracefully(resp)
var simpleResponse dto.OpenAITextResponse var simpleResponse dto.OpenAITextResponse
responseBody, err := io.ReadAll(resp.Body) responseBody, err := io.ReadAll(resp.Body)
...@@ -235,7 +236,7 @@ func OpenaiHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Respo ...@@ -235,7 +236,7 @@ func OpenaiHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Respo
responseBody = geminiRespStr responseBody = geminiRespStr
} }
common.IOCopyBytesGracefully(c, resp, responseBody) service.IOCopyBytesGracefully(c, resp, responseBody)
return &simpleResponse.Usage, nil return &simpleResponse.Usage, nil
} }
...@@ -247,7 +248,7 @@ func OpenaiTTSHandler(c *gin.Context, resp *http.Response, info *relaycommon.Rel ...@@ -247,7 +248,7 @@ func OpenaiTTSHandler(c *gin.Context, resp *http.Response, info *relaycommon.Rel
// if the upstream returns a specific status code, once the upstream has already written the header, // if the upstream returns a specific status code, once the upstream has already written the header,
// the subsequent failure of the response body should be regarded as a non-recoverable error, // the subsequent failure of the response body should be regarded as a non-recoverable error,
// and can be terminated directly. // and can be terminated directly.
defer common.CloseResponseBodyGracefully(resp) defer service.CloseResponseBodyGracefully(resp)
usage := &dto.Usage{} usage := &dto.Usage{}
usage.PromptTokens = info.PromptTokens usage.PromptTokens = info.PromptTokens
usage.TotalTokens = info.PromptTokens usage.TotalTokens = info.PromptTokens
...@@ -258,13 +259,13 @@ func OpenaiTTSHandler(c *gin.Context, resp *http.Response, info *relaycommon.Rel ...@@ -258,13 +259,13 @@ func OpenaiTTSHandler(c *gin.Context, resp *http.Response, info *relaycommon.Rel
c.Writer.WriteHeaderNow() c.Writer.WriteHeaderNow()
_, err := io.Copy(c.Writer, resp.Body) _, err := io.Copy(c.Writer, resp.Body)
if err != nil { if err != nil {
common.LogError(c, err.Error()) logger.LogError(c, err.Error())
} }
return usage return usage
} }
func OpenaiSTTHandler(c *gin.Context, resp *http.Response, info *relaycommon.RelayInfo, responseFormat string) (*types.NewAPIError, *dto.Usage) { func OpenaiSTTHandler(c *gin.Context, resp *http.Response, info *relaycommon.RelayInfo, responseFormat string) (*types.NewAPIError, *dto.Usage) {
defer common.CloseResponseBodyGracefully(resp) defer service.CloseResponseBodyGracefully(resp)
// count tokens by audio file duration // count tokens by audio file duration
audioTokens, err := countAudioTokens(c) audioTokens, err := countAudioTokens(c)
...@@ -276,7 +277,7 @@ func OpenaiSTTHandler(c *gin.Context, resp *http.Response, info *relaycommon.Rel ...@@ -276,7 +277,7 @@ func OpenaiSTTHandler(c *gin.Context, resp *http.Response, info *relaycommon.Rel
return types.NewOpenAIError(err, types.ErrorCodeReadResponseBodyFailed, http.StatusInternalServerError), nil return types.NewOpenAIError(err, types.ErrorCodeReadResponseBodyFailed, http.StatusInternalServerError), nil
} }
// 写入新的 response body // 写入新的 response body
common.IOCopyBytesGracefully(c, resp, responseBody) service.IOCopyBytesGracefully(c, resp, responseBody)
usage := &dto.Usage{} usage := &dto.Usage{}
usage.PromptTokens = audioTokens usage.PromptTokens = audioTokens
...@@ -386,7 +387,7 @@ func OpenaiRealtimeHandler(c *gin.Context, info *relaycommon.RelayInfo) (*types. ...@@ -386,7 +387,7 @@ func OpenaiRealtimeHandler(c *gin.Context, info *relaycommon.RelayInfo) (*types.
errChan <- fmt.Errorf("error counting text token: %v", err) errChan <- fmt.Errorf("error counting text token: %v", err)
return return
} }
common.LogInfo(c, fmt.Sprintf("type: %s, textToken: %d, audioToken: %d", realtimeEvent.Type, textToken, audioToken)) logger.LogInfo(c, fmt.Sprintf("type: %s, textToken: %d, audioToken: %d", realtimeEvent.Type, textToken, audioToken))
localUsage.TotalTokens += textToken + audioToken localUsage.TotalTokens += textToken + audioToken
localUsage.InputTokens += textToken + audioToken localUsage.InputTokens += textToken + audioToken
localUsage.InputTokenDetails.TextTokens += textToken localUsage.InputTokenDetails.TextTokens += textToken
...@@ -459,7 +460,7 @@ func OpenaiRealtimeHandler(c *gin.Context, info *relaycommon.RelayInfo) (*types. ...@@ -459,7 +460,7 @@ func OpenaiRealtimeHandler(c *gin.Context, info *relaycommon.RelayInfo) (*types.
errChan <- fmt.Errorf("error counting text token: %v", err) errChan <- fmt.Errorf("error counting text token: %v", err)
return return
} }
common.LogInfo(c, fmt.Sprintf("type: %s, textToken: %d, audioToken: %d", realtimeEvent.Type, textToken, audioToken)) logger.LogInfo(c, fmt.Sprintf("type: %s, textToken: %d, audioToken: %d", realtimeEvent.Type, textToken, audioToken))
localUsage.TotalTokens += textToken + audioToken localUsage.TotalTokens += textToken + audioToken
info.IsFirstRequest = false info.IsFirstRequest = false
localUsage.InputTokens += textToken + audioToken localUsage.InputTokens += textToken + audioToken
...@@ -474,9 +475,9 @@ func OpenaiRealtimeHandler(c *gin.Context, info *relaycommon.RelayInfo) (*types. ...@@ -474,9 +475,9 @@ func OpenaiRealtimeHandler(c *gin.Context, info *relaycommon.RelayInfo) (*types.
localUsage = &dto.RealtimeUsage{} localUsage = &dto.RealtimeUsage{}
// print now usage // print now usage
} }
common.LogInfo(c, fmt.Sprintf("realtime streaming sumUsage: %v", sumUsage)) logger.LogInfo(c, fmt.Sprintf("realtime streaming sumUsage: %v", sumUsage))
common.LogInfo(c, fmt.Sprintf("realtime streaming localUsage: %v", localUsage)) logger.LogInfo(c, fmt.Sprintf("realtime streaming localUsage: %v", localUsage))
common.LogInfo(c, fmt.Sprintf("realtime streaming localUsage: %v", localUsage)) logger.LogInfo(c, fmt.Sprintf("realtime streaming localUsage: %v", localUsage))
} else if realtimeEvent.Type == dto.RealtimeEventTypeSessionUpdated || realtimeEvent.Type == dto.RealtimeEventTypeSessionCreated { } else if realtimeEvent.Type == dto.RealtimeEventTypeSessionUpdated || realtimeEvent.Type == dto.RealtimeEventTypeSessionCreated {
realtimeSession := realtimeEvent.Session realtimeSession := realtimeEvent.Session
...@@ -491,7 +492,7 @@ func OpenaiRealtimeHandler(c *gin.Context, info *relaycommon.RelayInfo) (*types. ...@@ -491,7 +492,7 @@ func OpenaiRealtimeHandler(c *gin.Context, info *relaycommon.RelayInfo) (*types.
errChan <- fmt.Errorf("error counting text token: %v", err) errChan <- fmt.Errorf("error counting text token: %v", err)
return return
} }
common.LogInfo(c, fmt.Sprintf("type: %s, textToken: %d, audioToken: %d", realtimeEvent.Type, textToken, audioToken)) logger.LogInfo(c, fmt.Sprintf("type: %s, textToken: %d, audioToken: %d", realtimeEvent.Type, textToken, audioToken))
localUsage.TotalTokens += textToken + audioToken localUsage.TotalTokens += textToken + audioToken
localUsage.OutputTokens += textToken + audioToken localUsage.OutputTokens += textToken + audioToken
localUsage.OutputTokenDetails.TextTokens += textToken localUsage.OutputTokenDetails.TextTokens += textToken
...@@ -517,7 +518,7 @@ func OpenaiRealtimeHandler(c *gin.Context, info *relaycommon.RelayInfo) (*types. ...@@ -517,7 +518,7 @@ func OpenaiRealtimeHandler(c *gin.Context, info *relaycommon.RelayInfo) (*types.
case <-targetClosed: case <-targetClosed:
case err := <-errChan: case err := <-errChan:
//return service.OpenAIErrorWrapper(err, "realtime_error", http.StatusInternalServerError), nil //return service.OpenAIErrorWrapper(err, "realtime_error", http.StatusInternalServerError), nil
common.LogError(c, "realtime error: "+err.Error()) logger.LogError(c, "realtime error: "+err.Error())
case <-c.Done(): case <-c.Done():
} }
...@@ -553,7 +554,7 @@ func preConsumeUsage(ctx *gin.Context, info *relaycommon.RelayInfo, usage *dto.R ...@@ -553,7 +554,7 @@ func preConsumeUsage(ctx *gin.Context, info *relaycommon.RelayInfo, usage *dto.R
} }
func OpenaiHandlerWithUsage(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Response) (*dto.Usage, *types.NewAPIError) { func OpenaiHandlerWithUsage(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Response) (*dto.Usage, *types.NewAPIError) {
defer common.CloseResponseBodyGracefully(resp) defer service.CloseResponseBodyGracefully(resp)
responseBody, err := io.ReadAll(resp.Body) responseBody, err := io.ReadAll(resp.Body)
if err != nil { if err != nil {
...@@ -567,7 +568,7 @@ func OpenaiHandlerWithUsage(c *gin.Context, info *relaycommon.RelayInfo, resp *h ...@@ -567,7 +568,7 @@ func OpenaiHandlerWithUsage(c *gin.Context, info *relaycommon.RelayInfo, resp *h
} }
// 写入新的 response body // 写入新的 response body
common.IOCopyBytesGracefully(c, resp, responseBody) service.IOCopyBytesGracefully(c, resp, responseBody)
// Once we've written to the client, we should not return errors anymore // Once we've written to the client, we should not return errors anymore
// because the upstream has already consumed resources and returned content // because the upstream has already consumed resources and returned content
......
...@@ -6,6 +6,7 @@ import ( ...@@ -6,6 +6,7 @@ import (
"net/http" "net/http"
"one-api/common" "one-api/common"
"one-api/dto" "one-api/dto"
"one-api/logger"
relaycommon "one-api/relay/common" relaycommon "one-api/relay/common"
"one-api/relay/helper" "one-api/relay/helper"
"one-api/service" "one-api/service"
...@@ -16,7 +17,7 @@ import ( ...@@ -16,7 +17,7 @@ import (
) )
func OaiResponsesHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Response) (*dto.Usage, *types.NewAPIError) { func OaiResponsesHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Response) (*dto.Usage, *types.NewAPIError) {
defer common.CloseResponseBodyGracefully(resp) defer service.CloseResponseBodyGracefully(resp)
// read response body // read response body
var responsesResponse dto.OpenAIResponsesResponse var responsesResponse dto.OpenAIResponsesResponse
...@@ -33,7 +34,7 @@ func OaiResponsesHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http ...@@ -33,7 +34,7 @@ func OaiResponsesHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http
} }
// 写入新的 response body // 写入新的 response body
common.IOCopyBytesGracefully(c, resp, responseBody) service.IOCopyBytesGracefully(c, resp, responseBody)
// compute usage // compute usage
usage := dto.Usage{} usage := dto.Usage{}
...@@ -54,7 +55,7 @@ func OaiResponsesHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http ...@@ -54,7 +55,7 @@ func OaiResponsesHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http
func OaiResponsesStreamHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Response) (*dto.Usage, *types.NewAPIError) { 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") logger.LogError(c, "invalid response or response body")
return nil, types.NewError(fmt.Errorf("invalid response"), types.ErrorCodeBadResponse) return nil, types.NewError(fmt.Errorf("invalid response"), types.ErrorCodeBadResponse)
} }
......
...@@ -7,6 +7,7 @@ import ( ...@@ -7,6 +7,7 @@ import (
"one-api/common" "one-api/common"
"one-api/constant" "one-api/constant"
"one-api/dto" "one-api/dto"
"one-api/logger"
relaycommon "one-api/relay/common" relaycommon "one-api/relay/common"
"one-api/relay/helper" "one-api/relay/helper"
"one-api/service" "one-api/service"
...@@ -58,15 +59,15 @@ func palmStreamHandler(c *gin.Context, resp *http.Response) (*types.NewAPIError, ...@@ -58,15 +59,15 @@ func palmStreamHandler(c *gin.Context, resp *http.Response) (*types.NewAPIError,
go func() { go func() {
responseBody, err := io.ReadAll(resp.Body) responseBody, err := io.ReadAll(resp.Body)
if err != nil { if err != nil {
common.SysError("error reading stream response: " + err.Error()) logger.SysError("error reading stream response: " + err.Error())
stopChan <- true stopChan <- true
return return
} }
common.CloseResponseBodyGracefully(resp) service.CloseResponseBodyGracefully(resp)
var palmResponse PaLMChatResponse var palmResponse PaLMChatResponse
err = json.Unmarshal(responseBody, &palmResponse) err = json.Unmarshal(responseBody, &palmResponse)
if err != nil { if err != nil {
common.SysError("error unmarshalling stream response: " + err.Error()) logger.SysError("error unmarshalling stream response: " + err.Error())
stopChan <- true stopChan <- true
return return
} }
...@@ -78,7 +79,7 @@ func palmStreamHandler(c *gin.Context, resp *http.Response) (*types.NewAPIError, ...@@ -78,7 +79,7 @@ func palmStreamHandler(c *gin.Context, resp *http.Response) (*types.NewAPIError,
} }
jsonResponse, err := json.Marshal(fullTextResponse) jsonResponse, err := json.Marshal(fullTextResponse)
if err != nil { if err != nil {
common.SysError("error marshalling stream response: " + err.Error()) logger.SysError("error marshalling stream response: " + err.Error())
stopChan <- true stopChan <- true
return return
} }
...@@ -96,7 +97,7 @@ func palmStreamHandler(c *gin.Context, resp *http.Response) (*types.NewAPIError, ...@@ -96,7 +97,7 @@ func palmStreamHandler(c *gin.Context, resp *http.Response) (*types.NewAPIError,
return false return false
} }
}) })
common.CloseResponseBodyGracefully(resp) service.CloseResponseBodyGracefully(resp)
return nil, responseText return nil, responseText
} }
...@@ -105,7 +106,7 @@ func palmHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Respons ...@@ -105,7 +106,7 @@ func palmHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Respons
if err != nil { if err != nil {
return nil, types.NewOpenAIError(err, types.ErrorCodeReadResponseBodyFailed, http.StatusInternalServerError) return nil, types.NewOpenAIError(err, types.ErrorCodeReadResponseBodyFailed, http.StatusInternalServerError)
} }
common.CloseResponseBodyGracefully(resp) service.CloseResponseBodyGracefully(resp)
var palmResponse PaLMChatResponse var palmResponse PaLMChatResponse
err = json.Unmarshal(responseBody, &palmResponse) err = json.Unmarshal(responseBody, &palmResponse)
if err != nil { if err != nil {
...@@ -133,6 +134,6 @@ func palmHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Respons ...@@ -133,6 +134,6 @@ func palmHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Respons
} }
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)
common.IOCopyBytesGracefully(c, resp, jsonResponse) service.IOCopyBytesGracefully(c, resp, jsonResponse)
return &usage, nil return &usage, nil
} }
...@@ -4,9 +4,9 @@ import ( ...@@ -4,9 +4,9 @@ import (
"encoding/json" "encoding/json"
"io" "io"
"net/http" "net/http"
"one-api/common"
"one-api/dto" "one-api/dto"
relaycommon "one-api/relay/common" relaycommon "one-api/relay/common"
"one-api/service"
"one-api/types" "one-api/types"
"github.com/gin-gonic/gin" "github.com/gin-gonic/gin"
...@@ -17,7 +17,7 @@ func siliconflowRerankHandler(c *gin.Context, info *relaycommon.RelayInfo, resp ...@@ -17,7 +17,7 @@ func siliconflowRerankHandler(c *gin.Context, info *relaycommon.RelayInfo, resp
if err != nil { if err != nil {
return nil, types.NewOpenAIError(err, types.ErrorCodeReadResponseBodyFailed, http.StatusInternalServerError) return nil, types.NewOpenAIError(err, types.ErrorCodeReadResponseBodyFailed, http.StatusInternalServerError)
} }
common.CloseResponseBodyGracefully(resp) service.CloseResponseBodyGracefully(resp)
var siliconflowResp SFRerankResponse var siliconflowResp SFRerankResponse
err = json.Unmarshal(responseBody, &siliconflowResp) err = json.Unmarshal(responseBody, &siliconflowResp)
if err != nil { if err != nil {
...@@ -39,6 +39,6 @@ func siliconflowRerankHandler(c *gin.Context, info *relaycommon.RelayInfo, resp ...@@ -39,6 +39,6 @@ func siliconflowRerankHandler(c *gin.Context, info *relaycommon.RelayInfo, resp
} }
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)
common.IOCopyBytesGracefully(c, resp, jsonResponse) service.IOCopyBytesGracefully(c, resp, jsonResponse)
return usage, nil return usage, nil
} }
...@@ -11,6 +11,7 @@ import ( ...@@ -11,6 +11,7 @@ import (
"one-api/common" "one-api/common"
"one-api/constant" "one-api/constant"
"one-api/dto" "one-api/dto"
"one-api/logger"
"one-api/relay/channel" "one-api/relay/channel"
relaycommon "one-api/relay/common" relaycommon "one-api/relay/common"
"one-api/service" "one-api/service"
...@@ -139,7 +140,7 @@ func (a *TaskAdaptor) FetchTask(baseUrl, key string, body map[string]any) (*http ...@@ -139,7 +140,7 @@ func (a *TaskAdaptor) FetchTask(baseUrl, key string, body map[string]any) (*http
req, err := http.NewRequest("POST", requestUrl, bytes.NewBuffer(byteBody)) req, err := http.NewRequest("POST", requestUrl, bytes.NewBuffer(byteBody))
if err != nil { if err != nil {
common.SysError(fmt.Sprintf("Get Task error: %v", err)) logger.SysError(fmt.Sprintf("Get Task error: %v", err))
return nil, err return nil, err
} }
defer req.Body.Close() defer req.Body.Close()
......
...@@ -13,6 +13,7 @@ import ( ...@@ -13,6 +13,7 @@ import (
"one-api/common" "one-api/common"
"one-api/constant" "one-api/constant"
"one-api/dto" "one-api/dto"
"one-api/logger"
relaycommon "one-api/relay/common" relaycommon "one-api/relay/common"
"one-api/relay/helper" "one-api/relay/helper"
"one-api/service" "one-api/service"
...@@ -106,7 +107,7 @@ func tencentStreamHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *htt ...@@ -106,7 +107,7 @@ func tencentStreamHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *htt
var tencentResponse TencentChatResponse var tencentResponse TencentChatResponse
err := json.Unmarshal([]byte(data), &tencentResponse) err := json.Unmarshal([]byte(data), &tencentResponse)
if err != nil { if err != nil {
common.SysError("error unmarshalling stream response: " + err.Error()) logger.SysError("error unmarshalling stream response: " + err.Error())
continue continue
} }
...@@ -117,17 +118,17 @@ func tencentStreamHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *htt ...@@ -117,17 +118,17 @@ func tencentStreamHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *htt
err = helper.ObjectData(c, response) err = helper.ObjectData(c, response)
if err != nil { if err != nil {
common.SysError(err.Error()) logger.SysError(err.Error())
} }
} }
if err := scanner.Err(); err != nil { if err := scanner.Err(); err != nil {
common.SysError("error reading stream: " + err.Error()) logger.SysError("error reading stream: " + err.Error())
} }
helper.Done(c) helper.Done(c)
common.CloseResponseBodyGracefully(resp) service.CloseResponseBodyGracefully(resp)
return service.ResponseText2Usage(responseText, info.UpstreamModelName, info.PromptTokens), nil return service.ResponseText2Usage(responseText, info.UpstreamModelName, info.PromptTokens), nil
} }
...@@ -138,7 +139,7 @@ func tencentHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Resp ...@@ -138,7 +139,7 @@ func tencentHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Resp
if err != nil { if err != nil {
return nil, types.NewOpenAIError(err, types.ErrorCodeReadResponseBodyFailed, http.StatusInternalServerError) return nil, types.NewOpenAIError(err, types.ErrorCodeReadResponseBodyFailed, http.StatusInternalServerError)
} }
common.CloseResponseBodyGracefully(resp) service.CloseResponseBodyGracefully(resp)
err = json.Unmarshal(responseBody, &tencentSb) err = json.Unmarshal(responseBody, &tencentSb)
if err != nil { if err != nil {
return nil, types.NewOpenAIError(err, types.ErrorCodeBadResponseBody, http.StatusInternalServerError) return nil, types.NewOpenAIError(err, types.ErrorCodeBadResponseBody, http.StatusInternalServerError)
...@@ -156,7 +157,7 @@ func tencentHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Resp ...@@ -156,7 +157,7 @@ func tencentHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Resp
} }
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)
common.IOCopyBytesGracefully(c, resp, jsonResponse) service.IOCopyBytesGracefully(c, resp, jsonResponse)
return &fullTextResponse.Usage, nil return &fullTextResponse.Usage, nil
} }
......
...@@ -6,6 +6,7 @@ import ( ...@@ -6,6 +6,7 @@ import (
"net/http" "net/http"
"one-api/common" "one-api/common"
"one-api/dto" "one-api/dto"
"one-api/logger"
"one-api/relay/channel/openai" "one-api/relay/channel/openai"
relaycommon "one-api/relay/common" relaycommon "one-api/relay/common"
"one-api/relay/helper" "one-api/relay/helper"
...@@ -47,7 +48,7 @@ func xAIStreamHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Re ...@@ -47,7 +48,7 @@ func xAIStreamHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Re
var xAIResp *dto.ChatCompletionsStreamResponse var xAIResp *dto.ChatCompletionsStreamResponse
err := json.Unmarshal([]byte(data), &xAIResp) err := json.Unmarshal([]byte(data), &xAIResp)
if err != nil { if err != nil {
common.SysError("error unmarshalling stream response: " + err.Error()) logger.SysError("error unmarshalling stream response: " + err.Error())
return true return true
} }
...@@ -63,7 +64,7 @@ func xAIStreamHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Re ...@@ -63,7 +64,7 @@ func xAIStreamHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Re
_ = openai.ProcessStreamResponse(*openaiResponse, &responseTextBuilder, &toolCount) _ = openai.ProcessStreamResponse(*openaiResponse, &responseTextBuilder, &toolCount)
err = helper.ObjectData(c, openaiResponse) err = helper.ObjectData(c, openaiResponse)
if err != nil { if err != nil {
common.SysError(err.Error()) logger.SysError(err.Error())
} }
return true return true
}) })
...@@ -74,12 +75,12 @@ func xAIStreamHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Re ...@@ -74,12 +75,12 @@ func xAIStreamHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Re
} }
helper.Done(c) helper.Done(c)
common.CloseResponseBodyGracefully(resp) service.CloseResponseBodyGracefully(resp)
return usage, nil return usage, nil
} }
func xAIHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Response) (*dto.Usage, *types.NewAPIError) { func xAIHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Response) (*dto.Usage, *types.NewAPIError) {
defer common.CloseResponseBodyGracefully(resp) defer service.CloseResponseBodyGracefully(resp)
responseBody, err := io.ReadAll(resp.Body) responseBody, err := io.ReadAll(resp.Body)
if err != nil { if err != nil {
...@@ -101,7 +102,7 @@ func xAIHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Response ...@@ -101,7 +102,7 @@ func xAIHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Response
return nil, types.NewError(err, types.ErrorCodeBadResponseBody) return nil, types.NewError(err, types.ErrorCodeBadResponseBody)
} }
common.IOCopyBytesGracefully(c, resp, encodeJson) service.IOCopyBytesGracefully(c, resp, encodeJson)
return xaiResponse.Usage, nil return xaiResponse.Usage, nil
} }
...@@ -11,6 +11,7 @@ import ( ...@@ -11,6 +11,7 @@ import (
"one-api/common" "one-api/common"
"one-api/constant" "one-api/constant"
"one-api/dto" "one-api/dto"
"one-api/logger"
"one-api/relay/helper" "one-api/relay/helper"
"one-api/types" "one-api/types"
"strings" "strings"
...@@ -143,7 +144,7 @@ func xunfeiStreamHandler(c *gin.Context, textRequest dto.GeneralOpenAIRequest, a ...@@ -143,7 +144,7 @@ func xunfeiStreamHandler(c *gin.Context, textRequest dto.GeneralOpenAIRequest, a
response := streamResponseXunfei2OpenAI(&xunfeiResponse) response := streamResponseXunfei2OpenAI(&xunfeiResponse)
jsonResponse, err := json.Marshal(response) jsonResponse, err := json.Marshal(response)
if err != nil { if err != nil {
common.SysError("error marshalling stream response: " + err.Error()) logger.SysError("error marshalling stream response: " + err.Error())
return true return true
} }
c.Render(-1, common.CustomEvent{Data: "data: " + string(jsonResponse)}) c.Render(-1, common.CustomEvent{Data: "data: " + string(jsonResponse)})
...@@ -218,20 +219,20 @@ func xunfeiMakeRequest(textRequest dto.GeneralOpenAIRequest, domain, authUrl, ap ...@@ -218,20 +219,20 @@ func xunfeiMakeRequest(textRequest dto.GeneralOpenAIRequest, domain, authUrl, ap
for { for {
_, msg, err := conn.ReadMessage() _, msg, err := conn.ReadMessage()
if err != nil { if err != nil {
common.SysError("error reading stream response: " + err.Error()) logger.SysError("error reading stream response: " + err.Error())
break break
} }
var response XunfeiChatResponse var response XunfeiChatResponse
err = json.Unmarshal(msg, &response) err = json.Unmarshal(msg, &response)
if err != nil { if err != nil {
common.SysError("error unmarshalling stream response: " + err.Error()) logger.SysError("error unmarshalling stream response: " + err.Error())
break break
} }
dataChan <- response dataChan <- response
if response.Payload.Choices.Status == 2 { if response.Payload.Choices.Status == 2 {
err := conn.Close() err := conn.Close()
if err != nil { if err != nil {
common.SysError("error closing websocket connection: " + err.Error()) logger.SysError("error closing websocket connection: " + err.Error())
} }
break break
} }
...@@ -282,6 +283,6 @@ func getAPIVersion(c *gin.Context, modelName string) string { ...@@ -282,6 +283,6 @@ func getAPIVersion(c *gin.Context, modelName string) string {
return apiVersion return apiVersion
} }
apiVersion = "v1.1" apiVersion = "v1.1"
common.SysLog("api_version not found, using default: " + apiVersion) logger.SysLog("api_version not found, using default: " + apiVersion)
return apiVersion return apiVersion
} }
...@@ -8,8 +8,10 @@ import ( ...@@ -8,8 +8,10 @@ import (
"one-api/common" "one-api/common"
"one-api/constant" "one-api/constant"
"one-api/dto" "one-api/dto"
"one-api/logger"
relaycommon "one-api/relay/common" relaycommon "one-api/relay/common"
"one-api/relay/helper" "one-api/relay/helper"
"one-api/service"
"one-api/types" "one-api/types"
"strings" "strings"
"sync" "sync"
...@@ -38,7 +40,7 @@ func getZhipuToken(apikey string) string { ...@@ -38,7 +40,7 @@ func getZhipuToken(apikey string) string {
split := strings.Split(apikey, ".") split := strings.Split(apikey, ".")
if len(split) != 2 { if len(split) != 2 {
common.SysError("invalid zhipu key: " + apikey) logger.SysError("invalid zhipu key: " + apikey)
return "" return ""
} }
...@@ -186,7 +188,7 @@ func zhipuStreamHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http. ...@@ -186,7 +188,7 @@ func zhipuStreamHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.
response := streamResponseZhipu2OpenAI(data) response := streamResponseZhipu2OpenAI(data)
jsonResponse, err := json.Marshal(response) jsonResponse, err := json.Marshal(response)
if err != nil { if err != nil {
common.SysError("error marshalling stream response: " + err.Error()) logger.SysError("error marshalling stream response: " + err.Error())
return true return true
} }
c.Render(-1, common.CustomEvent{Data: "data: " + string(jsonResponse)}) c.Render(-1, common.CustomEvent{Data: "data: " + string(jsonResponse)})
...@@ -195,13 +197,13 @@ func zhipuStreamHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http. ...@@ -195,13 +197,13 @@ func zhipuStreamHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.
var zhipuResponse ZhipuStreamMetaResponse var zhipuResponse ZhipuStreamMetaResponse
err := json.Unmarshal([]byte(data), &zhipuResponse) err := json.Unmarshal([]byte(data), &zhipuResponse)
if err != nil { if err != nil {
common.SysError("error unmarshalling stream response: " + err.Error()) logger.SysError("error unmarshalling stream response: " + err.Error())
return true return true
} }
response, zhipuUsage := streamMetaResponseZhipu2OpenAI(&zhipuResponse) response, zhipuUsage := streamMetaResponseZhipu2OpenAI(&zhipuResponse)
jsonResponse, err := json.Marshal(response) jsonResponse, err := json.Marshal(response)
if err != nil { if err != nil {
common.SysError("error marshalling stream response: " + err.Error()) logger.SysError("error marshalling stream response: " + err.Error())
return true return true
} }
usage = zhipuUsage usage = zhipuUsage
...@@ -212,7 +214,7 @@ func zhipuStreamHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http. ...@@ -212,7 +214,7 @@ func zhipuStreamHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.
return false return false
} }
}) })
common.CloseResponseBodyGracefully(resp) service.CloseResponseBodyGracefully(resp)
return usage, nil return usage, nil
} }
...@@ -222,7 +224,7 @@ func zhipuHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Respon ...@@ -222,7 +224,7 @@ func zhipuHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Respon
if err != nil { if err != nil {
return nil, types.NewOpenAIError(err, types.ErrorCodeReadResponseBodyFailed, http.StatusInternalServerError) return nil, types.NewOpenAIError(err, types.ErrorCodeReadResponseBodyFailed, http.StatusInternalServerError)
} }
common.CloseResponseBodyGracefully(resp) service.CloseResponseBodyGracefully(resp)
err = json.Unmarshal(responseBody, &zhipuResponse) err = json.Unmarshal(responseBody, &zhipuResponse)
if err != nil { if err != nil {
return nil, types.NewOpenAIError(err, types.ErrorCodeBadResponseBody, http.StatusInternalServerError) return nil, types.NewOpenAIError(err, types.ErrorCodeBadResponseBody, http.StatusInternalServerError)
......
...@@ -10,6 +10,7 @@ import ( ...@@ -10,6 +10,7 @@ import (
"one-api/common" "one-api/common"
"one-api/constant" "one-api/constant"
"one-api/dto" "one-api/dto"
"one-api/logger"
"one-api/model" "one-api/model"
relaycommon "one-api/relay/common" relaycommon "one-api/relay/common"
relayconstant "one-api/relay/constant" relayconstant "one-api/relay/constant"
...@@ -214,7 +215,7 @@ func RelaySwapFace(c *gin.Context) *dto.MidjourneyResponse { ...@@ -214,7 +215,7 @@ func RelaySwapFace(c *gin.Context) *dto.MidjourneyResponse {
if mjResp.StatusCode == 200 && mjResp.Response.Code == 1 { if mjResp.StatusCode == 200 && mjResp.Response.Code == 1 {
err := service.PostConsumeQuota(relayInfo, priceData.Quota, 0, true) err := service.PostConsumeQuota(relayInfo, priceData.Quota, 0, true)
if err != nil { if err != nil {
common.SysError("error consuming token remain quota: " + err.Error()) logger.SysError("error consuming token remain quota: " + err.Error())
} }
tokenName := c.GetString("token_name") tokenName := c.GetString("token_name")
...@@ -300,7 +301,7 @@ func RelayMidjourneyTaskImageSeed(c *gin.Context) *dto.MidjourneyResponse { ...@@ -300,7 +301,7 @@ func RelayMidjourneyTaskImageSeed(c *gin.Context) *dto.MidjourneyResponse {
if err != nil { if err != nil {
return service.MidjourneyErrorWrapper(constant.MjRequestError, "unmarshal_response_body_failed") return service.MidjourneyErrorWrapper(constant.MjRequestError, "unmarshal_response_body_failed")
} }
common.IOCopyBytesGracefully(c, nil, respBody) service.IOCopyBytesGracefully(c, nil, respBody)
return nil return nil
} }
...@@ -521,7 +522,7 @@ func RelayMidjourneySubmit(c *gin.Context, relayMode int) *dto.MidjourneyRespons ...@@ -521,7 +522,7 @@ func RelayMidjourneySubmit(c *gin.Context, relayMode int) *dto.MidjourneyRespons
if consumeQuota && midjResponseWithStatus.StatusCode == 200 { if consumeQuota && midjResponseWithStatus.StatusCode == 200 {
err := service.PostConsumeQuota(relayInfo, priceData.Quota, 0, true) err := service.PostConsumeQuota(relayInfo, priceData.Quota, 0, true)
if err != nil { if err != nil {
common.SysError("error consuming token remain quota: " + err.Error()) logger.SysError("error consuming token remain quota: " + err.Error())
} }
tokenName := c.GetString("token_name") tokenName := c.GetString("token_name")
logContent := fmt.Sprintf("模型固定价格 %.2f,分组倍率 %.2f,操作 %s,ID %s", priceData.ModelPrice, priceData.GroupRatioInfo.GroupRatio, midjRequest.Action, midjResponse.Result) logContent := fmt.Sprintf("模型固定价格 %.2f,分组倍率 %.2f,操作 %s,ID %s", priceData.ModelPrice, priceData.GroupRatioInfo.GroupRatio, midjRequest.Action, midjResponse.Result)
...@@ -572,7 +573,7 @@ func RelayMidjourneySubmit(c *gin.Context, relayMode int) *dto.MidjourneyRespons ...@@ -572,7 +573,7 @@ func RelayMidjourneySubmit(c *gin.Context, relayMode int) *dto.MidjourneyRespons
//无实例账号自动禁用渠道(No available account instance) //无实例账号自动禁用渠道(No available account instance)
channel, err := model.GetChannelById(midjourneyTask.ChannelId, true) channel, err := model.GetChannelById(midjourneyTask.ChannelId, true)
if err != nil { if err != nil {
common.SysError("get_channel_null: " + err.Error()) logger.SysError("get_channel_null: " + err.Error())
} }
if channel.GetAutoBan() && common.AutomaticDisableChannelEnabled { if channel.GetAutoBan() && common.AutomaticDisableChannelEnabled {
model.UpdateChannelStatus(midjourneyTask.ChannelId, "", 2, "No available account instance") model.UpdateChannelStatus(midjourneyTask.ChannelId, "", 2, "No available account instance")
......
...@@ -2,7 +2,6 @@ package relay ...@@ -2,7 +2,6 @@ package relay
import ( import (
"bytes" "bytes"
"errors"
"fmt" "fmt"
"io" "io"
"net/http" "net/http"
...@@ -18,68 +17,26 @@ import ( ...@@ -18,68 +17,26 @@ import (
"github.com/gin-gonic/gin" "github.com/gin-gonic/gin"
) )
func getAndValidateClaudeRequest(c *gin.Context) (textRequest *dto.ClaudeRequest, err error) { func ClaudeHelper(c *gin.Context, info *relaycommon.RelayInfo) (newAPIError *types.NewAPIError) {
textRequest = &dto.ClaudeRequest{}
err = c.ShouldBindJSON(textRequest)
if err != nil {
return nil, err
}
if textRequest.Messages == nil || len(textRequest.Messages) == 0 {
return nil, errors.New("field messages is required")
}
if textRequest.Model == "" {
return nil, errors.New("field model is required")
}
return textRequest, nil
}
func ClaudeHelper(c *gin.Context) (newAPIError *types.NewAPIError) {
relayInfo := relaycommon.GenRelayInfoClaude(c) info.InitChannelMeta(c)
// get & validate textRequest 获取并验证文本请求 textRequest, ok := info.Request.(*dto.ClaudeRequest)
textRequest, err := getAndValidateClaudeRequest(c)
if err != nil {
return types.NewError(err, types.ErrorCodeInvalidRequest, types.ErrOptionWithSkipRetry())
}
if textRequest.Stream { if !ok {
relayInfo.IsStream = true common.FatalLog(fmt.Sprintf("invalid request type, expected dto.ClaudeRequest, got %T", info.Request))
} }
err = helper.ModelMappedHelper(c, relayInfo, textRequest) err := helper.ModelMappedHelper(c, info, textRequest)
if err != nil { if err != nil {
return types.NewError(err, types.ErrorCodeChannelModelMappedError, types.ErrOptionWithSkipRetry()) return types.NewError(err, types.ErrorCodeChannelModelMappedError, types.ErrOptionWithSkipRetry())
} }
promptTokens, err := getClaudePromptTokens(textRequest, relayInfo) adaptor := GetAdaptor(info.ApiType)
// count messages token error 计算promptTokens错误
if err != nil {
return types.NewError(err, types.ErrorCodeCountTokenFailed, types.ErrOptionWithSkipRetry())
}
priceData, err := helper.ModelPriceHelper(c, relayInfo, promptTokens, int(textRequest.MaxTokens))
if err != nil {
return types.NewError(err, types.ErrorCodeModelPriceError, types.ErrOptionWithSkipRetry())
}
// pre-consume quota 预消耗配额
preConsumedQuota, userQuota, newAPIError := preConsumeQuota(c, priceData.ShouldPreConsumedQuota, relayInfo)
if newAPIError != nil {
return newAPIError
}
defer func() {
if newAPIError != nil {
returnPreConsumedQuota(c, relayInfo, userQuota, preConsumedQuota)
}
}()
adaptor := GetAdaptor(relayInfo.ApiType)
if adaptor == nil { if adaptor == nil {
return types.NewError(fmt.Errorf("invalid api type: %d", relayInfo.ApiType), types.ErrorCodeInvalidApiType, types.ErrOptionWithSkipRetry()) return types.NewError(fmt.Errorf("invalid api type: %d", info.ApiType), types.ErrorCodeInvalidApiType, types.ErrOptionWithSkipRetry())
} }
adaptor.Init(relayInfo) adaptor.Init(info)
if textRequest.MaxTokens == 0 { if textRequest.MaxTokens == 0 {
textRequest.MaxTokens = uint(model_setting.GetClaudeSettings().GetDefaultMaxTokens(textRequest.Model)) textRequest.MaxTokens = uint(model_setting.GetClaudeSettings().GetDefaultMaxTokens(textRequest.Model))
...@@ -104,18 +61,18 @@ func ClaudeHelper(c *gin.Context) (newAPIError *types.NewAPIError) { ...@@ -104,18 +61,18 @@ func ClaudeHelper(c *gin.Context) (newAPIError *types.NewAPIError) {
textRequest.Temperature = common.GetPointer[float64](1.0) textRequest.Temperature = common.GetPointer[float64](1.0)
} }
textRequest.Model = strings.TrimSuffix(textRequest.Model, "-thinking") textRequest.Model = strings.TrimSuffix(textRequest.Model, "-thinking")
relayInfo.UpstreamModelName = textRequest.Model info.UpstreamModelName = textRequest.Model
} }
var requestBody io.Reader var requestBody io.Reader
if model_setting.GetGlobalSettings().PassThroughRequestEnabled || relayInfo.ChannelSetting.PassThroughBodyEnabled { if model_setting.GetGlobalSettings().PassThroughRequestEnabled || info.ChannelSetting.PassThroughBodyEnabled {
body, err := common.GetRequestBody(c) body, err := common.GetRequestBody(c)
if err != nil { if err != nil {
return types.NewErrorWithStatusCode(err, types.ErrorCodeReadRequestBodyFailed, http.StatusBadRequest, types.ErrOptionWithSkipRetry()) return types.NewErrorWithStatusCode(err, types.ErrorCodeReadRequestBodyFailed, http.StatusBadRequest, types.ErrOptionWithSkipRetry())
} }
requestBody = bytes.NewBuffer(body) requestBody = bytes.NewBuffer(body)
} else { } else {
convertedRequest, err := adaptor.ConvertClaudeRequest(c, relayInfo, textRequest) convertedRequest, err := adaptor.ConvertClaudeRequest(c, info, textRequest)
if err != nil { if err != nil {
return types.NewError(err, types.ErrorCodeConvertRequestFailed, types.ErrOptionWithSkipRetry()) return types.NewError(err, types.ErrorCodeConvertRequestFailed, types.ErrOptionWithSkipRetry())
} }
...@@ -125,10 +82,10 @@ func ClaudeHelper(c *gin.Context) (newAPIError *types.NewAPIError) { ...@@ -125,10 +82,10 @@ func ClaudeHelper(c *gin.Context) (newAPIError *types.NewAPIError) {
} }
// apply param override // apply param override
if len(relayInfo.ParamOverride) > 0 { if len(info.ParamOverride) > 0 {
reqMap := make(map[string]interface{}) reqMap := make(map[string]interface{})
_ = common.Unmarshal(jsonData, &reqMap) _ = common.Unmarshal(jsonData, &reqMap)
for key, value := range relayInfo.ParamOverride { for key, value := range info.ParamOverride {
reqMap[key] = value reqMap[key] = value
} }
jsonData, err = common.Marshal(reqMap) jsonData, err = common.Marshal(reqMap)
...@@ -145,14 +102,14 @@ func ClaudeHelper(c *gin.Context) (newAPIError *types.NewAPIError) { ...@@ -145,14 +102,14 @@ func ClaudeHelper(c *gin.Context) (newAPIError *types.NewAPIError) {
statusCodeMappingStr := c.GetString("status_code_mapping") statusCodeMappingStr := c.GetString("status_code_mapping")
var httpResp *http.Response var httpResp *http.Response
resp, err := adaptor.DoRequest(c, relayInfo, requestBody) resp, err := adaptor.DoRequest(c, info, requestBody)
if err != nil { if err != nil {
return types.NewOpenAIError(err, types.ErrorCodeDoRequestFailed, http.StatusInternalServerError) return types.NewOpenAIError(err, types.ErrorCodeDoRequestFailed, http.StatusInternalServerError)
} }
if resp != nil { if resp != nil {
httpResp = resp.(*http.Response) httpResp = resp.(*http.Response)
relayInfo.IsStream = relayInfo.IsStream || strings.HasPrefix(httpResp.Header.Get("Content-Type"), "text/event-stream") info.IsStream = info.IsStream || strings.HasPrefix(httpResp.Header.Get("Content-Type"), "text/event-stream")
if httpResp.StatusCode != http.StatusOK { if httpResp.StatusCode != http.StatusOK {
newAPIError = service.RelayErrorHandler(httpResp, false) newAPIError = service.RelayErrorHandler(httpResp, false)
// reset status code 重置状态码 // reset status code 重置状态码
...@@ -161,24 +118,14 @@ func ClaudeHelper(c *gin.Context) (newAPIError *types.NewAPIError) { ...@@ -161,24 +118,14 @@ func ClaudeHelper(c *gin.Context) (newAPIError *types.NewAPIError) {
} }
} }
usage, newAPIError := adaptor.DoResponse(c, httpResp, relayInfo) usage, newAPIError := adaptor.DoResponse(c, httpResp, info)
//log.Printf("usage: %v", usage) //log.Printf("usage: %v", usage)
if newAPIError != nil { if newAPIError != nil {
// reset status code 重置状态码 // reset status code 重置状态码
service.ResetStatusCode(newAPIError, statusCodeMappingStr) service.ResetStatusCode(newAPIError, statusCodeMappingStr)
return newAPIError return newAPIError
} }
service.PostClaudeConsumeQuota(c, relayInfo, usage.(*dto.Usage), preConsumedQuota, userQuota, priceData, "")
return nil
}
func getClaudePromptTokens(textRequest *dto.ClaudeRequest, info *relaycommon.RelayInfo) (int, error) { service.PostClaudeConsumeQuota(c, info, usage.(*dto.Usage))
var promptTokens int return nil
var err error
switch info.RelayMode {
default:
promptTokens, err = service.CountTokenClaudeRequest(*textRequest, info.UpstreamModelName)
}
info.PromptTokens = promptTokens
return promptTokens, err
} }
...@@ -8,6 +8,7 @@ import ( ...@@ -8,6 +8,7 @@ import (
"one-api/dto" "one-api/dto"
"one-api/relay/channel/xinference" "one-api/relay/channel/xinference"
relaycommon "one-api/relay/common" relaycommon "one-api/relay/common"
"one-api/service"
"one-api/types" "one-api/types"
"github.com/gin-gonic/gin" "github.com/gin-gonic/gin"
...@@ -18,7 +19,7 @@ func RerankHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Respo ...@@ -18,7 +19,7 @@ func RerankHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Respo
if err != nil { if err != nil {
return nil, types.NewOpenAIError(err, types.ErrorCodeReadResponseBodyFailed, http.StatusInternalServerError) return nil, types.NewOpenAIError(err, types.ErrorCodeReadResponseBodyFailed, http.StatusInternalServerError)
} }
common.CloseResponseBodyGracefully(resp) service.CloseResponseBodyGracefully(resp)
if common.DebugEnabled { if common.DebugEnabled {
println("reranker response body: ", string(responseBody)) println("reranker response body: ", string(responseBody))
} }
......
...@@ -8,7 +8,6 @@ import ( ...@@ -8,7 +8,6 @@ import (
"one-api/common" "one-api/common"
"one-api/dto" "one-api/dto"
relaycommon "one-api/relay/common" relaycommon "one-api/relay/common"
relayconstant "one-api/relay/constant"
"one-api/relay/helper" "one-api/relay/helper"
"one-api/service" "one-api/service"
"one-api/types" "one-api/types"
...@@ -16,69 +15,27 @@ import ( ...@@ -16,69 +15,27 @@ import (
"github.com/gin-gonic/gin" "github.com/gin-gonic/gin"
) )
func getEmbeddingPromptToken(embeddingRequest dto.EmbeddingRequest) int { func EmbeddingHelper(c *gin.Context, info *relaycommon.RelayInfo) (newAPIError *types.NewAPIError) {
token := service.CountTokenInput(embeddingRequest.Input, embeddingRequest.Model)
return token
}
func validateEmbeddingRequest(c *gin.Context, info *relaycommon.RelayInfo, embeddingRequest dto.EmbeddingRequest) error {
if embeddingRequest.Input == nil {
return fmt.Errorf("input is empty")
}
if info.RelayMode == relayconstant.RelayModeModerations && embeddingRequest.Model == "" {
embeddingRequest.Model = "omni-moderation-latest"
}
if info.RelayMode == relayconstant.RelayModeEmbeddings && embeddingRequest.Model == "" {
embeddingRequest.Model = c.Param("model")
}
return nil
}
func EmbeddingHelper(c *gin.Context) (newAPIError *types.NewAPIError) { info.InitChannelMeta(c)
relayInfo := relaycommon.GenRelayInfoEmbedding(c)
var embeddingRequest *dto.EmbeddingRequest embeddingRequest, ok := info.Request.(*dto.EmbeddingRequest)
err := common.UnmarshalBodyReusable(c, &embeddingRequest) if !ok {
if err != nil { common.FatalLog(fmt.Sprintf("invalid request type, expected dto.ClaudeRequest, got %T", info.Request))
common.LogError(c, fmt.Sprintf("getAndValidateTextRequest failed: %s", err.Error()))
return types.NewError(err, types.ErrorCodeInvalidRequest, types.ErrOptionWithSkipRetry())
} }
err = validateEmbeddingRequest(c, relayInfo, *embeddingRequest) err := helper.ModelMappedHelper(c, info, embeddingRequest)
if err != nil {
return types.NewError(err, types.ErrorCodeInvalidRequest, types.ErrOptionWithSkipRetry())
}
err = helper.ModelMappedHelper(c, relayInfo, embeddingRequest)
if err != nil { if err != nil {
return types.NewError(err, types.ErrorCodeChannelModelMappedError, types.ErrOptionWithSkipRetry()) return types.NewError(err, types.ErrorCodeChannelModelMappedError, types.ErrOptionWithSkipRetry())
} }
promptToken := getEmbeddingPromptToken(*embeddingRequest) adaptor := GetAdaptor(info.ApiType)
relayInfo.PromptTokens = promptToken
priceData, err := helper.ModelPriceHelper(c, relayInfo, promptToken, 0)
if err != nil {
return types.NewError(err, types.ErrorCodeModelPriceError, types.ErrOptionWithSkipRetry())
}
// pre-consume quota 预消耗配额
preConsumedQuota, userQuota, newAPIError := preConsumeQuota(c, priceData.ShouldPreConsumedQuota, relayInfo)
if newAPIError != nil {
return newAPIError
}
defer func() {
if newAPIError != nil {
returnPreConsumedQuota(c, relayInfo, userQuota, preConsumedQuota)
}
}()
adaptor := GetAdaptor(relayInfo.ApiType)
if adaptor == nil { if adaptor == nil {
return types.NewError(fmt.Errorf("invalid api type: %d", relayInfo.ApiType), types.ErrorCodeInvalidApiType, types.ErrOptionWithSkipRetry()) return types.NewError(fmt.Errorf("invalid api type: %d", info.ApiType), types.ErrorCodeInvalidApiType, types.ErrOptionWithSkipRetry())
} }
adaptor.Init(relayInfo) adaptor.Init(info)
convertedRequest, err := adaptor.ConvertEmbeddingRequest(c, relayInfo, *embeddingRequest) convertedRequest, err := adaptor.ConvertEmbeddingRequest(c, info, *embeddingRequest)
if err != nil { if err != nil {
return types.NewError(err, types.ErrorCodeConvertRequestFailed, types.ErrOptionWithSkipRetry()) return types.NewError(err, types.ErrorCodeConvertRequestFailed, types.ErrOptionWithSkipRetry())
} }
...@@ -88,7 +45,7 @@ func EmbeddingHelper(c *gin.Context) (newAPIError *types.NewAPIError) { ...@@ -88,7 +45,7 @@ func EmbeddingHelper(c *gin.Context) (newAPIError *types.NewAPIError) {
} }
requestBody := bytes.NewBuffer(jsonData) requestBody := bytes.NewBuffer(jsonData)
statusCodeMappingStr := c.GetString("status_code_mapping") statusCodeMappingStr := c.GetString("status_code_mapping")
resp, err := adaptor.DoRequest(c, relayInfo, requestBody) resp, err := adaptor.DoRequest(c, info, requestBody)
if err != nil { if err != nil {
return types.NewOpenAIError(err, types.ErrorCodeDoRequestFailed, http.StatusInternalServerError) return types.NewOpenAIError(err, types.ErrorCodeDoRequestFailed, http.StatusInternalServerError)
} }
...@@ -104,12 +61,12 @@ func EmbeddingHelper(c *gin.Context) (newAPIError *types.NewAPIError) { ...@@ -104,12 +61,12 @@ func EmbeddingHelper(c *gin.Context) (newAPIError *types.NewAPIError) {
} }
} }
usage, newAPIError := adaptor.DoResponse(c, httpResp, relayInfo) usage, newAPIError := adaptor.DoResponse(c, httpResp, info)
if newAPIError != nil { if newAPIError != nil {
// reset status code 重置状态码 // reset status code 重置状态码
service.ResetStatusCode(newAPIError, statusCodeMappingStr) service.ResetStatusCode(newAPIError, statusCodeMappingStr)
return newAPIError return newAPIError
} }
postConsumeQuota(c, relayInfo, usage.(*dto.Usage), preConsumedQuota, userQuota, priceData, "") postConsumeQuota(c, info, usage.(*dto.Usage), "")
return nil return nil
} }
...@@ -7,6 +7,7 @@ import ( ...@@ -7,6 +7,7 @@ import (
"net/http" "net/http"
"one-api/common" "one-api/common"
"one-api/dto" "one-api/dto"
"one-api/logger"
"one-api/types" "one-api/types"
"github.com/gin-gonic/gin" "github.com/gin-gonic/gin"
...@@ -100,7 +101,7 @@ func Done(c *gin.Context) { ...@@ -100,7 +101,7 @@ func Done(c *gin.Context) {
func WssString(c *gin.Context, ws *websocket.Conn, str string) error { func WssString(c *gin.Context, ws *websocket.Conn, str string) error {
if ws == nil { if ws == nil {
common.LogError(c, "websocket connection is nil") logger.LogError(c, "websocket connection is nil")
return errors.New("websocket connection is nil") return errors.New("websocket connection is nil")
} }
//common.LogInfo(c, fmt.Sprintf("sending message: %s", str)) //common.LogInfo(c, fmt.Sprintf("sending message: %s", str))
...@@ -113,7 +114,7 @@ func WssObject(c *gin.Context, ws *websocket.Conn, object interface{}) error { ...@@ -113,7 +114,7 @@ func WssObject(c *gin.Context, ws *websocket.Conn, object interface{}) error {
return fmt.Errorf("error marshalling object: %w", err) return fmt.Errorf("error marshalling object: %w", err)
} }
if ws == nil { if ws == nil {
common.LogError(c, "websocket connection is nil") logger.LogError(c, "websocket connection is nil")
return errors.New("websocket connection is nil") return errors.New("websocket connection is nil")
} }
//common.LogInfo(c, fmt.Sprintf("sending message: %s", jsonData)) //common.LogInfo(c, fmt.Sprintf("sending message: %s", jsonData))
......
...@@ -4,9 +4,10 @@ import ( ...@@ -4,9 +4,10 @@ import (
"encoding/json" "encoding/json"
"errors" "errors"
"fmt" "fmt"
common2 "one-api/common"
"one-api/dto" "one-api/dto"
common2 "one-api/logger"
"one-api/relay/common" "one-api/relay/common"
"one-api/types"
"github.com/gin-gonic/gin" "github.com/gin-gonic/gin"
) )
...@@ -54,29 +55,29 @@ func ModelMappedHelper(c *gin.Context, info *common.RelayInfo, request any) erro ...@@ -54,29 +55,29 @@ func ModelMappedHelper(c *gin.Context, info *common.RelayInfo, request any) erro
} }
if request != nil { if request != nil {
switch info.RelayFormat { switch info.RelayFormat {
case common.RelayFormatGemini: case types.RelayFormatGemini:
// Gemini 模型映射 // Gemini 模型映射
case common.RelayFormatClaude: case types.RelayFormatClaude:
if claudeRequest, ok := request.(*dto.ClaudeRequest); ok { if claudeRequest, ok := request.(*dto.ClaudeRequest); ok {
claudeRequest.Model = info.UpstreamModelName claudeRequest.Model = info.UpstreamModelName
} }
case common.RelayFormatOpenAIResponses: case types.RelayFormatOpenAIResponses:
if openAIResponsesRequest, ok := request.(*dto.OpenAIResponsesRequest); ok { if openAIResponsesRequest, ok := request.(*dto.OpenAIResponsesRequest); ok {
openAIResponsesRequest.Model = info.UpstreamModelName openAIResponsesRequest.Model = info.UpstreamModelName
} }
case common.RelayFormatOpenAIAudio: case types.RelayFormatOpenAIAudio:
if openAIAudioRequest, ok := request.(*dto.AudioRequest); ok { if openAIAudioRequest, ok := request.(*dto.AudioRequest); ok {
openAIAudioRequest.Model = info.UpstreamModelName openAIAudioRequest.Model = info.UpstreamModelName
} }
case common.RelayFormatOpenAIImage: case types.RelayFormatOpenAIImage:
if imageRequest, ok := request.(*dto.ImageRequest); ok { if imageRequest, ok := request.(*dto.ImageRequest); ok {
imageRequest.Model = info.UpstreamModelName imageRequest.Model = info.UpstreamModelName
} }
case common.RelayFormatRerank: case types.RelayFormatRerank:
if rerankRequest, ok := request.(*dto.RerankRequest); ok { if rerankRequest, ok := request.(*dto.RerankRequest); ok {
rerankRequest.Model = info.UpstreamModelName rerankRequest.Model = info.UpstreamModelName
} }
case common.RelayFormatEmbedding: case types.RelayFormatEmbedding:
if embeddingRequest, ok := request.(*dto.EmbeddingRequest); ok { if embeddingRequest, ok := request.(*dto.EmbeddingRequest); ok {
embeddingRequest.Model = info.UpstreamModelName embeddingRequest.Model = info.UpstreamModelName
} }
......
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