diff --git a/controller/midjourney.go b/controller/midjourney.go index bf52314a7581..dbc7fb9d7360 100644 --- a/controller/midjourney.go +++ b/controller/midjourney.go @@ -14,6 +14,7 @@ import ( "github.com/QuantumNous/new-api/model" "github.com/QuantumNous/new-api/service" "github.com/QuantumNous/new-api/setting" + "github.com/QuantumNous/new-api/setting/operation_setting" "github.com/QuantumNous/new-api/setting/system_setting" "github.com/gin-gonic/gin" @@ -332,11 +333,13 @@ func GetUserMidjourney(c *gin.Context) { items := model.GetAllUserTask(userId, pageInfo.GetStartIdx(), pageInfo.GetPageSize(), queryParams) total := model.CountAllUserTask(userId, queryParams) - if setting.MjForwardUrlEnabled { - for i, midjourney := range items { + for _, midjourney := range items { + if setting.MjForwardUrlEnabled { midjourney.ImageUrl = system_setting.ServerAddress + "/mj/image/" + midjourney.MjId - items[i] = midjourney } + // 面向用户的 MJ 任务列表同样是上游失败原因的出口,与 /mj/task 查询 + // (coverMidjourneyTaskDto)保持一致;管理员列表 GetAllMidjourney 保留原文。 + midjourney.FailReason = operation_setting.OverrideUpstreamMessage(midjourney.FailReason) } pageInfo.SetTotal(int(total)) pageInfo.SetItems(items) diff --git a/controller/midjourney_fail_reason_override_test.go b/controller/midjourney_fail_reason_override_test.go new file mode 100644 index 000000000000..09297cd33e1b --- /dev/null +++ b/controller/midjourney_fail_reason_override_test.go @@ -0,0 +1,78 @@ +package controller + +import ( + "net/http" + "net/http/httptest" + "testing" + + "github.com/QuantumNous/new-api/common" + "github.com/QuantumNous/new-api/model" + "github.com/QuantumNous/new-api/setting/operation_setting" + + "github.com/gin-gonic/gin" + "github.com/glebarez/sqlite" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "gorm.io/gorm" +) + +// The Midjourney task list is a user-facing outlet for the same upstream failure text +// the /mj/task fetch path already masks (relay.coverMidjourneyTaskDto). The user list +// must mask it while the admin list keeps the original for diagnosis. +func TestMidjourneyListFailReasonOverride(t *testing.T) { + enableTaskErrorOverride(t) + gin.SetMode(gin.TestMode) + + previousDB := model.DB + db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{}) + require.NoError(t, err) + require.NoError(t, db.AutoMigrate(&model.Midjourney{})) + model.DB = db + t.Cleanup(func() { model.DB = previousDB }) + + const userID = 7 + const upstreamReason = "insufficient credits, please top-up your account" + const localReason = "获取渠道信息失败,请联系管理员,渠道ID:3" + + require.NoError(t, db.Create(&model.Midjourney{ + UserId: userID, MjId: "mj-upstream", Status: "FAILURE", FailReason: upstreamReason, + }).Error) + require.NoError(t, db.Create(&model.Midjourney{ + UserId: userID, MjId: "mj-local", Status: "FAILURE", FailReason: localReason, + }).Error) + + failReasons := func(handler gin.HandlerFunc) map[string]string { + t.Helper() + recorder := httptest.NewRecorder() + c, _ := gin.CreateTestContext(recorder) + c.Request = httptest.NewRequest(http.MethodGet, "/api/mj/self", nil) + c.Set("id", userID) + + handler(c) + require.Equal(t, http.StatusOK, recorder.Code) + + var response struct { + Data struct { + Items []struct { + MjId string `json:"mj_id"` + FailReason string `json:"fail_reason"` + } `json:"items"` + } `json:"data"` + } + require.NoError(t, common.Unmarshal(recorder.Body.Bytes(), &response)) + reasons := make(map[string]string, len(response.Data.Items)) + for _, item := range response.Data.Items { + reasons[item.MjId] = item.FailReason + } + return reasons + } + + userReasons := failReasons(GetUserMidjourney) + assert.Equal(t, operation_setting.ErrorOverrideMessage, userReasons["mj-upstream"]) + // This site's own failure text must survive the user-facing path. + assert.Equal(t, localReason, userReasons["mj-local"]) + + adminReasons := failReasons(GetAllMidjourney) + assert.Equal(t, upstreamReason, adminReasons["mj-upstream"]) + assert.Equal(t, localReason, adminReasons["mj-local"]) +} diff --git a/controller/playground.go b/controller/playground.go index 1c9c8d3b77d9..e434c892aade 100644 --- a/controller/playground.go +++ b/controller/playground.go @@ -8,6 +8,7 @@ import ( "github.com/QuantumNous/new-api/model" relaycommon "github.com/QuantumNous/new-api/relay/common" "github.com/QuantumNous/new-api/relaykit/types" + "github.com/QuantumNous/new-api/setting/operation_setting" "github.com/gin-gonic/gin" ) @@ -17,6 +18,7 @@ func Playground(c *gin.Context) { defer func() { if newAPIError != nil { + operation_setting.OverrideUpstreamError(newAPIError) c.JSON(newAPIError.StatusCode, gin.H{ "error": newAPIError.ToOpenAIError(), }) diff --git a/controller/relay.go b/controller/relay.go index e48f60135e55..baa4e0f2ee11 100644 --- a/controller/relay.go +++ b/controller/relay.go @@ -92,7 +92,14 @@ func Relay(c *gin.Context, relayFormat types.RelayFormat) { defer func() { if newAPIError != nil { logger.LogError(c, fmt.Sprintf("relay error: %s", common.LocalLogPreview(newAPIError.Error()))) - newAPIError.SetMessage(common.MessageWithRequestId(newAPIError.Error(), requestId)) + // 必须在错误日志之后执行:日志与渠道禁用判定始终使用原始上游文案。 + // 覆写文案与 request id 一次写入:上游错误的 ToOpenAIError() 直接返回 RelayError, + // 不读 Err,SetMessage 改不到响应体,必须走 ReplaceMessage + if operation_setting.ShouldOverrideUpstreamError(newAPIError) { + newAPIError.ReplaceMessage(common.MessageWithRequestId(operation_setting.ErrorOverrideMessage, requestId)) + } else { + newAPIError.SetMessage(common.MessageWithRequestId(newAPIError.Error(), requestId)) + } switch relayFormat { case types.RelayFormatOpenAIRealtime: helper.WssError(c, ws, newAPIError.ToOpenAIError()) @@ -431,13 +438,26 @@ func processChannelError(c *gin.Context, channelError types.ChannelError, err *t adminInfo["multi_key_index"] = common.GetContextKeyInt(c, constant.ContextKeyChannelMultiKeyIndex) } service.AppendChannelAffinityAdminInfo(c, adminInfo) + // 错误日志的 Content 会通过 /api/log/self 回显给发起请求的用户,因此它和 HTTP 响应 + // 一样是对外出口。覆写生效时这里必须同步覆写,否则客户端在响应里看到 + // Service Unavailable,转头在日志页仍能读到上游账务原文。原文改记到 + // admin_info,model.formatUserLogs 会为普通用户剥离整个 admin_info。 + logContent := err.MaskSensitiveErrorWithStatusCode() + if operation_setting.ShouldOverrideUpstreamError(err) { + adminInfo["original_error"] = logContent + logContent = operation_setting.ErrorOverrideMessage + // 与 MaskSensitiveErrorWithStatusCode 保持一致:无状态码时不写 status_code=0 + if err.StatusCode != 0 { + logContent = fmt.Sprintf("status_code=%d, %s", err.StatusCode, logContent) + } + } other["admin_info"] = adminInfo startTime := common.GetContextKeyTime(c, constant.ContextKeyRequestStartTime) if startTime.IsZero() { startTime = time.Now() } useTimeSeconds := int(time.Since(startTime).Seconds()) - model.RecordErrorLog(c, userId, channelId, modelName, tokenName, err.MaskSensitiveErrorWithStatusCode(), tokenId, useTimeSeconds, common.GetContextKeyBool(c, constant.ContextKeyIsStream), userGroup, other) + model.RecordErrorLog(c, userId, channelId, modelName, tokenName, logContent, tokenId, useTimeSeconds, common.GetContextKeyBool(c, constant.ContextKeyIsStream), userGroup, other) } } @@ -475,13 +495,23 @@ func RelayMidjourney(c *gin.Context) { mjErr.Result = "当前分组负载已饱和,请稍后再试,或升级账户以提升服务质量。" statusCode = http.StatusTooManyRequests } + description := fmt.Sprintf("%s %s", mjErr.Description, mjErr.Result) + channelId := c.GetInt("channel_id") + logger.LogError(c, fmt.Sprintf("relay error (channel #%d, status code %d): %s", channelId, statusCode, description)) + // 这里一律不覆写:能走到 mjErr 的 MidjourneyResponse 全部是本站自产的(参数校验、 + // 额度不足、DB/IO 失败)。mj-proxy 各 handler 拿到上游响应后是把上游 body 原样 + // io.Copy 给客户端再 return nil 的,上游错误文案根本不经过这个分支。 + // 之前在此处按关键词覆写会把本站的 quota_not_enough(mjproxy_handler.go 的 + // RelaySwapFace / RelayMidjourneySubmit)误伤成 Service Unavailable, + // 让用户看不到自己额度不足的真实原因。 + // MJ 上游文案的出口是被代理的响应体本身与任务的 FailReason,前者要改写就得重写 + // 上游 JSON,会破坏 mj-proxy 协议兼容性,故不在本功能范围内;后者已在 + // TaskModel2Dto / coverMidjourneyTaskDto 的读取边界处理。 c.JSON(statusCode, gin.H{ - "description": fmt.Sprintf("%s %s", mjErr.Description, mjErr.Result), + "description": description, "type": "upstream_error", "code": mjErr.Code, }) - channelId := c.GetInt("channel_id") - logger.LogError(c, fmt.Sprintf("relay error (channel #%d, status code %d): %s", channelId, statusCode, fmt.Sprintf("%s %s", mjErr.Description, mjErr.Result))) } } @@ -617,7 +647,7 @@ func RelayTask(c *gin.Context) { processChannelError(c, *types.NewChannelError(channel.Id, channel.Type, channel.Name, channel.ChannelInfo.IsMultiKey, common.GetContextKeyString(c, constant.ContextKeyChannelKey), channel.GetAutoBan()), - types.NewOpenAIError(taskErr.Error, types.ErrorCodeBadResponseStatusCode, taskErr.StatusCode)) + service.APIErrorFromTaskError(taskErr)) } if taskFailoverEnabled { @@ -674,10 +704,24 @@ func respondTaskError(c *gin.Context, taskErr *taskdto.TaskError) { if taskErr.StatusCode == http.StatusTooManyRequests { taskErr.Message = "当前分组上游负载已饱和,请稍后再试" } + // 仅覆写文案确实取自上游任务平台的错误。这里必须用正向标记 FromUpstream 判定: + // 本站自产错误(读 body 失败、解析失败、预扣费额度不足)默认 LocalError == false, + // 用 !LocalError 反推会把它们一并掩盖。 + if taskErr.FromUpstream { + if overridden := operation_setting.OverrideUpstreamMessage(taskErr.Message); overridden != taskErr.Message { + // 覆写前先记录原始上游文案,后台日志与排障始终可见全文 + logger.LogError(c, fmt.Sprintf("task upstream error overridden: %s", common.LocalLogPreview(taskErr.Message))) + taskErr.Message = overridden + // Data 带 json:"data" 会一并返回客户端。当前 task 适配器不往里塞上游 body, + // 但覆写掉 Message 却留着一个可能承载上游原文的字段是自相矛盾的, + // 与 NewAPIError.ReplaceMessage 清 Metadata/Param 保持一致。 + taskErr.Data = nil + } + } c.JSON(taskErr.StatusCode, taskErr) } -func shouldRetryTaskRelay(c *gin.Context, channelId int, taskErr *taskdto.TaskError, retryTimes int) bool { +func shouldRetryTaskRelay(c *gin.Context, channelId int, taskErr *taskdto.TaskError, retryTimes int, failoverEnabled bool) bool { if taskErr == nil { return false } diff --git a/controller/task.go b/controller/task.go index a80f1a687aab..feec5fa42678 100644 --- a/controller/task.go +++ b/controller/task.go @@ -32,7 +32,8 @@ func GetAllTask(c *gin.Context) { items := model.TaskGetAllTasks(pageInfo.GetStartIdx(), pageInfo.GetPageSize(), queryParams) total := model.TaskCountAllTasks(queryParams) pageInfo.SetTotal(int(total)) - pageInfo.SetItems(tasksToDto(items, true)) + // 管理员视图不掩盖上游失败原因,排障需要原文 + pageInfo.SetItems(tasksToDto(items, true, false)) common.ApiSuccess(c, pageInfo) } @@ -56,11 +57,13 @@ func GetUserTask(c *gin.Context) { items := model.TaskGetAllUserTask(userId, pageInfo.GetStartIdx(), pageInfo.GetPageSize(), queryParams) total := model.TaskCountAllUserTask(userId, queryParams) pageInfo.SetTotal(int(total)) - pageInfo.SetItems(tasksToDto(items, false)) + pageInfo.SetItems(tasksToDto(items, false, true)) common.ApiSuccess(c, pageInfo) } -func tasksToDto(tasks []*model.Task, fillUser bool) []*dto.TaskDto { +// tasksToDto 转换任务列表。maskUpstreamFailReason 单独传参而不复用 fillUser:填充用户名 +// 与是否掩盖上游失败原因是两个无关维度,绑在一起会让后续加入新调用方时选错默认值。 +func tasksToDto(tasks []*model.Task, fillUser bool, maskUpstreamFailReason bool) []*dto.TaskDto { var userIdMap map[int]*model.UserBase if fillUser { userIdMap = make(map[int]*model.UserBase) @@ -82,7 +85,7 @@ func tasksToDto(tasks []*model.Task, fillUser bool) []*dto.TaskDto { task.Username = user.Username } } - result[i] = relay.TaskModel2Dto(task) + result[i] = relay.TaskModel2Dto(task, maskUpstreamFailReason) } return result } diff --git a/controller/task_error_override_test.go b/controller/task_error_override_test.go new file mode 100644 index 000000000000..f31a2f8813fa --- /dev/null +++ b/controller/task_error_override_test.go @@ -0,0 +1,152 @@ +package controller + +import ( + "errors" + "net/http" + "net/http/httptest" + "testing" + + taskdto "github.com/QuantumNous/new-api/dto" + "github.com/QuantumNous/new-api/relaykit/types" + "github.com/QuantumNous/new-api/service" + "github.com/QuantumNous/new-api/setting/operation_setting" + + "github.com/gin-gonic/gin" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +// enableTaskErrorOverride turns the override on with the shipped default keywords and +// restores the previous package state afterwards. +func enableTaskErrorOverride(t *testing.T) { + t.Helper() + origEnabled := operation_setting.ErrorOverrideEnabled + origKeywords := operation_setting.ErrorOverrideKeywords + t.Cleanup(func() { + operation_setting.ErrorOverrideEnabled = origEnabled + operation_setting.ErrorOverrideKeywords = origKeywords + }) + operation_setting.ErrorOverrideEnabled = true + operation_setting.ErrorOverrideKeywords = []string{"no available", "quota", "credits", "top-up"} +} + +// respondTaskError must key the override off the positive FromUpstream marker. Local +// errors default to LocalError == false, so gating on !LocalError would mask this +// site's own billing and infrastructure failures whenever their text happens to +// contain an override keyword. +func TestRespondTaskError_OverrideGate(t *testing.T) { + enableTaskErrorOverride(t) + gin.SetMode(gin.TestMode) + + cases := []struct { + name string + taskErr *taskdto.TaskError + wantMessage string + }{ + { + // Regression: PreConsumeBilling produces this locally. It contains "quota", + // so an inverted gate would hide the real reason the request failed. + name: "local subscription quota error is not overridden", + taskErr: service.TaskErrorFromAPIError(types.NewErrorWithStatusCode( + errors.New("订阅额度不足或未配置订阅: subscription quota insufficient"), + types.ErrorCodeInsufficientUserQuota, http.StatusForbidden)), + wantMessage: "订阅额度不足或未配置订阅: subscription quota insufficient", + }, + { + name: "local user quota error is not overridden", + taskErr: service.TaskErrorFromAPIError(types.NewErrorWithStatusCode( + errors.New("用户额度不足, 剩余额度: $0.00"), + types.ErrorCodeInsufficientUserQuota, http.StatusForbidden)), + wantMessage: "用户额度不足, 剩余额度: $0.00", + }, + { + name: "local no-available-key error is not overridden", + taskErr: service.TaskErrorWrapper( + errors.New("no available channel for model under group default"), + "channel_no_available_key", http.StatusServiceUnavailable), + wantMessage: "no available channel for model under group default", + }, + { + name: "upstream keyword match is overridden", + taskErr: service.TaskErrorWrapperUpstream( + errors.New("insufficient credits, please top up your account"), + "upstream_error", http.StatusPaymentRequired), + wantMessage: operation_setting.ErrorOverrideMessage, + }, + { + name: "upstream without keyword keeps its message", + taskErr: service.TaskErrorWrapperUpstream( + errors.New("invalid prompt"), "upstream_error", http.StatusBadRequest), + wantMessage: "invalid prompt", + }, + } + + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + recorder := httptest.NewRecorder() + c, _ := gin.CreateTestContext(recorder) + c.Request = httptest.NewRequest(http.MethodPost, "/v1/video/generations", nil) + + respondTaskError(c, tc.taskErr) + + assert.Equal(t, tc.wantMessage, tc.taskErr.Message) + }) + } +} + +// Data carries json:"data" and reaches the client. Overriding Message while leaving a +// field that may hold the upstream payload would defeat the override, so an overridden +// task error must clear it — and a non-overridden one must keep it. +func TestRespondTaskError_ClearsDataOnOverride(t *testing.T) { + enableTaskErrorOverride(t) + gin.SetMode(gin.TestMode) + + t.Run("overridden error drops Data", func(t *testing.T) { + taskErr := service.TaskErrorWrapperUpstream( + errors.New("insufficient credits, please top up"), + "upstream_error", http.StatusPaymentRequired) + taskErr.Data = map[string]string{"raw": "upstream-billing-detail"} + + recorder := httptest.NewRecorder() + c, _ := gin.CreateTestContext(recorder) + c.Request = httptest.NewRequest(http.MethodPost, "/v1/video/generations", nil) + + respondTaskError(c, taskErr) + + require.Equal(t, operation_setting.ErrorOverrideMessage, taskErr.Message) + assert.Nil(t, taskErr.Data) + }) + + t.Run("untouched error keeps Data", func(t *testing.T) { + taskErr := service.TaskErrorWrapperUpstream( + errors.New("invalid prompt"), "upstream_error", http.StatusBadRequest) + taskErr.Data = map[string]string{"field": "prompt"} + + recorder := httptest.NewRecorder() + c, _ := gin.CreateTestContext(recorder) + c.Request = httptest.NewRequest(http.MethodPost, "/v1/video/generations", nil) + + respondTaskError(c, taskErr) + + require.Equal(t, "invalid prompt", taskErr.Message) + assert.NotNil(t, taskErr.Data) + }) +} + +// A 429 rewrite happens before the override and its text is this site's own, so it +// must survive regardless of the upstream marker. +func TestRespondTaskError_RateLimitMessagePreserved(t *testing.T) { + enableTaskErrorOverride(t) + gin.SetMode(gin.TestMode) + + taskErr := service.TaskErrorWrapperUpstream( + errors.New("upstream quota exhausted"), "upstream_error", http.StatusTooManyRequests) + + recorder := httptest.NewRecorder() + c, _ := gin.CreateTestContext(recorder) + c.Request = httptest.NewRequest(http.MethodPost, "/v1/video/generations", nil) + + respondTaskError(c, taskErr) + + require.Equal(t, "当前分组上游负载已饱和,请稍后再试", taskErr.Message) +} diff --git a/dto/task.go b/dto/task.go index 4a9a8e2e6d18..1fe309b3542d 100644 --- a/dto/task.go +++ b/dto/task.go @@ -10,7 +10,11 @@ type TaskError struct { Data any `json:"data"` StatusCode int `json:"-"` LocalError bool `json:"-"` - Error error `json:"-"` + // FromUpstream 表示 Message 的文案取自上游任务平台的响应体,而非本站生成。 + // 错误信息覆写只作用于该标记为 true 的错误;LocalError 是重试判定用的独立维度, + // 不能反过来当作「来自上游」的依据(本站自产错误默认 LocalError == false)。 + FromUpstream bool `json:"-"` + Error error `json:"-"` } type TaskData interface { diff --git a/middleware/distributor.go b/middleware/distributor.go index 2b8a794776ff..67fa9d365f34 100644 --- a/middleware/distributor.go +++ b/middleware/distributor.go @@ -52,7 +52,7 @@ func Distribute() func(c *gin.Context) { } if aliasApplied { common.SetContextKey(c, constant.ContextKeyRequestedModelAlias, modelRequest.Model) - rewriteRequestBodyModel(c, strings.TrimSuffix(resolvedModel, ratio_setting.CompactModelSuffix)) + rewriteRequestBodyModel(c, resolvedModel) modelRequest.Model = resolvedModel } } diff --git a/model/channel_cache_failover_test.go b/model/channel_cache_failover_test.go index 80e6e311a865..ec2ef51f5e09 100644 --- a/model/channel_cache_failover_test.go +++ b/model/channel_cache_failover_test.go @@ -4,7 +4,7 @@ import ( "testing" "github.com/QuantumNous/new-api/common" - "github.com/QuantumNous/new-api/dto" + relaykitdto "github.com/QuantumNous/new-api/relaykit/dto" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" ) @@ -31,7 +31,7 @@ func setupFailoverChannelCache(t *testing.T, channels []*Channel, group string, oldAdvanced := channel2advancedCustomConfig channelsIDM = newIDM group2model2channels = newG2M - channel2advancedCustomConfig = make(map[int]*dto.AdvancedCustomConfig) + channel2advancedCustomConfig = make(map[int]*relaykitdto.AdvancedCustomConfig) channelSyncLock.Unlock() t.Cleanup(func() { diff --git a/model/option.go b/model/option.go index 89eb56f2a427..f7e81c28b94b 100644 --- a/model/option.go +++ b/model/option.go @@ -175,6 +175,8 @@ func InitOptionMap() { common.OptionMap["SensitiveWords"] = setting.SensitiveWordsToString() common.OptionMap["StreamCacheQueueLength"] = strconv.Itoa(setting.StreamCacheQueueLength) common.OptionMap["AutomaticDisableKeywords"] = operation_setting.AutomaticDisableKeywordsToString() + common.OptionMap["ErrorOverrideEnabled"] = strconv.FormatBool(operation_setting.ErrorOverrideEnabled) + common.OptionMap["ErrorOverrideKeywords"] = operation_setting.ErrorOverrideKeywordsToString() common.OptionMap["AutomaticDisableStatusCodes"] = operation_setting.AutomaticDisableStatusCodesToString() common.OptionMap["AutomaticRetryStatusCodes"] = operation_setting.AutomaticRetryStatusCodesToString() common.OptionMap["ExposeRatioEnabled"] = strconv.FormatBool(ratio_setting.IsExposeRatioEnabled()) @@ -373,6 +375,8 @@ func updateOptionMap(key string, value string) (err error) { operation_setting.DemoSiteEnabled = boolValue case "SelfUseModeEnabled": operation_setting.SelfUseModeEnabled = boolValue + case "ErrorOverrideEnabled": + operation_setting.ErrorOverrideEnabled = boolValue case "ChannelFailoverEnabled": operation_setting.ChannelFailoverEnabled = boolValue case "CheckSensitiveOnPromptEnabled": @@ -593,6 +597,8 @@ func updateOptionMap(key string, value string) (err error) { setting.SensitiveWordsFromString(value) case "AutomaticDisableKeywords": operation_setting.AutomaticDisableKeywordsFromString(value) + case "ErrorOverrideKeywords": + operation_setting.ErrorOverrideKeywordsFromString(value) case "AutomaticDisableStatusCodes": err = operation_setting.AutomaticDisableStatusCodesFromString(value) case "AutomaticRetryStatusCodes": diff --git a/relay/channel/ali/image.go b/relay/channel/ali/image.go index 6913fa346aa6..cf46725e487e 100644 --- a/relay/channel/ali/image.go +++ b/relay/channel/ali/image.go @@ -317,7 +317,7 @@ func aliImageHandler(a *Adaptor, c *gin.Context, resp *http.Response, info *rela return types.NewError(err, types.ErrorCodeBadResponse), nil } if aliResponse.Output.TaskStatus != "SUCCEEDED" { - return types.WithOpenAIError(types.OpenAIError{ + return types.WithUpstreamOpenAIError(types.OpenAIError{ Message: aliResponse.Output.Message, Type: "ali_error", Param: "", diff --git a/relay/channel/ali/rerank.go b/relay/channel/ali/rerank.go index ac2afbd3d3db..7bc3e358a2f6 100644 --- a/relay/channel/ali/rerank.go +++ b/relay/channel/ali/rerank.go @@ -46,7 +46,7 @@ func RerankHandler(c *gin.Context, resp *http.Response, info *relaycommon.RelayI } if aliResponse.Code != "" { - return types.WithOpenAIError(types.OpenAIError{ + return types.WithUpstreamOpenAIError(types.OpenAIError{ Message: aliResponse.Message, Type: aliResponse.Code, Param: aliResponse.RequestId, diff --git a/relay/channel/claude/relay-claude.go b/relay/channel/claude/relay-claude.go index 2f424b32abdd..b3615ab70b33 100644 --- a/relay/channel/claude/relay-claude.go +++ b/relay/channel/claude/relay-claude.go @@ -91,7 +91,7 @@ func HandleStreamResponseData(c *gin.Context, info *relaycommon.RelayInfo, claud return types.NewError(err, types.ErrorCodeBadResponseBody) } if claudeError := claudeResponse.GetClaudeError(); claudeError != nil && claudeError.Type != "" { - return types.WithClaudeError(*claudeError, http.StatusInternalServerError) + return types.WithUpstreamClaudeError(*claudeError, http.StatusInternalServerError) } if claudeResponse.StopReason != "" { maybeMarkClaudeRefusal(c, claudeResponse.StopReason) @@ -221,7 +221,7 @@ func HandleClaudeResponseData(c *gin.Context, info *relaycommon.RelayInfo, claud return types.NewError(err, types.ErrorCodeBadResponseBody) } if claudeError := claudeResponse.GetClaudeError(); claudeError != nil && claudeError.Type != "" { - return types.WithClaudeError(*claudeError, http.StatusInternalServerError) + return types.WithUpstreamClaudeError(*claudeError, http.StatusInternalServerError) } maybeMarkClaudeRefusal(c, claudeResponse.StopReason) if claudeInfo.Usage == nil { diff --git a/relay/channel/jimeng/image.go b/relay/channel/jimeng/image.go index 888531bdd99d..689e5ed1be6d 100644 --- a/relay/channel/jimeng/image.go +++ b/relay/channel/jimeng/image.go @@ -64,7 +64,7 @@ func jimengImageHandler(c *gin.Context, resp *http.Response, info *relaycommon.R // Check if the response indicates an error if jimengResponse.Code != 10000 { - return nil, types.WithOpenAIError(types.OpenAIError{ + return nil, types.WithUpstreamOpenAIError(types.OpenAIError{ Message: jimengResponse.Message, Type: "jimeng_error", Param: "", diff --git a/relay/channel/minimax/image.go b/relay/channel/minimax/image.go index c86cd47dba72..c7349ae28e9a 100644 --- a/relay/channel/minimax/image.go +++ b/relay/channel/minimax/image.go @@ -187,7 +187,7 @@ func miniMaxImageHandler(c *gin.Context, resp *http.Response, info *relaycommon. return nil, types.NewOpenAIError(err, types.ErrorCodeBadResponseBody, http.StatusInternalServerError) } if minimaxResponse.BaseResp.StatusCode != 0 { - return nil, types.WithOpenAIError(types.OpenAIError{ + return nil, types.WithUpstreamOpenAIError(types.OpenAIError{ Message: minimaxResponse.BaseResp.StatusMsg, Type: "minimax_image_error", Code: fmt.Sprintf("%d", minimaxResponse.BaseResp.StatusCode), diff --git a/relay/channel/openai/chat_via_responses.go b/relay/channel/openai/chat_via_responses.go index 25caeb5854bd..a6dfa071481b 100644 --- a/relay/channel/openai/chat_via_responses.go +++ b/relay/channel/openai/chat_via_responses.go @@ -38,7 +38,7 @@ func OaiResponsesToChatHandler(c *gin.Context, info *relaycommon.RelayInfo, resp } if oaiError := responsesResp.GetOpenAIError(); oaiError != nil && oaiError.Type != "" { - return nil, types.WithOpenAIError(*oaiError, resp.StatusCode) + return nil, types.WithUpstreamOpenAIError(*oaiError, resp.StatusCode) } chatResult, err := relayconvert.ConvertResponse(c, info, types.RelayFormatOpenAI, &responsesResp) @@ -124,7 +124,7 @@ func OaiResponsesToChatBufferedStreamHandler(c *gin.Context, info *relaycommon.R case "response.failed", "response.error": if streamResp.Response != nil { if oaiErr := streamResp.Response.GetOpenAIError(); oaiErr != nil && oaiErr.Type != "" { - streamErr = types.WithOpenAIError(*oaiErr, http.StatusInternalServerError) + streamErr = types.WithUpstreamOpenAIError(*oaiErr, http.StatusInternalServerError) break } } @@ -283,7 +283,7 @@ func OaiResponsesToChatStreamHandler(c *gin.Context, info *relaycommon.RelayInfo if streamResp.Type == "response.error" || streamResp.Type == "response.failed" { if streamResp.Response != nil { if oaiErr := streamResp.Response.GetOpenAIError(); oaiErr != nil && oaiErr.Type != "" { - streamErr = types.WithOpenAIError(*oaiErr, http.StatusInternalServerError) + streamErr = types.WithUpstreamOpenAIError(*oaiErr, http.StatusInternalServerError) sr.Stop(streamErr) return } diff --git a/relay/channel/openai/relay-openai.go b/relay/channel/openai/relay-openai.go index 9a0619eb27f5..3f9ceaff70d1 100644 --- a/relay/channel/openai/relay-openai.go +++ b/relay/channel/openai/relay-openai.go @@ -250,7 +250,7 @@ func OpenaiHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Respo } if oaiError := simpleResponse.GetOpenAIError(); oaiError != nil && oaiError.Type != "" { - return nil, types.WithOpenAIError(*oaiError, resp.StatusCode) + return nil, types.WithUpstreamOpenAIError(*oaiError, resp.StatusCode) } for _, choice := range simpleResponse.Choices { diff --git a/relay/channel/openai/relay_image.go b/relay/channel/openai/relay_image.go index 1e6be0dd4cbe..60682cc8cbaf 100644 --- a/relay/channel/openai/relay_image.go +++ b/relay/channel/openai/relay_image.go @@ -46,7 +46,7 @@ func OpenaiImageHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http. } if oaiError := usageResp.GetOpenAIError(); oaiError != nil && oaiError.Type != "" { - return nil, types.WithOpenAIError(*oaiError, resp.StatusCode) + return nil, types.WithUpstreamOpenAIError(*oaiError, resp.StatusCode) } updateOpenAIImageCount(info, gjson.GetBytes(responseBody, "data.#").Int()) @@ -247,7 +247,7 @@ func openaiImageJSONAsStreamHandler(c *gin.Context, info *relaycommon.RelayInfo, return nil, types.NewOpenAIError(err, types.ErrorCodeBadResponseBody, http.StatusInternalServerError) } if oaiError := usageResp.GetOpenAIError(); oaiError != nil && oaiError.Type != "" { - return nil, types.WithOpenAIError(*oaiError, resp.StatusCode) + return nil, types.WithUpstreamOpenAIError(*oaiError, resp.StatusCode) } normalizeOpenAIUsage(&usageResp.Usage) applyUsagePostProcessing(info, &usageResp.Usage, responseBody) diff --git a/relay/channel/openai/relay_responses.go b/relay/channel/openai/relay_responses.go index ceca1af3b381..3fa6f5f7d6a9 100644 --- a/relay/channel/openai/relay_responses.go +++ b/relay/channel/openai/relay_responses.go @@ -31,7 +31,7 @@ func OaiResponsesHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http return nil, types.NewOpenAIError(err, types.ErrorCodeBadResponseBody, http.StatusInternalServerError) } if oaiError := responsesResponse.GetOpenAIError(); oaiError != nil && oaiError.Type != "" { - return nil, types.WithOpenAIError(*oaiError, resp.StatusCode) + return nil, types.WithUpstreamOpenAIError(*oaiError, resp.StatusCode) } // 写入新的 response body diff --git a/relay/channel/openai/relay_responses_compact.go b/relay/channel/openai/relay_responses_compact.go index ff30d36e417b..4bbceb11a514 100644 --- a/relay/channel/openai/relay_responses_compact.go +++ b/relay/channel/openai/relay_responses_compact.go @@ -25,7 +25,7 @@ func OaiResponsesCompactionHandler(c *gin.Context, resp *http.Response) (*dto.Us return nil, types.NewOpenAIError(err, types.ErrorCodeBadResponseBody, http.StatusInternalServerError) } if oaiError := compactResp.GetOpenAIError(); oaiError != nil && oaiError.Type != "" { - return nil, types.WithOpenAIError(*oaiError, resp.StatusCode) + return nil, types.WithUpstreamOpenAIError(*oaiError, resp.StatusCode) } service.IOCopyBytesGracefully(c, resp, responseBody) diff --git a/relay/channel/openai/responses_via_chat.go b/relay/channel/openai/responses_via_chat.go index 53b9d33cbc0e..47880d0f2cea 100644 --- a/relay/channel/openai/responses_via_chat.go +++ b/relay/channel/openai/responses_via_chat.go @@ -32,7 +32,7 @@ func OaiChatToResponsesHandler(c *gin.Context, info *relaycommon.RelayInfo, resp return nil, types.NewOpenAIError(err, types.ErrorCodeBadResponseBody, http.StatusInternalServerError) } if oaiError := chatResp.GetOpenAIError(); oaiError != nil && oaiError.Type != "" { - return nil, types.WithOpenAIError(*oaiError, resp.StatusCode) + return nil, types.WithUpstreamOpenAIError(*oaiError, resp.StatusCode) } if responseID := helper.GetResponseID(c); responseID != "" { @@ -97,7 +97,7 @@ func OaiChatToResponsesStreamHandler(c *gin.Context, info *relaycommon.RelayInfo var errorResp dto.OpenAITextResponse if err := common.UnmarshalJsonStr(data, &errorResp); err == nil { if oaiError := errorResp.GetOpenAIError(); oaiError != nil && oaiError.Type != "" { - streamErr = types.WithOpenAIError(*oaiError, resp.StatusCode) + streamErr = types.WithUpstreamOpenAIError(*oaiError, resp.StatusCode) sr.Stop(streamErr) return } diff --git a/relay/channel/palm/relay-palm.go b/relay/channel/palm/relay-palm.go index 89ac14c7d010..fc8a24dc0f10 100644 --- a/relay/channel/palm/relay-palm.go +++ b/relay/channel/palm/relay-palm.go @@ -113,7 +113,7 @@ func palmHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Respons return nil, types.NewOpenAIError(err, types.ErrorCodeBadResponseBody, http.StatusInternalServerError) } if palmResponse.Error.Code != 0 || len(palmResponse.Candidates) == 0 { - return nil, types.WithOpenAIError(types.OpenAIError{ + return nil, types.WithUpstreamOpenAIError(types.OpenAIError{ Message: palmResponse.Error.Message, Type: palmResponse.Error.Status, Param: "", diff --git a/relay/channel/task/ali/adaptor.go b/relay/channel/task/ali/adaptor.go index 7452c614dbdb..2ec9558ffb70 100644 --- a/relay/channel/task/ali/adaptor.go +++ b/relay/channel/task/ali/adaptor.go @@ -489,13 +489,13 @@ func (a *TaskAdaptor) DoResponse(c *gin.Context, resp *http.Response, info *rela // 解析阿里响应 var aliResp AliVideoResponse if err := common.Unmarshal(responseBody, &aliResp); err != nil { - taskErr = service.TaskErrorWrapper(errors.Wrapf(err, "body: %s", responseBody), "unmarshal_response_body_failed", http.StatusInternalServerError) + taskErr = service.TaskErrorWrapperUpstream(errors.Wrapf(err, "body: %s", responseBody), "unmarshal_response_body_failed", http.StatusInternalServerError) return } // 检查错误 if aliResp.Code != "" { - taskErr = service.TaskErrorWrapper(fmt.Errorf("%s: %s", aliResp.Code, aliResp.Message), "ali_api_error", resp.StatusCode) + taskErr = service.TaskErrorWrapperUpstream(fmt.Errorf("%s: %s", aliResp.Code, aliResp.Message), "ali_api_error", resp.StatusCode) return } diff --git a/relay/channel/task/doubao/adaptor.go b/relay/channel/task/doubao/adaptor.go index 69302a676290..eade387ee2b9 100644 --- a/relay/channel/task/doubao/adaptor.go +++ b/relay/channel/task/doubao/adaptor.go @@ -219,7 +219,7 @@ func (a *TaskAdaptor) DoResponse(c *gin.Context, resp *http.Response, info *rela // Parse Doubao response var dResp responsePayload if err := common.Unmarshal(responseBody, &dResp); err != nil { - taskErr = service.TaskErrorWrapper(errors.Wrapf(err, "body: %s", responseBody), "unmarshal_response_body_failed", http.StatusInternalServerError) + taskErr = service.TaskErrorWrapperUpstream(errors.Wrapf(err, "body: %s", responseBody), "unmarshal_response_body_failed", http.StatusInternalServerError) return } diff --git a/relay/channel/task/hailuo/adaptor.go b/relay/channel/task/hailuo/adaptor.go index af9f5c57c52a..899fd0c1324c 100644 --- a/relay/channel/task/hailuo/adaptor.go +++ b/relay/channel/task/hailuo/adaptor.go @@ -89,12 +89,12 @@ func (a *TaskAdaptor) DoResponse(c *gin.Context, resp *http.Response, info *rela var hResp VideoResponse if err := common.Unmarshal(responseBody, &hResp); err != nil { - taskErr = service.TaskErrorWrapper(errors.Wrapf(err, "body: %s", responseBody), "unmarshal_response_body_failed", http.StatusInternalServerError) + taskErr = service.TaskErrorWrapperUpstream(errors.Wrapf(err, "body: %s", responseBody), "unmarshal_response_body_failed", http.StatusInternalServerError) return } if hResp.BaseResp.StatusCode != StatusSuccess { - taskErr = service.TaskErrorWrapper( + taskErr = service.TaskErrorWrapperUpstream( fmt.Errorf("hailuo api error: %s", hResp.BaseResp.StatusMsg), strconv.Itoa(hResp.BaseResp.StatusCode), http.StatusBadRequest, diff --git a/relay/channel/task/jimeng/adaptor.go b/relay/channel/task/jimeng/adaptor.go index 5e788d3415b5..197fcee1cd36 100644 --- a/relay/channel/task/jimeng/adaptor.go +++ b/relay/channel/task/jimeng/adaptor.go @@ -194,12 +194,12 @@ func (a *TaskAdaptor) DoResponse(c *gin.Context, resp *http.Response, info *rela // Parse Jimeng response var jResp responsePayload if err := common.Unmarshal(responseBody, &jResp); err != nil { - taskErr = service.TaskErrorWrapper(errors.Wrapf(err, "body: %s", responseBody), "unmarshal_response_body_failed", http.StatusInternalServerError) + taskErr = service.TaskErrorWrapperUpstream(errors.Wrapf(err, "body: %s", responseBody), "unmarshal_response_body_failed", http.StatusInternalServerError) return } if jResp.Code != 10000 { - taskErr = service.TaskErrorWrapper(fmt.Errorf("%s", jResp.Message), fmt.Sprintf("%d", jResp.Code), http.StatusInternalServerError) + taskErr = service.TaskErrorWrapperUpstream(fmt.Errorf("%s", jResp.Message), fmt.Sprintf("%d", jResp.Code), http.StatusInternalServerError) return } diff --git a/relay/channel/task/kling/adaptor.go b/relay/channel/task/kling/adaptor.go index 200c3c6829ee..41c88671f7aa 100644 --- a/relay/channel/task/kling/adaptor.go +++ b/relay/channel/task/kling/adaptor.go @@ -203,7 +203,9 @@ func (a *TaskAdaptor) DoResponse(c *gin.Context, resp *http.Response, info *rela return } if kResp.Code != 0 { + // LocalError 保持 true:该错误不重试。文案取自上游,需同时标记以便错误信息覆写生效。 taskErr = service.TaskErrorWrapperLocal(fmt.Errorf("%s", kResp.Message), "task_failed", http.StatusBadRequest) + taskErr.FromUpstream = true return } ov := dto.NewOpenAIVideo() diff --git a/relay/channel/task/sora/adaptor.go b/relay/channel/task/sora/adaptor.go index 7f81e5335ebb..9b6e8d4a42d8 100644 --- a/relay/channel/task/sora/adaptor.go +++ b/relay/channel/task/sora/adaptor.go @@ -236,7 +236,7 @@ func (a *TaskAdaptor) DoResponse(c *gin.Context, resp *http.Response, info *rela // Parse Sora response var dResp responseTask if err := common.Unmarshal(responseBody, &dResp); err != nil { - taskErr = service.TaskErrorWrapper(errors.Wrapf(err, "body: %s", responseBody), "unmarshal_response_body_failed", http.StatusInternalServerError) + taskErr = service.TaskErrorWrapperUpstream(errors.Wrapf(err, "body: %s", responseBody), "unmarshal_response_body_failed", http.StatusInternalServerError) return } diff --git a/relay/channel/task/suno/adaptor.go b/relay/channel/task/suno/adaptor.go index 35b5e423b7ff..14e0b49b7314 100644 --- a/relay/channel/task/suno/adaptor.go +++ b/relay/channel/task/suno/adaptor.go @@ -105,7 +105,7 @@ func (a *TaskAdaptor) DoResponse(c *gin.Context, resp *http.Response, info *rela return } if !sunoResponse.IsSuccess() { - taskErr = service.TaskErrorWrapper(fmt.Errorf("%s", sunoResponse.Message), sunoResponse.Code, http.StatusInternalServerError) + taskErr = service.TaskErrorWrapperUpstream(fmt.Errorf("%s", sunoResponse.Message), sunoResponse.Code, http.StatusInternalServerError) return } diff --git a/relay/channel/task/vidu/adaptor.go b/relay/channel/task/vidu/adaptor.go index 62e029bbe88e..dd03dc0d69df 100644 --- a/relay/channel/task/vidu/adaptor.go +++ b/relay/channel/task/vidu/adaptor.go @@ -172,7 +172,7 @@ func (a *TaskAdaptor) DoResponse(c *gin.Context, resp *http.Response, info *rela var vResp responsePayload err = common.Unmarshal(responseBody, &vResp) if err != nil { - taskErr = service.TaskErrorWrapper(errors.Wrap(err, fmt.Sprintf("%s", responseBody)), "unmarshal_response_failed", http.StatusInternalServerError) + taskErr = service.TaskErrorWrapperUpstream(errors.Wrap(err, fmt.Sprintf("%s", responseBody)), "unmarshal_response_failed", http.StatusInternalServerError) return } diff --git a/relay/channel/tencent/relay-tencent.go b/relay/channel/tencent/relay-tencent.go index 611c186defa7..34e39d1fefaa 100644 --- a/relay/channel/tencent/relay-tencent.go +++ b/relay/channel/tencent/relay-tencent.go @@ -145,7 +145,7 @@ func tencentHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Resp return nil, types.NewOpenAIError(err, types.ErrorCodeBadResponseBody, http.StatusInternalServerError) } if tencentSb.Response.Error.Code != 0 { - return nil, types.WithOpenAIError(types.OpenAIError{ + return nil, types.WithUpstreamOpenAIError(types.OpenAIError{ Message: tencentSb.Response.Error.Message, Code: tencentSb.Response.Error.Code, }, resp.StatusCode) diff --git a/relay/channel/zhipu/relay-zhipu.go b/relay/channel/zhipu/relay-zhipu.go index 0c280e2b705d..fae2f2f2332a 100644 --- a/relay/channel/zhipu/relay-zhipu.go +++ b/relay/channel/zhipu/relay-zhipu.go @@ -234,7 +234,7 @@ func zhipuHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Respon return nil, types.NewOpenAIError(err, types.ErrorCodeBadResponseBody, http.StatusInternalServerError) } if !zhipuResponse.Success { - return nil, types.WithOpenAIError(types.OpenAIError{ + return nil, types.WithUpstreamOpenAIError(types.OpenAIError{ Message: zhipuResponse.Msg, Code: zhipuResponse.Code, }, resp.StatusCode) diff --git a/relay/channel/zhipu_4v/image.go b/relay/channel/zhipu_4v/image.go index cdb35a82c8ad..bd97804c5dfd 100644 --- a/relay/channel/zhipu_4v/image.go +++ b/relay/channel/zhipu_4v/image.go @@ -67,7 +67,7 @@ func zhipu4vImageHandler(c *gin.Context, resp *http.Response, info *relaycommon. } if zhipuResp.Error != nil && zhipuResp.Error.Message != "" { - return nil, types.WithOpenAIError(types.OpenAIError{ + return nil, types.WithUpstreamOpenAIError(types.OpenAIError{ Message: zhipuResp.Error.Message, Type: "zhipu_image_error", Code: zhipuResp.Error.Code, diff --git a/relay/mjproxy_handler.go b/relay/mjproxy_handler.go index daa0585a0797..f64dc9e98df3 100644 --- a/relay/mjproxy_handler.go +++ b/relay/mjproxy_handler.go @@ -21,6 +21,7 @@ import ( "github.com/QuantumNous/new-api/relay/helper" "github.com/QuantumNous/new-api/service" "github.com/QuantumNous/new-api/setting" + "github.com/QuantumNous/new-api/setting/operation_setting" "github.com/QuantumNous/new-api/setting/system_setting" "github.com/gin-gonic/gin" @@ -161,7 +162,9 @@ func coverMidjourneyTaskDto(c *gin.Context, originTask *model.Midjourney) (midjo midjourneyTask.VideoUrl = originTask.VideoUrl } midjourneyTask.Status = originTask.Status - midjourneyTask.FailReason = originTask.FailReason + // FailReason 落库时保留上游原文(排障与管理员视图需要),在这个面向用户的读取边界 + // 应用覆写。coverMidjourneyTaskDto 只服务 /mj/task/... 查询,没有管理员调用方。 + midjourneyTask.FailReason = operation_setting.OverrideUpstreamMessage(originTask.FailReason) midjourneyTask.Action = originTask.Action midjourneyTask.Description = originTask.Description midjourneyTask.Prompt = originTask.Prompt diff --git a/relay/relay_task.go b/relay/relay_task.go index fb384d18937a..ebd9f0a95892 100644 --- a/relay/relay_task.go +++ b/relay/relay_task.go @@ -19,7 +19,10 @@ import ( relayconstant "github.com/QuantumNous/new-api/relay/constant" "github.com/QuantumNous/new-api/relay/helper" "github.com/QuantumNous/new-api/service" + "github.com/QuantumNous/new-api/setting/operation_setting" "github.com/gin-gonic/gin" + "github.com/tidwall/gjson" + "github.com/tidwall/sjson" ) type TaskSubmitResult struct { @@ -223,7 +226,7 @@ func RelayTaskSubmit(c *gin.Context, info *relaycommon.RelayInfo) (*TaskSubmitRe } if resp != nil && resp.StatusCode != http.StatusOK { responseBody, _ := io.ReadAll(resp.Body) - return nil, service.TaskErrorWrapper(fmt.Errorf("%s", string(responseBody)), "fail_to_fetch_task", resp.StatusCode) + return nil, service.TaskErrorWrapperUpstream(fmt.Errorf("%s", string(responseBody)), "fail_to_fetch_task", resp.StatusCode) } // 10. 返回 OtherRatios 给下游(header 必须在 DoResponse 写 body 之前设置) @@ -335,7 +338,7 @@ func sunoFetchRespBodyBuilder(c *gin.Context) (respBody []byte, taskResp *dto.Ta return } for _, task := range taskModels { - tasks = append(tasks, TaskModel2Dto(task)) + tasks = append(tasks, TaskModel2Dto(task, true)) } } else { tasks = make([]any, 0) @@ -363,7 +366,7 @@ func sunoFetchByIDRespBodyBuilder(c *gin.Context) (respBody []byte, taskResp *dt respBody, err = common.Marshal(dto.TaskResponse[any]{ Code: "success", - Data: TaskModel2Dto(originTask), + Data: TaskModel2Dto(originTask, true), }) return } @@ -406,7 +409,7 @@ func videoFetchByIDRespBodyBuilder(c *gin.Context) (respBody []byte, taskResp *d taskResp = service.TaskErrorWrapper(err, "convert_to_openai_video_failed", http.StatusInternalServerError) return } - respBody = openAIVideoData + respBody = overrideOpenAIVideoUpstreamError(openAIVideoData) return } taskResp = service.TaskErrorWrapperLocal(fmt.Errorf("not_implemented:%s", originTask.Platform), "not_implemented", http.StatusNotImplemented) @@ -416,7 +419,7 @@ func videoFetchByIDRespBodyBuilder(c *gin.Context) (respBody []byte, taskResp *d // 通用 TaskDto 格式 respBody, err = common.Marshal(dto.TaskResponse[any]{ Code: "success", - Data: TaskModel2Dto(originTask), + Data: TaskModel2Dto(originTask, true), }) if err != nil { taskResp = service.TaskErrorWrapper(err, "marshal_response_failed", http.StatusInternalServerError) @@ -547,7 +550,44 @@ func mapTaskStatusToSimple(status model.TaskStatus) string { } } -func TaskModel2Dto(task *model.Task) *dto.TaskDto { +// overrideOpenAIVideoUpstreamError 覆写 OpenAI Video 响应体里的上游错误文案。 +// /v1/videos/{id} 走各 adaptor 的 ConvertToOpenAIVideo,响应体由落库的上游原始数据构建 +// (service/task_polling.go 的 task.Data),不经过 TaskModel2Dto,是上游失败文案的另一个 +// 用户可见出口。 +// +// 直接改写 JSON 而不是反序列化成 dto.OpenAIVideo 再序列化:sora 的 converter 把上游对象整体 +// 透传,结构体往返会丢掉上游多返回的字段。与 relay 链路一致,只替换 message,error.code 保留。 +func overrideOpenAIVideoUpstreamError(respBody []byte) []byte { + message := gjson.GetBytes(respBody, "error.message") + if message.Type != gjson.String || message.String() == "" { + return respBody + } + overridden := operation_setting.OverrideUpstreamMessage(message.String()) + if overridden == message.String() { + return respBody + } + masked, err := sjson.SetBytes(respBody, "error.message", overridden) + if err != nil { + // 改写失败时宁可不返回上游原文 + return []byte(fmt.Sprintf(`{"error":{"message":%q}}`, overridden)) + } + return masked +} + +// TaskModel2Dto 把任务模型转为对外 DTO。maskUpstreamFailReason 为 true 时对 FailReason +// 应用错误信息覆写:异步任务把上游失败原因落库(service/task_polling.go 的 +// task.FailReason = taskResult.Reason),用户随后通过任务查询接口读到它,这是同一份上游 +// 文案的另一个出口。管理员视图传 false,始终看原文。 +// +// 与 relay/task 链路不同,这里只能按关键词门控:Task 表没有来源标记列,补一列要跨三种 +// 数据库做迁移,代价与收益不匹配。本站自产的 FailReason 是「任务超时(%d分钟)」这类 +// 中文文案(sweepTimedOutTasks)和 upstream returned error 这类固定串(FailTaskInfo), +// 都不含默认关键词;但管理员配置过宽的关键词时本站文案仍可能被误伤。 +func TaskModel2Dto(task *model.Task, maskUpstreamFailReason bool) *dto.TaskDto { + failReason := task.FailReason + if maskUpstreamFailReason { + failReason = operation_setting.OverrideUpstreamMessage(failReason) + } return &dto.TaskDto{ ID: task.ID, CreatedAt: task.CreatedAt, @@ -560,7 +600,7 @@ func TaskModel2Dto(task *model.Task) *dto.TaskDto { Quota: task.Quota, Action: task.Action, Status: string(task.Status), - FailReason: task.FailReason, + FailReason: failReason, ResultURL: task.GetResultURL(), SubmitTime: task.SubmitTime, StartTime: task.StartTime, diff --git a/relay/task_fail_reason_override_test.go b/relay/task_fail_reason_override_test.go new file mode 100644 index 000000000000..0af1b7adf10e --- /dev/null +++ b/relay/task_fail_reason_override_test.go @@ -0,0 +1,119 @@ +package relay + +import ( + "testing" + + "github.com/QuantumNous/new-api/model" + "github.com/QuantumNous/new-api/setting/operation_setting" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "github.com/tidwall/gjson" +) + +// Async tasks persist the upstream failure reason (service/task_polling.go sets +// task.FailReason = taskResult.Reason) and users read it back through the task query +// endpoints. That is a second outlet for the same upstream text the relay response +// override hides, so TaskModel2Dto must mask it for user-facing callers while the admin +// view keeps the original for diagnosis. +func TestTaskModel2Dto_FailReasonOverride(t *testing.T) { + origEnabled := operation_setting.ErrorOverrideEnabled + origKeywords := operation_setting.ErrorOverrideKeywords + t.Cleanup(func() { + operation_setting.ErrorOverrideEnabled = origEnabled + operation_setting.ErrorOverrideKeywords = origKeywords + }) + operation_setting.ErrorOverrideEnabled = true + operation_setting.ErrorOverrideKeywords = []string{"no available", "quota", "credits", "top-up"} + + upstreamReason := "insufficient credits, please top up your account" + + t.Run("user view masks an upstream reason", func(t *testing.T) { + task := &model.Task{TaskID: "t-1", FailReason: upstreamReason} + assert.Equal(t, operation_setting.ErrorOverrideMessage, + TaskModel2Dto(task, true).FailReason) + }) + + t.Run("admin view keeps the original reason", func(t *testing.T) { + task := &model.Task{TaskID: "t-1", FailReason: upstreamReason} + assert.Equal(t, upstreamReason, TaskModel2Dto(task, false).FailReason) + }) + + // Local reasons produced by this site must survive the user-facing path. These are + // the strings sweepTimedOutTasks and FailTaskInfo actually write. + t.Run("local reasons are not masked", func(t *testing.T) { + for _, reason := range []string{ + "任务超时(30分钟)", + "任务超时(旧系统遗留任务,不进行退款,请联系管理员)", + "upstream returned error", + "upstream returned unrecognized message", + } { + task := &model.Task{TaskID: "t-2", FailReason: reason} + assert.Equal(t, reason, TaskModel2Dto(task, true).FailReason, reason) + } + }) + + t.Run("empty reason stays empty", func(t *testing.T) { + task := &model.Task{TaskID: "t-3"} + require.Empty(t, TaskModel2Dto(task, true).FailReason) + }) + + t.Run("disabled override keeps the upstream reason", func(t *testing.T) { + operation_setting.ErrorOverrideEnabled = false + defer func() { operation_setting.ErrorOverrideEnabled = true }() + + task := &model.Task{TaskID: "t-4", FailReason: upstreamReason} + assert.Equal(t, upstreamReason, TaskModel2Dto(task, true).FailReason) + }) +} + +// /v1/videos/{id} builds its response from the stored upstream payload through each +// adaptor's ConvertToOpenAIVideo, bypassing TaskModel2Dto entirely. It is therefore a +// third outlet for the same upstream failure text and must be masked at the JSON level: +// the sora converter passes the whole upstream object through, so a struct round-trip +// would silently drop fields the upstream returned. +func TestOverrideOpenAIVideoUpstreamError(t *testing.T) { + origEnabled := operation_setting.ErrorOverrideEnabled + origKeywords := operation_setting.ErrorOverrideKeywords + t.Cleanup(func() { + operation_setting.ErrorOverrideEnabled = origEnabled + operation_setting.ErrorOverrideKeywords = origKeywords + }) + operation_setting.ErrorOverrideEnabled = true + operation_setting.ErrorOverrideKeywords = []string{"no available", "quota", "credits", "top-up"} + + t.Run("masks the message and keeps every other upstream field", func(t *testing.T) { + body := []byte(`{"id":"video_1","status":"failed","seconds":"8","error":{"code":"billing_hard_limit_reached","message":"You have insufficient credits, please top-up","upstream_only":"keep me"}}`) + + got := overrideOpenAIVideoUpstreamError(body) + + assert.Equal(t, operation_setting.ErrorOverrideMessage, gjson.GetBytes(got, "error.message").String()) + // Status code equivalents and any extra upstream fields survive the rewrite. + assert.Equal(t, "billing_hard_limit_reached", gjson.GetBytes(got, "error.code").String()) + assert.Equal(t, "keep me", gjson.GetBytes(got, "error.upstream_only").String()) + assert.Equal(t, "video_1", gjson.GetBytes(got, "id").String()) + assert.Equal(t, "failed", gjson.GetBytes(got, "status").String()) + assert.Equal(t, "8", gjson.GetBytes(got, "seconds").String()) + }) + + t.Run("leaves bodies without a matching message byte-identical", func(t *testing.T) { + for _, body := range []string{ + `{"id":"video_2","status":"completed"}`, + `{"id":"video_3","status":"failed","error":{"code":"moderation_blocked","message":"content policy violation"}}`, + `{"id":"video_4","error":null}`, + `{"id":"video_5","error":"insufficient credits"}`, + `not json at all, please top-up`, + ``, + } { + assert.Equal(t, body, string(overrideOpenAIVideoUpstreamError([]byte(body))), body) + } + }) + + t.Run("disabled override keeps the upstream message", func(t *testing.T) { + operation_setting.ErrorOverrideEnabled = false + defer func() { operation_setting.ErrorOverrideEnabled = true }() + + body := `{"error":{"message":"You have insufficient credits, please top-up"}}` + assert.Equal(t, body, string(overrideOpenAIVideoUpstreamError([]byte(body)))) + }) +} diff --git a/relaykit/types/error.go b/relaykit/types/error.go index 387fdad76948..0691143e7108 100644 --- a/relaykit/types/error.go +++ b/relaykit/types/error.go @@ -92,6 +92,7 @@ type NewAPIError struct { RelayError any skipRetry bool recordErrorLog *bool + fromUpstream bool errorType ErrorType errorCode ErrorCode StatusCode int @@ -177,6 +178,44 @@ func (e *NewAPIError) SetMessage(message string) { e.Err = errors.New(message) } +// MarkUpstreamOrigin records that this error's message text came from an upstream +// provider response rather than being produced locally by the gateway. Callers that +// decode an upstream error body must set this so the message can be treated as +// upstream-controlled text downstream. +func (e *NewAPIError) MarkUpstreamOrigin() { + if e == nil { + return + } + e.fromUpstream = true +} + +// ReplaceMessage rewrites the user-facing message on every representation of the +// error while preserving StatusCode, type and code. Unlike SetMessage it also +// updates RelayError, because ToOpenAIError/ToClaudeError return the provider +// payload verbatim for upstream errors and never read Err. +// +// Metadata and Param are cleared alongside the message: metadata carries the raw +// provider error body, and Param carries provider-supplied detail such as an +// upstream request id (see the Ali rerank adaptor). Leaving either in place would +// keep leaking the upstream trail that replacing the message is meant to hide. +func (e *NewAPIError) ReplaceMessage(message string) { + if e == nil { + return + } + e.Err = errors.New(message) + e.Metadata = nil + switch relayError := e.RelayError.(type) { + case OpenAIError: + relayError.Message = message + relayError.Metadata = nil + relayError.Param = "" + e.RelayError = relayError + case ClaudeError: + relayError.Message = message + e.RelayError = relayError + } +} + func (e *NewAPIError) ToOpenAIError() OpenAIError { var result OpenAIError switch e.errorType { @@ -363,6 +402,23 @@ func WithClaudeError(claudeError ClaudeError, statusCode int, ops ...NewAPIError return e } +// WithUpstreamOpenAIError builds an error from an OpenAI-shaped payload decoded out +// of an upstream provider response, recording that the message text is upstream +// controlled. Channel adaptors that surface a provider error must use this instead of +// WithOpenAIError, which is also used for locally constructed errors. +func WithUpstreamOpenAIError(openAIError OpenAIError, statusCode int, ops ...NewAPIErrorOptions) *NewAPIError { + e := WithOpenAIError(openAIError, statusCode, ops...) + e.MarkUpstreamOrigin() + return e +} + +// WithUpstreamClaudeError is the Claude-shaped counterpart of WithUpstreamOpenAIError. +func WithUpstreamClaudeError(claudeError ClaudeError, statusCode int, ops ...NewAPIErrorOptions) *NewAPIError { + e := WithClaudeError(claudeError, statusCode, ops...) + e.MarkUpstreamOrigin() + return e +} + func IsChannelError(err *NewAPIError) bool { if err == nil { return false @@ -378,6 +434,16 @@ func IsSkipRetryError(err *NewAPIError) bool { return err.skipRetry } +// IsFromUpstreamError reports whether the error message originated from an upstream +// provider response. Errors produced by the gateway itself return false. +func IsFromUpstreamError(err *NewAPIError) bool { + if err == nil { + return false + } + + return err.fromUpstream +} + func ErrOptionWithSkipRetry() NewAPIErrorOptions { return func(e *NewAPIError) { e.skipRetry = true diff --git a/relaykit/types/error_test.go b/relaykit/types/error_test.go new file mode 100644 index 000000000000..87749dcc978c --- /dev/null +++ b/relaykit/types/error_test.go @@ -0,0 +1,143 @@ +package types + +import ( + "encoding/json" + "errors" + "net/http" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func TestMarkUpstreamOrigin_NilSafe(t *testing.T) { + var e *NewAPIError + require.NotPanics(t, func() { e.MarkUpstreamOrigin() }) + assert.False(t, IsFromUpstreamError(e)) +} + +func TestReplaceMessage_NilSafe(t *testing.T) { + var e *NewAPIError + require.NotPanics(t, func() { e.ReplaceMessage("x") }) +} + +func TestIsFromUpstreamError_Nil(t *testing.T) { + assert.False(t, IsFromUpstreamError(nil)) +} + +func TestWithUpstreamOpenAIError_MarksOrigin(t *testing.T) { + e := WithUpstreamOpenAIError(OpenAIError{ + Message: "Your credit balance is too low, please top up", + Type: "upstream_error", + Code: "credits_exhausted", + }, http.StatusPaymentRequired) + + require.True(t, IsFromUpstreamError(e)) + assert.Equal(t, http.StatusPaymentRequired, e.StatusCode) + assert.Equal(t, ErrorTypeOpenAIError, e.GetErrorType()) + assert.Equal(t, ErrorCode("credits_exhausted"), e.GetErrorCode()) +} + +func TestWithOpenAIError_DoesNotMarkOrigin(t *testing.T) { + e := WithOpenAIError(OpenAIError{ + Message: "some local message", + Type: "upstream_error", + }, http.StatusBadRequest) + assert.False(t, IsFromUpstreamError(e)) +} + +func TestWithUpstreamClaudeError_MarksOrigin(t *testing.T) { + e := WithUpstreamClaudeError(ClaudeError{ + Type: "error", + Message: "You exceeded your current quota", + }, http.StatusTooManyRequests) + + require.True(t, IsFromUpstreamError(e)) + assert.Equal(t, ErrorTypeClaudeError, e.GetErrorType()) + assert.Equal(t, http.StatusTooManyRequests, e.StatusCode) +} + +func TestWithClaudeError_DoesNotMarkOrigin(t *testing.T) { + e := WithClaudeError(ClaudeError{ + Type: "error", + Message: "local", + }, http.StatusBadRequest) + assert.False(t, IsFromUpstreamError(e)) +} + +// ReplaceMessage must rewrite every user-facing representation while preserving +// StatusCode, type, code, and clearing upstream metadata. +func TestReplaceMessage_OpenAIError(t *testing.T) { + metadata := json.RawMessage(`{"raw":"upstream-secret"}`) + e := WithUpstreamOpenAIError(OpenAIError{ + Message: "No available channel for model gpt-4o under group default", + Type: "upstream_error", + Code: "no_available_channel", + Param: "upstream-request-id-123", + Metadata: metadata, + }, http.StatusServiceUnavailable) + e.Metadata = metadata + + e.ReplaceMessage("Service Unavailable") + + // Err mirrors the new message. + assert.Equal(t, "Service Unavailable", e.Error()) + + // ToOpenAIError returns the RelayError payload verbatim, so it must reflect the new message. + openAI := e.ToOpenAIError() + assert.Equal(t, "Service Unavailable", openAI.Message) + assert.Nil(t, openAI.Metadata) + // Param carries provider detail such as an upstream request id and must not survive. + assert.Equal(t, "", openAI.Param) + + // Preserved fields. + assert.Equal(t, http.StatusServiceUnavailable, e.StatusCode) + assert.Equal(t, ErrorTypeOpenAIError, e.GetErrorType()) + assert.Equal(t, ErrorCode("no_available_channel"), e.GetErrorCode()) + assert.Nil(t, e.Metadata) +} + +func TestReplaceMessage_ClaudeError(t *testing.T) { + e := WithUpstreamClaudeError(ClaudeError{ + Type: "error", + Message: "You exceeded your current quota", + }, http.StatusTooManyRequests) + + e.ReplaceMessage("Service Unavailable") + + assert.Equal(t, "Service Unavailable", e.Error()) + claude := e.ToClaudeError() + assert.Equal(t, "Service Unavailable", claude.Message) + assert.Equal(t, "error", claude.Type) + assert.Equal(t, http.StatusTooManyRequests, e.StatusCode) + assert.Equal(t, ErrorTypeClaudeError, e.GetErrorType()) +} + +func TestReplaceMessage_DefaultErrorType(t *testing.T) { + e := NewErrorWithStatusCode(errors.New("upstream quota exhausted"), + ErrorCodeBadResponseStatusCode, http.StatusBadGateway) + e.MarkUpstreamOrigin() + require.True(t, IsFromUpstreamError(e)) + + e.ReplaceMessage("Service Unavailable") + + assert.Equal(t, "Service Unavailable", e.Error()) + openAI := e.ToOpenAIError() + assert.Equal(t, "Service Unavailable", openAI.Message) + assert.Equal(t, http.StatusBadGateway, e.StatusCode) + assert.Equal(t, ErrorTypeNewAPIError, e.GetErrorType()) +} + +// NewError preserves a deeply wrapped *NewAPIError, so the upstream marker must +// survive wrapping carried by errors.As. +func TestUpstreamMarker_SurvivesNewErrorWrap(t *testing.T) { + upstream := WithUpstreamOpenAIError(OpenAIError{ + Message: "Please top up your credits", + Type: "upstream_error", + }, http.StatusPaymentRequired) + + wrapped := NewError(upstream, ErrorCodeBadResponseStatusCode) + + require.True(t, IsFromUpstreamError(wrapped)) + assert.Equal(t, http.StatusPaymentRequired, wrapped.StatusCode) +} diff --git a/service/error.go b/service/error.go index f14f1bbab660..1f21155994f2 100644 --- a/service/error.go +++ b/service/error.go @@ -89,8 +89,10 @@ func RelayErrorHandler(ctx context.Context, resp *http.Response, showBodyWhenFai responseBody, err := io.ReadAll(resp.Body) if err != nil { + // 读取上游响应体失败是本站基础设施错误,文案并非来自上游,不应标记上游来源 return } + CloseResponseBodyGracefully(resp) var errResponse dto.GeneralErrorResponse responseBodyText := string(responseBody) @@ -105,8 +107,12 @@ func RelayErrorHandler(ctx context.Context, resp *http.Response, showBodyWhenFai err = common.Unmarshal(responseBody, &errResponse) if err != nil { if showBodyWhenFail { + // 文案内嵌上游响应体原文,属于上游来源 newApiErr.Err = buildErrWithBody("") + newApiErr.MarkUpstreamOrigin() } else { + // 响应体解析失败且不回显 body 时,文案完全由本站生成(只含状态码), + // 不标记上游来源,避免覆写关键词误伤本站文案 logger.LogError(ctx, fmt.Sprintf("bad response status code %d, body: %s", resp.StatusCode, responseBodyPreview)) newApiErr.Err = fmt.Errorf("bad response status code %d", resp.StatusCode) } @@ -117,7 +123,7 @@ func RelayErrorHandler(ctx context.Context, resp *http.Response, showBodyWhenFai // General format error (OpenAI, Anthropic, Gemini, etc.) oaiError := errResponse.TryToOpenAIError() if oaiError != nil { - newApiErr = types.WithOpenAIError(*oaiError, resp.StatusCode) + newApiErr = types.WithUpstreamOpenAIError(*oaiError, resp.StatusCode) if showBodyWhenFail { newApiErr.Err = buildErrWithBody(newApiErr.Error()) } @@ -131,6 +137,8 @@ func RelayErrorHandler(ctx context.Context, resp *http.Response, showBodyWhenFai logger.LogError(ctx, fmt.Sprintf("bad response status code %d with empty error message, body: %s", resp.StatusCode, responseBodyPreview)) } newApiErr = types.NewOpenAIError(errors.New(message), types.ErrorCodeBadResponseStatusCode, resp.StatusCode) + // message 取自上游响应体(errResponse.ToMessage()),标记为上游来源 + newApiErr.MarkUpstreamOrigin() if showBodyWhenFail { newApiErr.Err = buildErrWithBody(newApiErr.Error()) } @@ -197,6 +205,15 @@ func TaskErrorWrapperLocal(err error, code string, statusCode int) *taskdto.Task return openaiErr } +// TaskErrorWrapperUpstream 包装文案取自上游任务平台响应体的错误,标记后可被错误信息覆写识别。 +// 适配器把上游返回的 message/code 直接塞进错误时必须用它,而不是 TaskErrorWrapper —— 后者 +// 同样服务于本站自产的读取/解析失败。 +func TaskErrorWrapperUpstream(err error, code string, statusCode int) *taskdto.TaskError { + taskErr := TaskErrorWrapper(err, code, statusCode) + taskErr.FromUpstream = true + return taskErr +} + func TaskErrorWrapper(err error, code string, statusCode int) *taskdto.TaskError { text := err.Error() lowerText := strings.ToLower(text) @@ -217,14 +234,38 @@ func TaskErrorWrapper(err error, code string, statusCode int) *taskdto.TaskError } // TaskErrorFromAPIError 将 PreConsumeBilling 返回的 NewAPIError 转换为 TaskError。 +// 预扣费错误(额度不足、订阅未配置等)由本站产生,FromUpstream 保持 false, +// 文案不会被错误信息覆写掩盖,用户仍能看到真实的额度原因。 func TaskErrorFromAPIError(apiErr *types.NewAPIError) *taskdto.TaskError { if apiErr == nil { return nil } return &taskdto.TaskError{ - Code: string(apiErr.GetErrorCode()), - Message: apiErr.Err.Error(), - StatusCode: apiErr.StatusCode, - Error: apiErr.Err, + Code: string(apiErr.GetErrorCode()), + Message: apiErr.Err.Error(), + StatusCode: apiErr.StatusCode, + FromUpstream: types.IsFromUpstreamError(apiErr), + Error: apiErr.Err, + } +} + +// APIErrorFromTaskError 是 TaskErrorFromAPIError 的反向转换,供任务链路复用 relay 的渠道 +// 错误处理(processChannelError:渠道自动禁用判定与用户可见的错误日志)。 +// +// FromUpstream 必须一并带过去:错误日志的 Content 会通过 /api/log/self 回显给用户, +// 是和响应体同级的对外出口。丢掉标记会让任务链路的日志始终写入上游原文,客户端在响应里 +// 看到 Service Unavailable,转头在日志页仍能读到上游账务细节。 +func APIErrorFromTaskError(taskErr *taskdto.TaskError) *types.NewAPIError { + if taskErr == nil { + return nil + } + err := taskErr.Error + if err == nil { + err = errors.New(taskErr.Message) + } + apiErr := types.NewOpenAIError(err, types.ErrorCodeBadResponseStatusCode, taskErr.StatusCode) + if taskErr.FromUpstream { + apiErr.MarkUpstreamOrigin() } + return apiErr } diff --git a/service/error_test.go b/service/error_test.go index 266d2a875560..0efba762cb2d 100644 --- a/service/error_test.go +++ b/service/error_test.go @@ -3,6 +3,7 @@ package service import ( "bytes" "context" + "errors" "fmt" "io" "net/http" @@ -10,6 +11,7 @@ import ( "testing" "github.com/QuantumNous/new-api/common" + taskdto "github.com/QuantumNous/new-api/dto" "github.com/QuantumNous/new-api/relaykit/types" "github.com/gin-gonic/gin" "github.com/stretchr/testify/require" @@ -150,6 +152,141 @@ func TestRelayErrorHandlerKeepsInvalidJSONBodyInDebugLog(t *testing.T) { require.Contains(t, logBuffer.String(), body) } +// TestRelayErrorHandlerIOErrorNotMarkedUpstream verifies that a local I/O failure +// while reading the upstream body is not marked as upstream-origin. The defer-based +// marker must only cover paths whose message text is sourced from the response body. +func TestRelayErrorHandlerIOErrorNotMarkedUpstream(t *testing.T) { + t.Parallel() + + // A reader that always errors on Read, so io.ReadAll fails before any + // upstream-derived message is constructed. + resp := &http.Response{ + StatusCode: http.StatusBadGateway, + Body: io.NopCloser(errReader{}), + } + + newAPIError := RelayErrorHandler(context.Background(), resp, false) + + require.NotNil(t, newAPIError) + require.False(t, types.IsFromUpstreamError(newAPIError), + "I/O failure reading upstream body is a local error and must not be marked upstream") +} + +// The upstream marker must follow the text, not merely the position in the function: +// a body that fails to parse yields a message this site builds from the status code +// alone, so it stays local unless the raw body is echoed back into it. +func TestRelayErrorHandlerMarksUpstreamPerPath(t *testing.T) { + cases := []struct { + name string + body string + showBodyWhenFail bool + wantUpstream bool + }{ + { + name: "unparseable body without echo stays local", + body: "502 Bad Gateway", + showBodyWhenFail: false, + wantUpstream: false, + }, + { + name: "unparseable body echoed into message is upstream", + body: "502 Bad Gateway", + showBodyWhenFail: true, + wantUpstream: true, + }, + { + name: "structured provider error is upstream", + body: `{"error":{"message":"You exceeded your current quota","type":"insufficient_quota"}}`, + showBodyWhenFail: false, + wantUpstream: true, + }, + { + name: "plain message body is upstream", + body: `{"message":"Please top up your credits"}`, + showBodyWhenFail: false, + wantUpstream: true, + }, + } + + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + withDebugEnabled(t, false) + + resp := &http.Response{ + StatusCode: http.StatusBadGateway, + Body: io.NopCloser(strings.NewReader(tc.body)), + } + + newAPIError := RelayErrorHandler(context.Background(), resp, tc.showBodyWhenFail) + + require.NotNil(t, newAPIError) + require.Equal(t, tc.wantUpstream, types.IsFromUpstreamError(newAPIError)) + }) + } +} + +// The task relay loop reuses processChannelError for channel auto-disable and for the +// user-visible error log, so it converts TaskError back into NewAPIError. The upstream +// marker has to survive that conversion: without it the log always records the raw +// upstream text, and a user who sees "Service Unavailable" in the response can still +// read the upstream billing details on the log page. +func TestAPIErrorFromTaskError(t *testing.T) { + t.Parallel() + + t.Run("nil task error converts to nil", func(t *testing.T) { + t.Parallel() + require.Nil(t, APIErrorFromTaskError(nil)) + }) + + t.Run("upstream marker and status code survive the conversion", func(t *testing.T) { + t.Parallel() + + taskErr := TaskErrorWrapperUpstream(errors.New("insufficient credits"), "fail_to_fetch_task", http.StatusPaymentRequired) + apiErr := APIErrorFromTaskError(taskErr) + + require.NotNil(t, apiErr) + require.True(t, types.IsFromUpstreamError(apiErr)) + require.Equal(t, http.StatusPaymentRequired, apiErr.StatusCode) + require.Equal(t, "insufficient credits", apiErr.Error()) + }) + + t.Run("locally produced task errors stay local", func(t *testing.T) { + t.Parallel() + + taskErr := TaskErrorWrapperLocal(errors.New("video_id is required"), "invalid_request", http.StatusBadRequest) + apiErr := APIErrorFromTaskError(taskErr) + + require.NotNil(t, apiErr) + require.False(t, types.IsFromUpstreamError(apiErr)) + require.Equal(t, "video_id is required", apiErr.Error()) + }) + + // TaskErrorFromAPIError leaves Error unset when the source carried no wrapped error, + // so the message is the only text available. + t.Run("falls back to Message when Error is nil", func(t *testing.T) { + t.Parallel() + + apiErr := APIErrorFromTaskError(&taskdto.TaskError{ + Message: "upstream rejected the request", + StatusCode: http.StatusBadGateway, + FromUpstream: true, + }) + + require.NotNil(t, apiErr) + require.Equal(t, "upstream rejected the request", apiErr.Error()) + require.True(t, types.IsFromUpstreamError(apiErr)) + }) +} + +// errReader is an io.ReadCloser whose Read always returns an error. +type errReader struct{} + +func (errReader) Read(p []byte) (int, error) { + return 0, io.ErrUnexpectedEOF +} + +func (errReader) Close() error { return nil } + func withDebugEnabled(t *testing.T, enabled bool) { t.Helper() diff --git a/service/task_billing.go b/service/task_billing.go index f310a8230c63..329c1af2859d 100644 --- a/service/task_billing.go +++ b/service/task_billing.go @@ -10,6 +10,7 @@ import ( "github.com/QuantumNous/new-api/logger" "github.com/QuantumNous/new-api/model" relaycommon "github.com/QuantumNous/new-api/relay/common" + "github.com/QuantumNous/new-api/setting/operation_setting" "github.com/QuantumNous/new-api/setting/ratio_setting" "github.com/QuantumNous/new-api/types" "github.com/gin-gonic/gin" @@ -184,6 +185,13 @@ func RefundTaskQuota(ctx context.Context, task *model.Task, reason string) bool // 3. 记录日志 other := taskBillingOther(task) other["task_id"] = task.TaskID + // 退款日志会通过 /api/log/self 回显给用户,reason 与任务 FailReason 是同一份上游文案, + // 因此和 relay 的错误日志一样在写库前覆写,原文改记到 admin_info + // (model.formatUserLogs 会为普通用户剥离整个 admin_info)。 + if masked := operation_setting.OverrideUpstreamMessage(reason); masked != reason { + other["admin_info"] = map[string]interface{}{"original_reason": reason} + reason = masked + } other["reason"] = reason model.RecordTaskBillingLog(model.RecordTaskBillingLogParams{ UserId: task.UserId, diff --git a/service/task_billing_test.go b/service/task_billing_test.go index 53e3f680d01c..a7a0739c5c77 100644 --- a/service/task_billing_test.go +++ b/service/task_billing_test.go @@ -12,6 +12,7 @@ import ( "github.com/QuantumNous/new-api/common" "github.com/QuantumNous/new-api/model" relaycommon "github.com/QuantumNous/new-api/relay/common" + "github.com/QuantumNous/new-api/setting/operation_setting" "github.com/QuantumNous/new-api/types" "github.com/glebarez/sqlite" "github.com/shopspring/decimal" @@ -422,6 +423,73 @@ func TestRefundTaskQuota_FundingFailureKeepsPendingMarker(t *testing.T) { assert.Equal(t, int64(0), countLogs(t)) } +// The refund reason is the same upstream text the task query endpoints mask, and the +// refund log is echoed to the user through /api/log/self. It therefore has to be +// overridden before the row is written, with the original kept under admin_info — +// model.formatUserLogs strips that whole key for non-admin views. +func TestRefundTaskQuota_MasksUpstreamReason(t *testing.T) { + origEnabled := operation_setting.ErrorOverrideEnabled + origKeywords := operation_setting.ErrorOverrideKeywords + t.Cleanup(func() { + operation_setting.ErrorOverrideEnabled = origEnabled + operation_setting.ErrorOverrideKeywords = origKeywords + }) + operation_setting.ErrorOverrideEnabled = true + operation_setting.ErrorOverrideKeywords = []string{"no available", "quota", "credits", "top-up"} + + const upstreamReason = "insufficient credits, please top-up your account" + const localReason = "任务超时(30分钟)" + + cases := []struct { + name string + userID int + reason string + wantReason string + wantOriginal string + }{ + { + name: "upstream reason is masked and the original moves to admin_info", + userID: 11, + reason: upstreamReason, + wantReason: operation_setting.ErrorOverrideMessage, + wantOriginal: upstreamReason, + }, + { + name: "local reason is written verbatim without admin_info", + userID: 12, + reason: localReason, + wantReason: localReason, + }, + } + + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + truncate(t) + const channelID, preConsumed = 11, 700 + seedUser(t, tc.userID, 10000) + seedChannel(t, channelID) + + task := makeTask(tc.userID, channelID, preConsumed, 0, BillingSourceWallet, 0) + require.NoError(t, model.DB.Create(task).Error) + require.True(t, RefundTaskQuota(context.Background(), task, tc.reason)) + + log := getLastLog(t) + require.NotNil(t, log) + other, err := common.StrToMap(log.Other) + require.NoError(t, err) + assert.Equal(t, tc.wantReason, other["reason"]) + + if tc.wantOriginal == "" { + assert.NotContains(t, other, "admin_info") + return + } + adminInfo, ok := other["admin_info"].(map[string]interface{}) + require.True(t, ok, "admin_info should carry the original reason") + assert.Equal(t, tc.wantOriginal, adminInfo["original_reason"]) + }) + } +} + // =========================================================================== // RecalculateTaskQuota tests // =========================================================================== diff --git a/service/violation_fee.go b/service/violation_fee.go index f51533629d76..26c34bfb66d5 100644 --- a/service/violation_fee.go +++ b/service/violation_fee.go @@ -45,7 +45,12 @@ func WrapAsViolationFeeGrokCSAM(err *types.NewAPIError) *types.NewAPIError { oai := err.ToOpenAIError() oai.Type = string(types.ErrorCodeViolationFeeGrokCSAM) oai.Code = string(types.ErrorCodeViolationFeeGrokCSAM) - return types.WithOpenAIError(oai, err.StatusCode, types.ErrOptionWithSkipRetry()) + wrapped := types.WithOpenAIError(oai, err.StatusCode, types.ErrOptionWithSkipRetry()) + // 重新包装不应改变错误来源,否则上游错误会被误判为本站错误 + if types.IsFromUpstreamError(err) { + wrapped.MarkUpstreamOrigin() + } + return wrapped } // NormalizeViolationFeeError ensures: @@ -64,7 +69,11 @@ func NormalizeViolationFeeError(err *types.NewAPIError) *types.NewAPIError { if IsViolationFeeCode(err.GetErrorCode()) { oai := err.ToOpenAIError() - return types.WithOpenAIError(oai, err.StatusCode, types.ErrOptionWithSkipRetry()) + wrapped := types.WithOpenAIError(oai, err.StatusCode, types.ErrOptionWithSkipRetry()) + if types.IsFromUpstreamError(err) { + wrapped.MarkUpstreamOrigin() + } + return wrapped } return err diff --git a/setting/model_setting/model_alias.go b/setting/model_setting/model_alias.go index 078690c01e5b..0d7f4e5ae4e9 100644 --- a/setting/model_setting/model_alias.go +++ b/setting/model_setting/model_alias.go @@ -7,7 +7,6 @@ import ( "github.com/QuantumNous/new-api/common" "github.com/QuantumNous/new-api/setting/config" - "github.com/QuantumNous/new-api/setting/ratio_setting" "github.com/QuantumNous/new-api/types" ) @@ -27,18 +26,13 @@ func init() { } // ResolveModelAlias 把全局别名解析为真实模型名,支持链式映射(a->b->c)并检测环。 -// 自映射(a->a)视为未配置别名。带 compact 后缀的模型名先剥离后缀参与解析,结果再补回后缀。 +// 自映射(a->a)视为未配置别名。 // 返回(解析后的模型名, 是否发生了别名替换, error)。 func ResolveModelAlias(requested string) (string, bool, error) { if requested == "" || modelAliasSetting.Mapping.Len() == 0 { return requested, false, nil } - baseName := requested - hasCompactSuffix := strings.HasSuffix(requested, ratio_setting.CompactModelSuffix) - if hasCompactSuffix { - baseName = strings.TrimSuffix(requested, ratio_setting.CompactModelSuffix) - } - current := baseName + current := requested visited := map[string]bool{current: true} applied := false for { @@ -56,9 +50,6 @@ func ResolveModelAlias(requested string) (string, bool, error) { if !applied { return requested, false, nil } - if hasCompactSuffix { - current = ratio_setting.WithCompactModelSuffix(current) - } return current, true, nil } diff --git a/setting/model_setting/model_alias_test.go b/setting/model_setting/model_alias_test.go index 4d04a53745fe..6dcfefbc815f 100644 --- a/setting/model_setting/model_alias_test.go +++ b/setting/model_setting/model_alias_test.go @@ -74,12 +74,22 @@ func TestResolveModelAlias(t *testing.T) { wantApplied: false, }, { - name: "compact suffix stripped and restored", - mapping: map[string]string{"cinax": "cinax-pro"}, + // 别名按完整模型名整体匹配,不对名字做任何分段拆解。 + name: "alias matches the full model name verbatim", + mapping: map[string]string{"cinax-openai-compact": "cinax-pro-openai-compact"}, requested: "cinax-openai-compact", wantResolved: "cinax-pro-openai-compact", wantApplied: true, }, + { + // 基名映射不会命中更长的模型名:解析不再剥离任何后缀, + // 带后缀的模型必须显式配置自己的映射。 + name: "base-name mapping does not match a longer model name", + mapping: map[string]string{"cinax": "cinax-pro"}, + requested: "cinax-openai-compact", + wantResolved: "cinax-openai-compact", + wantApplied: false, + }, { name: "empty requested model returns as-is", mapping: map[string]string{"cinax": "cinax-pro"}, diff --git a/setting/operation_setting/error_override_setting.go b/setting/operation_setting/error_override_setting.go new file mode 100644 index 000000000000..60e41185e8c4 --- /dev/null +++ b/setting/operation_setting/error_override_setting.go @@ -0,0 +1,75 @@ +package operation_setting + +import ( + "strings" + + "github.com/QuantumNous/new-api/relaykit/types" +) + +// ErrorOverrideMessage 覆写后统一返回给客户端的错误文案 +const ErrorOverrideMessage = "Service Unavailable" + +// 错误信息覆写:命中关键词的上游错误统一替换文案,避免向客户端暴露上游账务状态与转售链路。 +// 仅作用于上游传递到本站的错误;本站自身产生的错误(用户额度不足、无可用渠道等)不受影响。 +var ErrorOverrideEnabled = false + +// 关键词统一以小写存储,匹配时对错误文案做小写化后比对 +var ErrorOverrideKeywords = []string{"no available", "quota", "credits", "top-up"} + +func ErrorOverrideKeywordsToString() string { + return strings.Join(ErrorOverrideKeywords, "\n") +} + +func ErrorOverrideKeywordsFromString(s string) { + ErrorOverrideKeywords = []string{} + ak := strings.Split(s, "\n") + for _, k := range ak { + k = strings.TrimSpace(k) + k = strings.ToLower(k) + if k != "" { + ErrorOverrideKeywords = append(ErrorOverrideKeywords, k) + } + } +} + +// ShouldOverrideUpstreamError 报告该错误写给客户端时是否会被覆写,但不修改错误本身。 +// 供需要在覆写前分叉行为的调用方使用(例如错误日志要区分记录原文还是覆写后的文案)。 +func ShouldOverrideUpstreamError(err *types.NewAPIError) bool { + if !types.IsFromUpstreamError(err) { + return false + } + return shouldOverrideErrorMessage(err.Error()) +} + +// OverrideUpstreamError 在错误写给客户端前覆写其文案,返回是否已覆写。必须在错误日志记录 +// 之后调用,后台日志与渠道自动禁用判定始终使用原始上游文案。状态码、error.type、error.code +// 保持不变。 +func OverrideUpstreamError(err *types.NewAPIError) bool { + if !ShouldOverrideUpstreamError(err) { + return false + } + err.ReplaceMessage(ErrorOverrideMessage) + return true +} + +// OverrideUpstreamMessage 供不经过 NewAPIError 的链路使用(任务平台、Midjourney)。 +// 调用方需自行确认文案来自上游。 +func OverrideUpstreamMessage(message string) string { + if !shouldOverrideErrorMessage(message) { + return message + } + return ErrorOverrideMessage +} + +func shouldOverrideErrorMessage(message string) bool { + if !ErrorOverrideEnabled || message == "" { + return false + } + lowerMessage := strings.ToLower(message) + for _, keyword := range ErrorOverrideKeywords { + if strings.Contains(lowerMessage, keyword) { + return true + } + } + return false +} diff --git a/setting/operation_setting/error_override_setting_test.go b/setting/operation_setting/error_override_setting_test.go new file mode 100644 index 000000000000..de0ceaf8e475 --- /dev/null +++ b/setting/operation_setting/error_override_setting_test.go @@ -0,0 +1,240 @@ +package operation_setting + +import ( + "errors" + "net/http" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "github.com/QuantumNous/new-api/relaykit/types" +) + +// restoreOverrideState resets the package-level override settings to the defaults used +// by each test and reinstalls them after the test runs, so tests stay isolated. +func restoreOverrideState(t *testing.T) { + t.Helper() + origEnabled := ErrorOverrideEnabled + origKeywords := ErrorOverrideKeywords + t.Cleanup(func() { + ErrorOverrideEnabled = origEnabled + ErrorOverrideKeywords = origKeywords + }) + ErrorOverrideEnabled = true + ErrorOverrideKeywords = []string{"no available", "quota", "credits", "top-up"} +} + +func TestOverrideUpstreamMessage_DisabledKeepsMessage(t *testing.T) { + orig := ErrorOverrideEnabled + t.Cleanup(func() { ErrorOverrideEnabled = orig }) + ErrorOverrideEnabled = false + + msg := "Your credit balance is too low, please top up" + assert.Equal(t, msg, OverrideUpstreamMessage(msg)) + // Empty message is never overridden even when enabled. + assert.Equal(t, "", OverrideUpstreamMessage("")) +} + +func TestOverrideUpstreamMessage_MatchesKeywords(t *testing.T) { + restoreOverrideState(t) + + cases := []struct { + name string + message string + want string + }{ + {"credits", "insufficient credits", ErrorOverrideMessage}, + {"top-up", "Please top-up your account", ErrorOverrideMessage}, + {"quota", "You exceeded your current quota", ErrorOverrideMessage}, + {"no available", "No available channel for model gpt-4o under group default", ErrorOverrideMessage}, + {"case-insensitive credits", "CREDITS exhausted", ErrorOverrideMessage}, + {"case-insensitive quota", "QUOTA exceeded", ErrorOverrideMessage}, + {"no match", "Internal Server Error", "Internal Server Error"}, + {"empty", "", ""}, + } + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + assert.Equal(t, tc.want, OverrideUpstreamMessage(tc.message)) + }) + } +} + +func TestOverrideUpstreamError_LocalErrorsNotOverridden(t *testing.T) { + restoreOverrideState(t) + + // 本地错误:用户额度不足(由本站产生,未标记上游来源)。 + localQuota := types.NewErrorWithStatusCode( + errors.New("用户额度不足, 剩余额度: $0.00"), + types.ErrorCodeInsufficientUserQuota, http.StatusPaymentRequired) + assert.False(t, OverrideUpstreamError(localQuota)) + assert.Equal(t, "用户额度不足, 剩余额度: $0.00", localQuota.Error()) + + // 本地错误:distributor 的「无可用渠道」文案虽然命中 "no available" 关键词, + // 但因未标记上游来源,不应被覆写。 + localNoChannel := types.NewErrorWithStatusCode( + errors.New("No available channel for model gpt-4o under group default (distributor)"), + types.ErrorCodeBadResponseStatusCode, http.StatusServiceUnavailable) + assert.False(t, OverrideUpstreamError(localNoChannel)) + assert.Contains(t, localNoChannel.Error(), "No available channel") +} + +func TestOverrideUpstreamError_UpstreamErrorsOverridden(t *testing.T) { + restoreOverrideState(t) + + cases := []struct { + name string + build func() *types.NewAPIError + message string + }{ + { + name: "openai credits", + build: func() *types.NewAPIError { + return types.WithUpstreamOpenAIError(types.OpenAIError{ + Message: "insufficient credits balance, please top up", + Type: "upstream_error", + Code: "credits_exhausted", + }, http.StatusPaymentRequired) + }, + message: "insufficient credits balance, please top up", + }, + { + name: "claude quota", + build: func() *types.NewAPIError { + return types.WithUpstreamClaudeError(types.ClaudeError{ + Type: "error", + Message: "You exceeded your current quota", + }, http.StatusTooManyRequests) + }, + message: "You exceeded your current quota", + }, + { + name: "openai no available", + build: func() *types.NewAPIError { + return types.WithUpstreamOpenAIError(types.OpenAIError{ + Message: "No available channel for model gpt-4o under group default", + Type: "upstream_error", + Code: "no_available_channel", + }, http.StatusServiceUnavailable) + }, + message: "No available channel for model gpt-4o under group default", + }, + } + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + err := tc.build() + require.True(t, types.IsFromUpstreamError(err)) + require.True(t, OverrideUpstreamError(err)) + + // Every user-facing representation is rewritten. + assert.Equal(t, ErrorOverrideMessage, err.Error()) + openAI := err.ToOpenAIError() + assert.Equal(t, ErrorOverrideMessage, openAI.Message) + assert.Nil(t, openAI.Metadata) + }) + } +} + +func TestOverrideUpstreamError_PreservesStatusCodeTypeCode(t *testing.T) { + restoreOverrideState(t) + + err := types.WithUpstreamOpenAIError(types.OpenAIError{ + Message: "Please top up your credits", + Type: "upstream_error", + Code: "credits_exhausted", + }, http.StatusPaymentRequired) + originalCode := err.GetErrorCode() + originalType := err.GetErrorType() + + require.True(t, OverrideUpstreamError(err)) + + assert.Equal(t, http.StatusPaymentRequired, err.StatusCode) + assert.Equal(t, originalType, err.GetErrorType()) + assert.Equal(t, originalCode, err.GetErrorCode()) +} + +func TestOverrideUpstreamError_DisabledDoesNothing(t *testing.T) { + orig := ErrorOverrideEnabled + t.Cleanup(func() { ErrorOverrideEnabled = orig }) + ErrorOverrideEnabled = false + + err := types.WithUpstreamOpenAIError(types.OpenAIError{ + Message: "quota exhausted", + Type: "upstream_error", + }, http.StatusTooManyRequests) + assert.False(t, OverrideUpstreamError(err)) + assert.Equal(t, "quota exhausted", err.Error()) +} + +func TestOverrideUpstreamError_NilSafe(t *testing.T) { + restoreOverrideState(t) + require.NotPanics(t, func() { OverrideUpstreamError(nil) }) +} + +// ShouldOverrideUpstreamError answers the same question as OverrideUpstreamError but +// must not mutate the error. processChannelError relies on that to decide what to write +// into the user-visible error log while the untouched error still drives retry and +// channel-disable decisions. +func TestShouldOverrideUpstreamError_DoesNotMutate(t *testing.T) { + restoreOverrideState(t) + + upstream := types.WithUpstreamOpenAIError(types.OpenAIError{ + Message: "insufficient credits, please top up", + Type: "upstream_error", + Code: "credits_exhausted", + }, http.StatusPaymentRequired) + + require.True(t, ShouldOverrideUpstreamError(upstream)) + // Asking twice must stay stable, and the message must survive untouched. + require.True(t, ShouldOverrideUpstreamError(upstream)) + assert.Equal(t, "insufficient credits, please top up", upstream.Error()) + assert.Equal(t, "insufficient credits, please top up", upstream.ToOpenAIError().Message) + + local := types.NewErrorWithStatusCode( + errors.New("用户额度不足, 剩余额度: $0.00"), + types.ErrorCodeInsufficientUserQuota, http.StatusPaymentRequired) + assert.False(t, ShouldOverrideUpstreamError(local)) + require.NotPanics(t, func() { ShouldOverrideUpstreamError(nil) }) + assert.False(t, ShouldOverrideUpstreamError(nil)) +} + +func TestOverrideUpstreamMessage_NoKeywords(t *testing.T) { + origEnabled := ErrorOverrideEnabled + origKeywords := ErrorOverrideKeywords + t.Cleanup(func() { + ErrorOverrideEnabled = origEnabled + ErrorOverrideKeywords = origKeywords + }) + ErrorOverrideEnabled = true + ErrorOverrideKeywords = nil + + msg := "Your credits are too low" + assert.Equal(t, msg, OverrideUpstreamMessage(msg)) +} + +func TestErrorOverrideKeywordsFromString_ParsesAndNormalizes(t *testing.T) { + orig := ErrorOverrideKeywords + t.Cleanup(func() { ErrorOverrideKeywords = orig }) + + ErrorOverrideKeywordsFromString("Credits\n Quota \n\n\ntop-up\n ") + + require.Equal(t, []string{"credits", "quota", "top-up"}, ErrorOverrideKeywords) +} + +func TestErrorOverrideKeywordsFromString_EmptyStringClears(t *testing.T) { + orig := ErrorOverrideKeywords + t.Cleanup(func() { ErrorOverrideKeywords = orig }) + + ErrorOverrideKeywordsFromString("") + require.Empty(t, ErrorOverrideKeywords) +} + +func TestErrorOverrideKeywordsToString_RoundTrip(t *testing.T) { + orig := ErrorOverrideKeywords + t.Cleanup(func() { ErrorOverrideKeywords = orig }) + + ErrorOverrideKeywords = []string{"no available", "quota", "credits", "top-up"} + s := ErrorOverrideKeywordsToString() + ErrorOverrideKeywordsFromString(s) + require.Equal(t, []string{"no available", "quota", "credits", "top-up"}, ErrorOverrideKeywords) +} diff --git a/web/src/features/keys/components/__tests__/api-key-group-cell.test.tsx b/web/src/features/keys/components/__tests__/api-key-group-cell.test.tsx index 5cb64ae57f7b..6da76150afc0 100644 --- a/web/src/features/keys/components/__tests__/api-key-group-cell.test.tsx +++ b/web/src/features/keys/components/__tests__/api-key-group-cell.test.tsx @@ -100,7 +100,7 @@ describe('API key group table cell', () => { domWindow.close() }) - test('renders two unclipped rings and a localized Auto ratio when API data uses a nonlocalized string', async () => { + test('renders an unclipped ring and a localized Auto ratio when API data uses a nonlocalized string', async () => { const container = document.createElement('div') document.body.append(container) const root = createRoot(container) @@ -123,12 +123,14 @@ describe('API key group table cell', () => { assert.equal(badgeCell.classList.contains('overflow-visible'), true) assert.equal(badgeCell.classList.contains('overflow-hidden'), false) + // AutoGroupBadge 在 ApiKeyGroupCell 中被停用(由 Cross-group StatusBadge 承担该位置), + // 因此 auto 分组只剩 GroupRatioBadge 一个带流光边框的 frame。 const frames = container.querySelectorAll('[data-auto-group-frame]') const movingRings = container.querySelectorAll( '[data-auto-group-flow-border]' ) - assert.equal(frames.length, 2) - assert.equal(movingRings.length, 2) + assert.equal(frames.length, 1) + assert.equal(movingRings.length, 1) for (const frame of frames) { assert.equal(frame.classList.contains('relative'), true) assert.equal(frame.classList.contains('overflow-visible'), true) @@ -155,7 +157,7 @@ describe('API key group table cell', () => { container.remove() }) - test('keeps static Auto frames but omits both moving layers for reduced motion', async () => { + test('keeps the static Auto frame but omits the moving layer for reduced motion', async () => { const container = document.createElement('div') document.body.append(container) const root = createRoot(container) @@ -166,7 +168,7 @@ describe('API key group table cell', () => { assert.equal( container.querySelectorAll('[data-auto-group-frame]').length, - 2 + 1 ) assert.equal( container.querySelectorAll('[data-auto-group-flow-border]').length, @@ -177,7 +179,7 @@ describe('API key group table cell', () => { container.remove() }) - test('shows only the Auto badge when ratio data is unavailable', async () => { + test('shows only the Cross-group badge when ratio data is unavailable', async () => { const container = document.createElement('div') document.body.append(container) const root = createRoot(container) @@ -186,19 +188,21 @@ describe('API key group table cell', () => { root.render() ) + // 无 ratio 时 GroupRatioBadge 返回 null,auto 分组不渲染任何 frame, + // 只保留 Cross-group 状态徽章。 assert.equal( container.querySelectorAll('[data-auto-group-frame]').length, - 1 + 0 ) assert.equal( container.querySelectorAll('[data-auto-group-flow-border]').length, - 1 + 0 ) assert.equal( container.querySelector('[data-auto-group-effect="ratio"]'), null ) - assert.equal(container.textContent?.includes('Auto'), true) + assert.equal(container.textContent?.includes('Cross-group'), true) assert.equal(container.textContent?.includes('Ratio'), false) await act(async () => root.unmount()) diff --git a/web/src/features/system-settings/general/system-behavior-section.tsx b/web/src/features/system-settings/general/system-behavior-section.tsx index 4413a2a02b1b..f59f2a100e30 100644 --- a/web/src/features/system-settings/general/system-behavior-section.tsx +++ b/web/src/features/system-settings/general/system-behavior-section.tsx @@ -26,9 +26,12 @@ import { FormControl, FormDescription, FormField, + FormItem, FormLabel, + FormMessage, } from '@/components/ui/form' import { Switch } from '@/components/ui/switch' +import { Textarea } from '@/components/ui/textarea' import { SettingsForm, @@ -44,6 +47,8 @@ const behaviorSchema = z.object({ DefaultCollapseSidebar: z.boolean(), DemoSiteEnabled: z.boolean(), SelfUseModeEnabled: z.boolean(), + ErrorOverrideEnabled: z.boolean(), + ErrorOverrideKeywords: z.string(), ChannelFailoverEnabled: z.boolean(), }) @@ -147,6 +152,54 @@ export function SystemBehaviorSection({ )} /> + ( + + + {t('Error Message Override')} + + {t( + 'Replace upstream error messages that match a keyword with a generic message, so upstream billing state is not exposed. Errors produced by this site are never replaced.' + )} + + + + + + + )} + /> + + {form.watch('ErrorOverrideEnabled') && ( + ( + + {t('Override keywords')} + +