diff --git a/MY_COMMITS_BACKUP.md b/MY_COMMITS_BACKUP.md index 470f70b8c21b..7c2b62e92d0b 100644 --- a/MY_COMMITS_BACKUP.md +++ b/MY_COMMITS_BACKUP.md @@ -8,7 +8,7 @@ e8c836d70 fix(web): improve form validation error focus #5163 ``` -更新时间:2026-06-25。已逐个查看当前 fork 独有提交,并按“当前工作区最终仍保留的差异”重新整理。Redis Sentinel 支持和 Sentinel 故障转移重试逻辑已经从当前工作区移除,不再作为保留改动记录。 +更新时间:2026-07-23。已逐个查看当前 fork 独有提交,并按“当前工作区最终仍保留的差异”重新整理。Redis Sentinel 支持和 Sentinel 故障转移重试逻辑已经从当前工作区移除,不再作为保留改动记录。 ## 当前保留的改动 @@ -226,6 +226,73 @@ e8c836d70 fix(web): improve form validation error focus #5163 | ---- | ---- | | `service/channel_affinity_usage_cache_test.go` | atomic 唯一测试 key | +### 12. VChart 图表崩溃与深色模式修复 + +修复仪表盘图表在生产构建下运行时崩溃(`TypeError: Cannot read properties of undefined (reading 'createCanvas')`)以及修好后图表深色模式不生效的问题。根因是 VisActor 系列包在依赖树里存在多份物理副本,导致多个互不相通的单例。 + +保留行为: + +- **浏览器环境注册(两个前端)**:VChart 的浏览器环境注册是带副作用的,但 `@visactor/vchart` 的 `package.json` `sideEffects` 未列出它,生产构建 tree-shaking 会摇掉,运行时没有 env 被激活,`application.global.envContribution` 为 undefined,首个图表 `createCanvas` 崩溃。改为显式且不可被摇树的 `VChart.useRegisters([registerBrowserEnv])`,在任何图表挂载前注册。 +- **vrender 单例去重(classic)**:classic 依赖树里 `@visactor/vrender-core`(+`vrender-kits`/`vutils`,均 0.17.17)存在两份物理副本(`react-vchart/` 与 `vchart/` 各嵌套一份)。`vrender-core` 通过 `application.global` 维护渲染环境单例,多份副本 = 多个互不相通的单例,`registerBrowserEnv` 注册到一份、`` 渲染读另一份,env 仍为 undefined。通过 rsbuild alias 强制指向 classic 自带 vchart 内嵌的那份 0.17.17。 +- **vchart / ThemeManager 单例去重(classic)**:`ThemeManager` 挂在 `@visactor/vchart` 类上,classic 树里有两份 `@visactor/vchart`——应用经 `react-vchart` → classic 的 vchart 渲染,而 `@visactor/vchart-semi-theme` 内部 `import VChart from '@visactor/vchart'` 命中它自己嵌套的另一份。`initVChartSemiTheme` 把 `semiDesignDark` 注册/切换到后者的 ThemeManager,渲染读前者,导致深色主题不生效(图表只显示浅色)。将 `@visactor/vchart` 本身也 alias 到 classic 自带的那份(1.8.11),使渲染引擎、浏览器环境、Semi 主题共用同一组单例。 +- 深色链路:classic Theme context 设置 `document.body[theme-mode="dark"]` → `vchart-semi-theme` observer(`isWatchingThemeSwitch:true`)→ `setCurrentTheme('semiDesignDark')` → 图表经同一 ThemeManager 渲染。 +- **注意**:alias 不能指向 workspace 顶层 hoist 的 `@visactor/vchart` 2.1.2 / `vrender-core` 1.1.4——那是给 default 新前端 vchart 2.x 用的,与 classic 的 1.8.11 大版本不兼容。 + +构建产物自检(classic):`getCommonCanvas`、`isBrowserBound`、`setActiveEnvContribution`、`setCurrentTheme` 在 `dist/static/js/*.js` 中应各只出现一次。 + +主要文件: + +| 文件 | 说明 | +| ---- | ---- | +| `web/default/src/lib/vchart.ts` | default 前端显式注册浏览器环境 `VChart.useRegisters([registerBrowserEnv])` | +| `web/default/src/main.tsx` | 入口 side-effect 导入 `@/lib/vchart`,防止被 tree-shaking 丢弃 | +| `web/classic/src/constants/dashboard.constants.js` | classic 前端显式注册浏览器环境 | +| `web/classic/rsbuild.config.ts` | `visactorDedupeAlias`:将 `@visactor/vchart` 及其内嵌 `vrender-core`/`vrender-kits`/`vrender-components`/`vutils` 全部指向 classic 自带的单一副本 | + +对应提交: + +```text +724ecece2 fix(web): register VChart browser env to prevent createCanvas crash +3c2cf2321 fix(web/classic): dedupe vrender to one copy so VChart env registration applies +7c0cac115 fix(web/classic): dedupe @visactor/vchart so Semi dark chart theme applies +``` + +### 13. 上游响应语义校验与安全重试 + +修复 OpenAI Chat/Completions/Responses、Claude 和 Gemini 上游返回 HTTP 200、JSON/SSE 外壳合法但没有任何可消费语义输出时,被误当成成功响应并计费的问题,同时覆盖纯工具调用和流式边缘场景。 + +保留行为: + +- 统一按“语义输出”而不是响应体非空判断成功:文本、reasoning、refusal/content filter、音频/图片/代码结果和结构完整的工具调用均可构成有效输出;只有 role、usage、ping、start/stop、空 candidate/choice/content 等协议外壳不算输出。 +- OpenAI Chat/Completions 同时校验非流式和流式响应;工具调用必须有函数名,聚合后的 arguments 必须是完整 JSON,并正确统计并行工具调用。 +- OpenAI Responses API 校验 `completed`/`incomplete`/`failed` 状态、输出项及 function/custom tool call;流式工具调用按 item id 与 output index 关联,拒绝缺名或参数截断。 +- Claude 不再把任意合法 SSE 事件或 `stop_reason=tool_use` 的空 content 当作成功;工具流必须有 `tool_use` 块、完整 JSON 输入以及终止事件。 +- Gemini 拒绝空 candidate、空 parts、只有 usage 的流片段和无终止原因的截断流;安全过滤仍按明确拒绝处理,纯 function call 保持有效。 +- 在首个语义输出前缓存流式协议外壳,使真正的空回仍可在未写入客户端时触发渠道间重试;一旦响应已经写出则禁止切换渠道,避免把两个上游流拼接到同一客户端响应,并发送对应协议的终止错误事件。 +- 每次渠道重试前重置 response count、stream status、首包时间、thinking/Claude 转换状态和 Responses 内置工具计数,防止上一渠道状态污染下一渠道。 +- 删除不再使用的 40k stars light 海报 PNG/SVG。 + +主要文件: + +| 文件 | 说明 | +| ---- | ---- | +| `relay/responsevalidator/validator.go` | OpenAI、Responses、Claude、Gemini 的统一语义与终止状态校验 | +| `relay/channel/openai/relay-openai.go` | Chat/Completions 非流式及 SSE 校验、首个语义输出前缓存 | +| `relay/channel/openai/relay_responses.go` | Responses API 非流式及事件流校验 | +| `relay/channel/claude/relay-claude.go` | Claude 消息及事件流校验 | +| `relay/channel/gemini/relay-gemini.go` | Gemini 转 OpenAI 响应的空回、工具调用及终止校验 | +| `relay/channel/gemini/relay-gemini-native.go` | Gemini 原生响应的对应校验 | +| `controller/relay.go` | 已写响应禁止跨渠道重试,流内返回协议错误 | +| `relay/common/relay_info.go` | 渠道重试前清理响应状态 | +| `relay/helper/common.go` | OpenAI、Responses、Claude SSE 终止错误输出 | +| `relay/responsevalidator/validator_test.go` | 空外壳、纯工具调用、参数截断、过滤和终止状态测试 | + +对应提交: + +```text +55ac05637 fix(relay): reject semantically empty upstream responses +``` + ## 已移除或不再保留的改动 ### Redis Sentinel 支持 @@ -267,8 +334,4 @@ Redis Sentinel ## 当前工作区状态说明 -截至本文档更新时,当前工作区包含三类未提交改动: - -- 移除 Redis Sentinel 支持:`common/redis.go` -- 移除 Sentinel failover 重试逻辑:`middleware/rate-limit.go` -- 更新本地修改备份文档:`MY_COMMITS_BACKUP.md` +截至 2026-07-23,本次上游空响应语义校验、相关测试、海报删除和本备份文档均已纳入提交;没有为这些改动保留未提交文件。 diff --git a/controller/relay.go b/controller/relay.go index ee24100d5346..7d5cad2ff13d 100644 --- a/controller/relay.go +++ b/controller/relay.go @@ -74,6 +74,7 @@ func Relay(c *gin.Context, relayFormat types.RelayFormat) { var ( newAPIError *types.NewAPIError ws *websocket.Conn + relayInfo *relaycommon.RelayInfo ) if relayFormat == types.RelayFormatOpenAIRealtime { @@ -90,6 +91,13 @@ func Relay(c *gin.Context, relayFormat types.RelayFormat) { if newAPIError != nil { logger.LogError(c, fmt.Sprintf("relay error: %s", common.LocalLogPreview(newAPIError.Error()))) newAPIError.SetMessage(common.MessageWithRequestId(newAPIError.Error(), requestId)) + if relayInfo != nil && relayInfo.IsStream && c.Writer.Written() { + helper.StreamError(c, relayFormat, newAPIError) + return + } + if c.Writer.Written() { + return + } switch relayFormat { case types.RelayFormatOpenAIRealtime: helper.WssError(c, ws, newAPIError.ToOpenAIError()) @@ -117,7 +125,7 @@ func Relay(c *gin.Context, relayFormat types.RelayFormat) { return } - relayInfo, err := relaycommon.GenRelayInfo(c, relayFormat, request, ws) + relayInfo, err = relaycommon.GenRelayInfo(c, relayFormat, request, ws) if err != nil { newAPIError = types.NewError(err, types.ErrorCodeGenRelayInfoFailed) return @@ -190,6 +198,9 @@ func Relay(c *gin.Context, relayFormat types.RelayFormat) { for ; retryParam.GetRetry() <= common.RetryTimes; retryParam.IncreaseRetry() { relayInfo.RetryIndex = retryParam.GetRetry() + if relayInfo.RetryIndex > 0 { + relayInfo.ResetResponseStateForRetry() + } channel, channelErr := getChannel(c, relayInfo, retryParam) if channelErr != nil { logger.LogError(c, channelErr.Error()) @@ -326,6 +337,9 @@ func shouldRetry(c *gin.Context, openaiErr *types.NewAPIError, retryTimes int) b if openaiErr == nil { return false } + if c != nil && c.Writer != nil && c.Writer.Written() { + return false + } if service.ShouldSkipRetryAfterChannelAffinityFailure(c) { return false } diff --git a/controller/relay_retry_test.go b/controller/relay_retry_test.go new file mode 100644 index 000000000000..45957fc408e5 --- /dev/null +++ b/controller/relay_retry_test.go @@ -0,0 +1,26 @@ +package controller + +import ( + "net/http" + "net/http/httptest" + "testing" + + "github.com/QuantumNous/new-api/types" + "github.com/gin-gonic/gin" + "github.com/stretchr/testify/require" +) + +func TestShouldRetryRejectsCommittedResponse(t *testing.T) { + gin.SetMode(gin.TestMode) + recorder := httptest.NewRecorder() + context, _ := gin.CreateTestContext(recorder) + _, err := context.Writer.Write([]byte("partial")) + require.NoError(t, err) + + apiErr := types.NewOpenAIError( + http.ErrHandlerTimeout, + types.ErrorCodeBadResponse, + http.StatusInternalServerError, + ) + require.False(t, shouldRetry(context, apiErr, 1)) +} diff --git a/dto/openai_request.go b/dto/openai_request.go index fd0bed0ea4c4..fdb675458d5b 100644 --- a/dto/openai_request.go +++ b/dto/openai_request.go @@ -291,6 +291,9 @@ type Message struct { Prefix *bool `json:"prefix,omitempty"` ReasoningContent *string `json:"reasoning_content,omitempty"` Reasoning *string `json:"reasoning,omitempty"` + Refusal *string `json:"refusal,omitempty"` + Audio json.RawMessage `json:"audio,omitempty"` + FunctionCall json.RawMessage `json:"function_call,omitempty"` ToolCalls json.RawMessage `json:"tool_calls,omitempty"` ToolCallId string `json:"tool_call_id,omitempty"` parsedContent []MediaContent diff --git a/dto/openai_response.go b/dto/openai_response.go index e503f5f91c69..4dbd04410377 100644 --- a/dto/openai_response.go +++ b/dto/openai_response.go @@ -34,6 +34,7 @@ type TextResponse struct { type OpenAITextResponseChoice struct { Index int `json:"index"` Message `json:"message"` + Text string `json:"text,omitempty"` FinishReason string `json:"finish_reason"` } @@ -89,6 +90,9 @@ type ChatCompletionsStreamResponseChoiceDelta struct { Content *string `json:"content,omitempty"` ReasoningContent *string `json:"reasoning_content,omitempty"` Reasoning *string `json:"reasoning,omitempty"` + Refusal *string `json:"refusal,omitempty"` + Audio json.RawMessage `json:"audio,omitempty"` + FunctionCall json.RawMessage `json:"function_call,omitempty"` Role string `json:"role,omitempty"` ToolCalls []ToolCallResponse `json:"tool_calls,omitempty"` } @@ -366,6 +370,7 @@ func ResponsesArgumentsString(arguments json.RawMessage) string { type ResponsesOutputContent struct { Type string `json:"type"` Text string `json:"text"` + Refusal string `json:"refusal,omitempty"` Annotations []interface{} `json:"annotations"` } diff --git a/output/posters/newapi-40k-stars-light.png b/output/posters/newapi-40k-stars-light.png deleted file mode 100644 index c13fa2de5a2e..000000000000 Binary files a/output/posters/newapi-40k-stars-light.png and /dev/null differ diff --git a/output/posters/newapi-40k-stars-light.svg b/output/posters/newapi-40k-stars-light.svg deleted file mode 100644 index 965b71e59a3b..000000000000 --- a/output/posters/newapi-40k-stars-light.svg +++ /dev/null @@ -1,111 +0,0 @@ - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - NewAPI - - 40K - - Stars - - Thank you, builders - - - - - NewAPI - - - - - newapi.ai - - - - - - - - diff --git a/relay/channel/claude/relay-claude.go b/relay/channel/claude/relay-claude.go index 9d1beae3b153..6f45b5ae11aa 100644 --- a/relay/channel/claude/relay-claude.go +++ b/relay/channel/claude/relay-claude.go @@ -15,6 +15,7 @@ import ( relaycommon "github.com/QuantumNous/new-api/relay/common" "github.com/QuantumNous/new-api/relay/helper" "github.com/QuantumNous/new-api/relay/reasonmap" + "github.com/QuantumNous/new-api/relay/responsevalidator" "github.com/QuantumNous/new-api/service" "github.com/QuantumNous/new-api/setting/model_setting" "github.com/QuantumNous/new-api/setting/reasoning" @@ -910,22 +911,63 @@ func ClaudeStreamHandler(c *gin.Context, resp *http.Response, info *relaycommon. Usage: &dto.Usage{}, } var streamErr *types.NewAPIError - var hasValidResponse bool + streamState := responsevalidator.NewClaudeStreamState() + pendingData := make([]string, 0, 3) + streamCommitted := false helper.StreamScannerHandler(c, resp, info, func(data string, sr *helper.StreamResult) { + var event dto.ClaudeResponse + if err := common.UnmarshalJsonStr(data, &event); err != nil { + streamErr = types.NewError(err, types.ErrorCodeBadResponseBody) + sr.Stop(err) + return + } + if claudeError := event.GetClaudeError(); claudeError != nil && claudeError.Type != "" { + streamErr = types.WithClaudeError(*claudeError, http.StatusInternalServerError) + sr.Stop(streamErr) + return + } + if err := streamState.Observe(&event); err != nil { + streamErr = types.NewError(err, types.ErrorCodeBadResponseBody) + sr.Stop(err) + return + } + if !streamCommitted { + pendingData = append(pendingData, data) + if !streamState.Valid() { + return + } + for _, pending := range pendingData { + streamErr = HandleStreamResponseData(c, info, claudeInfo, pending) + if streamErr != nil { + sr.Stop(streamErr) + return + } + } + pendingData = nil + streamCommitted = true + return + } streamErr = HandleStreamResponseData(c, info, claudeInfo, data) if streamErr != nil { sr.Stop(streamErr) - return } - hasValidResponse = true }) if streamErr != nil { return nil, streamErr } - - // 检查流式响应是否收到有效内容,如果没有则返回错误以触发重试 - if !hasValidResponse { - return nil, types.NewOpenAIError(fmt.Errorf("empty response from Claude API"), types.ErrorCodeEmptyResponse, http.StatusInternalServerError) + if err := streamState.Validate(); err != nil { + code := types.ErrorCodeBadResponse + if !streamState.Valid() { + code = types.ErrorCodeEmptyResponse + } + return nil, types.NewOpenAIError(err, code, http.StatusInternalServerError) + } + if !streamCommitted { + for _, pending := range pendingData { + if apiErr := HandleStreamResponseData(c, info, claudeInfo, pending); apiErr != nil { + return nil, apiErr + } + } } HandleStreamFinalResponse(c, info, claudeInfo) @@ -943,9 +985,11 @@ func HandleClaudeResponseData(c *gin.Context, info *relaycommon.RelayInfo, claud } // 记录拒绝原因(上游功能) maybeMarkClaudeRefusal(c, claudeResponse.StopReason) - // 检查是否有有效的内容返回,如果没有则返回错误以触发重试 - stopReason := strings.TrimSpace(claudeResponse.StopReason) - if len(claudeResponse.Content) == 0 && claudeResponse.Completion == "" && stopReason != "tool_use" && stopReason != "tool_calls" { + validity, validationErr := responsevalidator.ClaudeResponse(&claudeResponse) + if validationErr != nil { + return types.NewOpenAIError(validationErr, types.ErrorCodeBadResponseBody, http.StatusInternalServerError) + } + if !validity.Valid() { return types.NewOpenAIError(fmt.Errorf("empty response from Claude API"), types.ErrorCodeEmptyResponse, http.StatusInternalServerError) } if claudeInfo.Usage == nil { diff --git a/relay/channel/claude/relay_claude_tool_use_empty_content_test.go b/relay/channel/claude/relay_claude_tool_use_empty_content_test.go index 6e96b0ae6cd6..4f4168dd11d3 100644 --- a/relay/channel/claude/relay_claude_tool_use_empty_content_test.go +++ b/relay/channel/claude/relay_claude_tool_use_empty_content_test.go @@ -13,7 +13,7 @@ import ( "github.com/stretchr/testify/require" ) -func TestHandleClaudeResponseData_AllowsToolUseStopReasonWithEmptyContent(t *testing.T) { +func TestHandleClaudeResponseData_RejectsToolUseStopReasonWithEmptyContent(t *testing.T) { t.Parallel() gin.SetMode(gin.TestMode) @@ -38,5 +38,6 @@ func TestHandleClaudeResponseData_AllowsToolUseStopReasonWithEmptyContent(t *tes resp := &http.Response{StatusCode: http.StatusOK, Header: make(http.Header)} apiErr := HandleClaudeResponseData(c, info, claudeInfo, resp, data) - require.Nil(t, apiErr) + require.NotNil(t, apiErr) + require.Equal(t, types.ErrorCodeEmptyResponse, apiErr.GetErrorCode()) } diff --git a/relay/channel/gemini/relay-gemini-native.go b/relay/channel/gemini/relay-gemini-native.go index 005151e34d97..5cd5ba84ef28 100644 --- a/relay/channel/gemini/relay-gemini-native.go +++ b/relay/channel/gemini/relay-gemini-native.go @@ -12,6 +12,7 @@ import ( "github.com/QuantumNous/new-api/logger" relaycommon "github.com/QuantumNous/new-api/relay/common" "github.com/QuantumNous/new-api/relay/helper" + "github.com/QuantumNous/new-api/relay/responsevalidator" "github.com/QuantumNous/new-api/service" "github.com/QuantumNous/new-api/types" @@ -35,18 +36,18 @@ func GeminiTextGenerationHandler(c *gin.Context, info *relaycommon.RelayInfo, re if err != nil { return nil, types.NewOpenAIError(err, types.ErrorCodeBadResponseBody, http.StatusInternalServerError) } + if geminiResponse.PromptFeedback != nil && geminiResponse.PromptFeedback.BlockReason != nil { + common.SetContextKey(c, constant.ContextKeyAdminRejectReason, fmt.Sprintf("gemini_block_reason=%s", *geminiResponse.PromptFeedback.BlockReason)) + return nil, types.NewOpenAIError(errors.New("request blocked by Gemini API: "+*geminiResponse.PromptFeedback.BlockReason), types.ErrorCodePromptBlocked, http.StatusBadRequest, types.ErrOptionWithSkipRetry()) + } - // 检查是否有候选返回,如果没有则返回错误以触发重试 - if len(geminiResponse.Candidates) == 0 { - if geminiResponse.PromptFeedback != nil && geminiResponse.PromptFeedback.BlockReason != nil { - // 记录拒绝原因(上游功能) - common.SetContextKey(c, constant.ContextKeyAdminRejectReason, fmt.Sprintf("gemini_block_reason=%s", *geminiResponse.PromptFeedback.BlockReason)) - return nil, types.NewOpenAIError(errors.New("request blocked by Gemini API: "+*geminiResponse.PromptFeedback.BlockReason), types.ErrorCodePromptBlocked, http.StatusBadRequest, types.ErrOptionWithSkipRetry()) - } else { - // 记录空响应原因(上游功能) - common.SetContextKey(c, constant.ContextKeyAdminRejectReason, "gemini_empty_candidates") - return nil, types.NewOpenAIError(errors.New("empty response from Gemini API"), types.ErrorCodeEmptyResponse, http.StatusInternalServerError) - } + validity, validationErr := responsevalidator.GeminiResponse(&geminiResponse) + if validationErr != nil { + return nil, types.NewOpenAIError(validationErr, types.ErrorCodeBadResponseBody, http.StatusInternalServerError) + } + if !validity.Valid() { + common.SetContextKey(c, constant.ContextKeyAdminRejectReason, "gemini_empty_candidates") + return nil, types.NewOpenAIError(errors.New("empty response from Gemini API"), types.ErrorCodeEmptyResponse, http.StatusInternalServerError) } // 计算使用量(基于 UsageMetadata) @@ -91,7 +92,28 @@ func NativeGeminiEmbeddingHandler(c *gin.Context, resp *http.Response, info *rel func GeminiTextGenerationStreamHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Response) (*dto.Usage, *types.NewAPIError) { helper.SetEventStreamHeaders(c) + var validationErr error + var validationAPIError *types.NewAPIError + sawTerminal := false usage, err := geminiStreamHandler(c, info, resp, func(data string, geminiResponse *dto.GeminiChatResponse) bool { + if geminiResponse.PromptFeedback != nil && geminiResponse.PromptFeedback.BlockReason != nil { + validationAPIError = types.NewOpenAIError( + errors.New("request blocked by Gemini API: "+*geminiResponse.PromptFeedback.BlockReason), + types.ErrorCodePromptBlocked, + http.StatusBadRequest, + types.ErrOptionWithSkipRetry(), + ) + return false + } + validity, currentErr := responsevalidator.GeminiResponse(geminiResponse) + if currentErr != nil { + validationErr = currentErr + return false + } + sawTerminal = sawTerminal || validity.Terminal + if !validity.Valid() { + return true + } err := helper.StringData(c, data) if err != nil { logger.LogError(c, "failed to write stream data: "+err.Error()) @@ -104,10 +126,19 @@ func GeminiTextGenerationStreamHandler(c *gin.Context, info *relaycommon.RelayIn if err != nil { return usage, err } + if validationAPIError != nil { + return nil, validationAPIError + } + if validationErr != nil { + return nil, types.NewOpenAIError(validationErr, types.ErrorCodeBadResponseBody, http.StatusInternalServerError) + } if info.SendResponseCount == 0 { return nil, types.NewOpenAIError(errors.New("empty response from Gemini API"), types.ErrorCodeEmptyResponse, http.StatusInternalServerError) } + if !sawTerminal { + return nil, types.NewOpenAIError(errors.New("Gemini stream ended before a terminal finish reason"), types.ErrorCodeBadResponse, http.StatusInternalServerError) + } return usage, nil } diff --git a/relay/channel/gemini/relay-gemini.go b/relay/channel/gemini/relay-gemini.go index 403912c3242f..10db98a89a6b 100644 --- a/relay/channel/gemini/relay-gemini.go +++ b/relay/channel/gemini/relay-gemini.go @@ -19,6 +19,7 @@ import ( "github.com/QuantumNous/new-api/relay/channel/openai" relaycommon "github.com/QuantumNous/new-api/relay/common" "github.com/QuantumNous/new-api/relay/helper" + "github.com/QuantumNous/new-api/relay/responsevalidator" "github.com/QuantumNous/new-api/service" "github.com/QuantumNous/new-api/setting/model_setting" "github.com/QuantumNous/new-api/setting/reasoning" @@ -1345,10 +1346,12 @@ func geminiStreamHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http var usage = &dto.Usage{} var imageCount int responseText := strings.Builder{} + var scanErr error helper.StreamScannerHandler(c, resp, info, func(data string, sr *helper.StreamResult) { var geminiResponse dto.GeminiChatResponse if err := common.UnmarshalJsonStr(data, &geminiResponse); err != nil { + scanErr = err sr.Stop(fmt.Errorf("unmarshal: %w", err)) return } @@ -1385,6 +1388,9 @@ func geminiStreamHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http sr.Stop(fmt.Errorf("gemini callback stopped")) } }) + if scanErr != nil { + return nil, types.NewOpenAIError(scanErr, types.ErrorCodeBadResponseBody, http.StatusInternalServerError) + } if imageCount != 0 { if usage.CompletionTokens == 0 { @@ -1409,10 +1415,28 @@ func GeminiChatStreamHandler(c *gin.Context, info *relaycommon.RelayInfo, resp * finishReason := constant.FinishReasonStop toolCallIndexByChoice := make(map[int]map[string]int) nextToolCallIndexByChoice := make(map[int]int) + var validationErr error + var validationAPIError *types.NewAPIError + sawTerminal := false usage, err := geminiStreamHandler(c, info, resp, func(data string, geminiResponse *dto.GeminiChatResponse) bool { - // 跳过空的 Candidates 响应 - if len(geminiResponse.Candidates) == 0 { + if geminiResponse.PromptFeedback != nil && geminiResponse.PromptFeedback.BlockReason != nil { + validationAPIError = types.NewOpenAIError( + errors.New("request blocked by Gemini API: "+*geminiResponse.PromptFeedback.BlockReason), + types.ErrorCodePromptBlocked, + http.StatusBadRequest, + types.ErrOptionWithSkipRetry(), + ) + return false + } + validity, currentErr := responsevalidator.GeminiResponse(geminiResponse) + if currentErr != nil { + validationErr = currentErr + return false + } + sawTerminal = sawTerminal || validity.Terminal + // Usage-only and protocol-shell chunks are not downstream output. + if !validity.Valid() { return true } @@ -1499,10 +1523,19 @@ func GeminiChatStreamHandler(c *gin.Context, info *relaycommon.RelayInfo, resp * if err != nil { return usage, err } + if validationAPIError != nil { + return nil, validationAPIError + } + if validationErr != nil { + return nil, types.NewOpenAIError(validationErr, types.ErrorCodeBadResponseBody, http.StatusInternalServerError) + } if info.SendResponseCount == 0 { return nil, types.NewOpenAIError(errors.New("empty response from Gemini API"), types.ErrorCodeEmptyResponse, http.StatusInternalServerError) } + if !sawTerminal { + return nil, types.NewOpenAIError(errors.New("Gemini stream ended before a terminal finish reason"), types.ErrorCodeBadResponse, http.StatusInternalServerError) + } response := helper.GenerateFinalUsageResponse(id, createAt, info.UpstreamModelName, *usage) if info.RelayFormat == types.RelayFormatClaude && info.ClaudeConvertInfo != nil && !info.ClaudeConvertInfo.Done { @@ -1528,16 +1561,20 @@ func GeminiChatHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.R if err != nil { return nil, types.NewOpenAIError(err, types.ErrorCodeBadResponseBody, http.StatusInternalServerError) } - if len(geminiResponse.Candidates) == 0 { - if geminiResponse.PromptFeedback != nil && geminiResponse.PromptFeedback.BlockReason != nil { - common.SetContextKey(c, constant.ContextKeyAdminRejectReason, fmt.Sprintf("gemini_block_reason=%s", *geminiResponse.PromptFeedback.BlockReason)) - return nil, types.NewOpenAIError( - errors.New("request blocked by Gemini API: "+*geminiResponse.PromptFeedback.BlockReason), - types.ErrorCodePromptBlocked, - http.StatusBadRequest, - types.ErrOptionWithSkipRetry(), - ) - } + if geminiResponse.PromptFeedback != nil && geminiResponse.PromptFeedback.BlockReason != nil { + common.SetContextKey(c, constant.ContextKeyAdminRejectReason, fmt.Sprintf("gemini_block_reason=%s", *geminiResponse.PromptFeedback.BlockReason)) + return nil, types.NewOpenAIError( + errors.New("request blocked by Gemini API: "+*geminiResponse.PromptFeedback.BlockReason), + types.ErrorCodePromptBlocked, + http.StatusBadRequest, + types.ErrOptionWithSkipRetry(), + ) + } + validity, validationErr := responsevalidator.GeminiResponse(&geminiResponse) + if validationErr != nil { + return nil, types.NewOpenAIError(validationErr, types.ErrorCodeBadResponseBody, http.StatusInternalServerError) + } + if !validity.Valid() { common.SetContextKey(c, constant.ContextKeyAdminRejectReason, "gemini_empty_candidates") return nil, types.NewOpenAIError( errors.New("empty response from Gemini API"), diff --git a/relay/channel/openai/relay-openai.go b/relay/channel/openai/relay-openai.go index 6f8df0f74622..f78f0b7402af 100644 --- a/relay/channel/openai/relay-openai.go +++ b/relay/channel/openai/relay-openai.go @@ -14,6 +14,7 @@ import ( relaycommon "github.com/QuantumNous/new-api/relay/common" relayconstant "github.com/QuantumNous/new-api/relay/constant" "github.com/QuantumNous/new-api/relay/helper" + "github.com/QuantumNous/new-api/relay/responsevalidator" "github.com/QuantumNous/new-api/service" "github.com/QuantumNous/new-api/types" @@ -119,30 +120,95 @@ func OaiStreamHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Re var usage = &dto.Usage{} var lastStreamData string var secondLastStreamData string // 存储倒数第二个stream data,用于音频模型 + var streamErr *types.NewAPIError + streamState := responsevalidator.NewOpenAIStreamState() + pendingData := make([]string, 0, 2) + streamCommitted := false // 检查是否为音频模型 isAudioModel := strings.Contains(strings.ToLower(model), "audio") helper.StreamScannerHandler(c, resp, info, func(data string, sr *helper.StreamResult) { - if lastStreamData != "" { - if err := HandleStreamFormat(c, info, lastStreamData, info.ChannelSetting.ForceFormat, info.ChannelSetting.ThinkingToContent); err != nil { - common.SysLog("error handling stream format: " + err.Error()) - sr.Error(err) + var chunk dto.ChatCompletionsStreamResponse + if info.RelayMode == relayconstant.RelayModeCompletions { + var completionChunk dto.CompletionsStreamResponse + if err := common.UnmarshalJsonStr(data, &completionChunk); err != nil { + streamErr = types.NewOpenAIError(err, types.ErrorCodeBadResponseBody, http.StatusInternalServerError) + sr.Stop(err) + return } - } - if len(data) > 0 { - // 对音频模型,保存倒数第二个stream data - if isAudioModel && lastStreamData != "" { - secondLastStreamData = lastStreamData + if err := streamState.ObserveCompletion(&completionChunk); err != nil { + streamErr = types.NewOpenAIError(err, types.ErrorCodeBadResponseBody, http.StatusInternalServerError) + sr.Stop(err) + return + } + processCompletionsStreamResponse(completionChunk, &responseTextBuilder) + } else { + if err := common.UnmarshalJsonStr(data, &chunk); err != nil { + streamErr = types.NewOpenAIError(err, types.ErrorCodeBadResponseBody, http.StatusInternalServerError) + sr.Stop(err) + return + } + if err := streamState.Observe(&chunk); err != nil { + streamErr = types.NewOpenAIError(err, types.ErrorCodeBadResponseBody, http.StatusInternalServerError) + sr.Stop(err) + return } + if err := ProcessStreamResponse(chunk, &responseTextBuilder, &toolCount); err != nil { + streamErr = types.NewOpenAIError(err, types.ErrorCodeBadResponseBody, http.StatusInternalServerError) + sr.Stop(err) + return + } + } - lastStreamData = data - if err := processTokenData(info.RelayMode, data, &responseTextBuilder, &toolCount); err != nil { - logger.LogError(c, "error processing stream token data: "+err.Error()) - sr.Error(err) + if isAudioModel && lastStreamData != "" { + secondLastStreamData = lastStreamData + } + lastStreamData = data + if !streamCommitted { + pendingData = append(pendingData, data) + if !streamState.Valid() { + return } + for _, pending := range pendingData { + if err := HandleStreamFormat(c, info, pending, info.ChannelSetting.ForceFormat, info.ChannelSetting.ThinkingToContent); err != nil { + streamErr = types.NewOpenAIError(err, types.ErrorCodeBadResponseBody, http.StatusInternalServerError) + sr.Stop(err) + return + } + } + pendingData = nil + streamCommitted = true + return + } + if err := HandleStreamFormat(c, info, data, info.ChannelSetting.ForceFormat, info.ChannelSetting.ThinkingToContent); err != nil { + streamErr = types.NewOpenAIError(err, types.ErrorCodeBadResponseBody, http.StatusInternalServerError) + sr.Stop(err) } }) + if streamErr != nil { + return nil, streamErr + } + done := info.StreamStatus != nil && info.StreamStatus.EndReason == relaycommon.StreamEndReasonDone + if err := streamState.Validate(done); err != nil { + code := types.ErrorCodeBadResponse + if !streamState.Valid() { + code = types.ErrorCodeEmptyResponse + } + return nil, types.NewOpenAIError(err, code, http.StatusInternalServerError) + } + if !streamCommitted { + for _, pending := range pendingData { + if err := HandleStreamFormat(c, info, pending, info.ChannelSetting.ForceFormat, info.ChannelSetting.ThinkingToContent); err != nil { + return nil, types.NewOpenAIError(err, types.ErrorCodeBadResponseBody, http.StatusInternalServerError) + } + } + streamCommitted = true + pendingData = nil + } + if semanticToolCount := streamState.ToolCount(); semanticToolCount > toolCount { + toolCount = semanticToolCount + } // 对音频模型,从倒数第二个stream data中提取usage信息 if isAudioModel && secondLastStreamData != "" { @@ -169,12 +235,6 @@ func OaiStreamHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Re logger.LogError(c, fmt.Sprintf("error handling last response: %s, lastStreamData: [%s]", err.Error(), lastStreamData)) } - if info.RelayFormat == types.RelayFormatOpenAI { - if shouldSendLastResp { - _ = sendStreamData(c, info, lastStreamData, info.ChannelSetting.ForceFormat, info.ChannelSetting.ThinkingToContent) - } - } - if !containStreamUsage { usage = service.ResponseText2Usage(c, responseTextBuilder.String(), info.UpstreamModelName, info.GetEstimatePromptTokens()) usage.CompletionTokens += toolCount * 7 @@ -182,11 +242,6 @@ func OaiStreamHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Re applyUsagePostProcessing(info, usage, common.StringToByteSlice(lastStreamData)) - // 检查流式响应是否收到有效内容,如果没有则返回错误以触发重试 - if lastStreamData == "" && responseTextBuilder.Len() == 0 && !containStreamUsage { - return nil, types.NewOpenAIError(fmt.Errorf("empty response from upstream API"), types.ErrorCodeEmptyResponse, http.StatusInternalServerError) - } - HandleFinalResponse(c, info, lastStreamData, responseId, createAt, model, systemFingerprint, usage, containStreamUsage) return usage, nil @@ -226,14 +281,29 @@ func OpenaiHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Respo return nil, types.WithOpenAIError(*oaiError, resp.StatusCode) } - // 检查是否有有效的 choices 返回,如果没有则返回错误以触发重试 - // 仅对聊天补全和补全模式检查,嵌入等模式的响应没有 choices 字段 - if (info.RelayMode == relayconstant.RelayModeChatCompletions || info.RelayMode == relayconstant.RelayModeCompletions) && len(simpleResponse.Choices) == 0 { + // Chat/completions must contain semantic output, a complete tool call, or + // an explicit filtering/refusal result. Some compatible upstreams return a + // Responses API object here, so convert that form before rejecting it. + if info.RelayMode == relayconstant.RelayModeChatCompletions || info.RelayMode == relayconstant.RelayModeCompletions { + validity, validationErr := responsevalidator.OpenAIChatResponse(&simpleResponse) + if validationErr != nil { + return nil, types.NewOpenAIError(validationErr, types.ErrorCodeBadResponseBody, http.StatusInternalServerError) + } + if validity.Valid() { + goto responseValidated + } var responsesResp dto.OpenAIResponsesResponse if respErr := common.Unmarshal(responseBody, &responsesResp); respErr == nil && responsesResp.Object != "" && len(responsesResp.Output) > 0 { chatId := helper.GetResponseID(c) chatResp, usage, convErr := service.ResponsesResponseToChatCompletionsResponse(&responsesResp, chatId) if convErr == nil && chatResp != nil { + convertedValidity, validationErr := responsevalidator.OpenAIChatResponse(chatResp) + if validationErr != nil { + return nil, types.NewOpenAIError(validationErr, types.ErrorCodeBadResponseBody, http.StatusInternalServerError) + } + if !convertedValidity.Valid() { + return nil, types.NewOpenAIError(fmt.Errorf("converted Responses payload contained no semantic output"), types.ErrorCodeEmptyResponse, http.StatusInternalServerError) + } if usage == nil || usage.TotalTokens == 0 { text := service.ExtractOutputTextFromResponses(&responsesResp) usage = service.ResponseText2Usage(c, text, info.UpstreamModelName, info.GetEstimatePromptTokens()) @@ -261,6 +331,7 @@ func OpenaiHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Respo return nil, types.NewOpenAIError(fmt.Errorf("empty response from upstream API"), types.ErrorCodeEmptyResponse, http.StatusInternalServerError) } +responseValidated: // 记录 content_filter 拒绝原因(上游功能) for _, choice := range simpleResponse.Choices { if choice.FinishReason == constant.FinishReasonContentFilter { diff --git a/relay/channel/openai/relay_responses.go b/relay/channel/openai/relay_responses.go index 2665b8d027e9..46905780eca4 100644 --- a/relay/channel/openai/relay_responses.go +++ b/relay/channel/openai/relay_responses.go @@ -11,6 +11,7 @@ import ( "github.com/QuantumNous/new-api/logger" relaycommon "github.com/QuantumNous/new-api/relay/common" "github.com/QuantumNous/new-api/relay/helper" + "github.com/QuantumNous/new-api/relay/responsevalidator" "github.com/QuantumNous/new-api/service" "github.com/QuantumNous/new-api/types" @@ -33,6 +34,13 @@ func OaiResponsesHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http if oaiError := responsesResponse.GetOpenAIError(); oaiError != nil && oaiError.Type != "" { return nil, types.WithOpenAIError(*oaiError, resp.StatusCode) } + validity, validationErr := responsevalidator.ResponsesResponse(&responsesResponse) + if validationErr != nil { + return nil, types.NewOpenAIError(validationErr, types.ErrorCodeBadResponseBody, http.StatusInternalServerError) + } + if !validity.Valid() { + return nil, types.NewOpenAIError(fmt.Errorf("empty response from Responses API"), types.ErrorCodeEmptyResponse, http.StatusInternalServerError) + } if responsesResponse.HasImageGenerationCall() { c.Set("image_generation_call", true) @@ -78,6 +86,14 @@ func OaiResponsesStreamHandler(c *gin.Context, info *relaycommon.RelayInfo, resp var usage = &dto.Usage{} var responseTextBuilder strings.Builder + var streamErr *types.NewAPIError + streamState := responsevalidator.NewResponsesStreamState() + type pendingEvent struct { + data string + response dto.ResponsesStreamResponse + } + pending := make([]pendingEvent, 0, 3) + streamCommitted := false helper.StreamScannerHandler(c, resp, info, func(data string, sr *helper.StreamResult) { @@ -85,10 +101,28 @@ func OaiResponsesStreamHandler(c *gin.Context, info *relaycommon.RelayInfo, resp var streamResponse dto.ResponsesStreamResponse if err := common.UnmarshalJsonStr(data, &streamResponse); err != nil { logger.LogError(c, "failed to unmarshal stream response: "+err.Error()) - sr.Error(err) + streamErr = types.NewOpenAIError(err, types.ErrorCodeBadResponseBody, http.StatusInternalServerError) + sr.Stop(err) return } - sendResponsesStreamData(c, streamResponse, data) + if err := streamState.Observe(&streamResponse); err != nil { + streamErr = types.NewOpenAIError(err, types.ErrorCodeBadResponse, http.StatusInternalServerError) + sr.Stop(err) + return + } + if !streamCommitted { + pending = append(pending, pendingEvent{data: data, response: streamResponse}) + if !streamState.Valid() { + return + } + for _, item := range pending { + sendResponsesStreamData(c, item.response, item.data) + } + pending = nil + streamCommitted = true + } else { + sendResponsesStreamData(c, streamResponse, data) + } switch streamResponse.Type { case "response.completed": if streamResponse.Response != nil { @@ -129,6 +163,21 @@ func OaiResponsesStreamHandler(c *gin.Context, info *relaycommon.RelayInfo, resp } } }) + if streamErr != nil { + return nil, streamErr + } + if err := streamState.Validate(); err != nil { + code := types.ErrorCodeBadResponse + if !streamState.Valid() { + code = types.ErrorCodeEmptyResponse + } + return nil, types.NewOpenAIError(err, code, http.StatusInternalServerError) + } + if !streamCommitted { + for _, item := range pending { + sendResponsesStreamData(c, item.response, item.data) + } + } if usage.CompletionTokens == 0 { // 计算输出文本的 token 数量 diff --git a/relay/common/relay_info.go b/relay/common/relay_info.go index 9f460ce5c6a7..d08d11be7f2a 100644 --- a/relay/common/relay_info.go +++ b/relay/common/relay_info.go @@ -193,6 +193,28 @@ type RelayInfo struct { *TaskRelayInfo } +func (info *RelayInfo) ResetResponseStateForRetry() { + if info == nil { + return + } + info.SendResponseCount = 0 + info.ReceivedResponseCount = 0 + info.StreamStatus = nil + info.FirstResponseTime = time.Time{} + info.isFirstResponse = true + info.ThinkingContentInfo = ThinkingContentInfo{} + if info.ClaudeConvertInfo != nil { + info.ClaudeConvertInfo = &ClaudeConvertInfo{LastMessagesType: LastMessageTypeNone} + } + if info.ResponsesUsageInfo != nil { + for _, tool := range info.ResponsesUsageInfo.BuiltInTools { + if tool != nil { + tool.CallCount = 0 + } + } + } +} + func (info *RelayInfo) InitChannelMeta(c *gin.Context) { channelType := common.GetContextKeyInt(c, constant.ContextKeyChannelType) paramOverride := common.GetContextKeyStringMap(c, constant.ContextKeyChannelParamOverride) diff --git a/relay/common/relay_info_test.go b/relay/common/relay_info_test.go index e53ec804ca06..058684daebbd 100644 --- a/relay/common/relay_info_test.go +++ b/relay/common/relay_info_test.go @@ -2,11 +2,48 @@ package common import ( "testing" + "time" "github.com/QuantumNous/new-api/types" "github.com/stretchr/testify/require" ) +func TestRelayInfoResetResponseStateForRetry(t *testing.T) { + t.Parallel() + + info := &RelayInfo{ + SendResponseCount: 3, + ReceivedResponseCount: 4, + FirstResponseTime: time.Now(), + isFirstResponse: false, + StreamStatus: &StreamStatus{}, + ThinkingContentInfo: ThinkingContentInfo{ + HasSentThinkingContent: true, + }, + ClaudeConvertInfo: &ClaudeConvertInfo{ + LastMessagesType: LastMessageTypeTools, + Done: true, + }, + ResponsesUsageInfo: &ResponsesUsageInfo{ + BuiltInTools: map[string]*BuildInToolInfo{ + "web_search": {CallCount: 2}, + }, + }, + } + + info.ResetResponseStateForRetry() + + require.Zero(t, info.SendResponseCount) + require.Zero(t, info.ReceivedResponseCount) + require.True(t, info.FirstResponseTime.IsZero()) + require.True(t, info.isFirstResponse) + require.Nil(t, info.StreamStatus) + require.False(t, info.ThinkingContentInfo.HasSentThinkingContent) + require.Equal(t, LastMessageTypeNone, info.ClaudeConvertInfo.LastMessagesType) + require.False(t, info.ClaudeConvertInfo.Done) + require.Zero(t, info.ResponsesUsageInfo.BuiltInTools["web_search"].CallCount) +} + func TestRelayInfoGetFinalRequestRelayFormatPrefersExplicitFinal(t *testing.T) { info := &RelayInfo{ RelayFormat: types.RelayFormatOpenAI, diff --git a/relay/helper/common.go b/relay/helper/common.go index 5b118aef8118..2dca408bcbed 100644 --- a/relay/helper/common.go +++ b/relay/helper/common.go @@ -137,6 +137,40 @@ func Done(c *gin.Context) { _ = StringData(c, "[DONE]") } +// StreamError reports a terminal error after an SSE response has already been +// committed. At that point the HTTP status cannot be changed and a transparent +// channel retry would mix two upstream streams. +func StreamError(c *gin.Context, format types.RelayFormat, apiErr *types.NewAPIError) { + if c == nil || apiErr == nil { + return + } + switch format { + case types.RelayFormatClaude: + payload, err := common.Marshal(map[string]any{ + "type": "error", + "error": apiErr.ToClaudeError(), + }) + if err == nil { + c.Render(-1, common.CustomEvent{Data: "event: error\n"}) + c.Render(-1, common.CustomEvent{Data: "data: " + string(payload)}) + _ = FlushWriter(c) + } + case types.RelayFormatOpenAIResponses, types.RelayFormatOpenAIResponsesCompaction: + payload, err := common.Marshal(map[string]any{ + "type": "error", + "error": apiErr.ToOpenAIError(), + }) + if err == nil { + c.Render(-1, common.CustomEvent{Data: "event: error\n"}) + c.Render(-1, common.CustomEvent{Data: "data: " + string(payload)}) + _ = FlushWriter(c) + } + default: + _ = ObjectData(c, map[string]any{"error": apiErr.ToOpenAIError()}) + Done(c) + } +} + func WssString(c *gin.Context, ws *websocket.Conn, str string) error { if ws == nil { logger.LogError(c, "websocket connection is nil") diff --git a/relay/responsevalidator/validator.go b/relay/responsevalidator/validator.go new file mode 100644 index 000000000000..b3024d52c11f --- /dev/null +++ b/relay/responsevalidator/validator.go @@ -0,0 +1,559 @@ +package responsevalidator + +import ( + "encoding/json" + "fmt" + "strings" + + "github.com/QuantumNous/new-api/common" + "github.com/QuantumNous/new-api/dto" +) + +type Result struct { + Output bool + Filtered bool + Terminal bool +} + +func (r Result) Valid() bool { + return r.Output || r.Filtered +} + +func rawPresent(raw json.RawMessage) bool { + value := strings.TrimSpace(string(raw)) + return value != "" && value != "null" && value != "{}" && value != "[]" +} + +func validArguments(arguments string) bool { + arguments = strings.TrimSpace(arguments) + if arguments == "" { + return true + } + var value any + return common.Unmarshal([]byte(arguments), &value) == nil +} + +func ValidateToolCalls(calls []dto.ToolCallResponse) error { + for i, call := range calls { + if strings.TrimSpace(call.Function.Name) == "" { + return fmt.Errorf("tool call %d has no function name", i) + } + if !validArguments(call.Function.Arguments) { + return fmt.Errorf("tool call %d has invalid JSON arguments", i) + } + } + return nil +} + +func parseToolCalls(raw json.RawMessage) ([]dto.ToolCallResponse, error) { + if !rawPresent(raw) { + return nil, nil + } + var calls []dto.ToolCallResponse + if err := common.Unmarshal(raw, &calls); err != nil { + return nil, err + } + return calls, ValidateToolCalls(calls) +} + +func validLegacyFunctionCall(raw json.RawMessage) (bool, error) { + if !rawPresent(raw) { + return false, nil + } + var call struct { + Name string `json:"name"` + Arguments string `json:"arguments"` + } + if err := common.Unmarshal(raw, &call); err != nil { + return false, err + } + if strings.TrimSpace(call.Name) == "" { + return false, fmt.Errorf("legacy function call has no function name") + } + if !validArguments(call.Arguments) { + return false, fmt.Errorf("legacy function call has invalid JSON arguments") + } + return true, nil +} + +func OpenAIChatResponse(response *dto.OpenAITextResponse) (Result, error) { + var result Result + if response == nil { + return result, fmt.Errorf("response is nil") + } + for _, choice := range response.Choices { + if choice.Text != "" || choice.Message.StringContent() != "" || + choice.Message.GetReasoningContent() != "" || + (choice.Message.Refusal != nil && *choice.Message.Refusal != "") || + rawPresent(choice.Message.Audio) { + result.Output = true + } + if strings.EqualFold(choice.FinishReason, "content_filter") { + result.Filtered = true + } + if choice.FinishReason != "" { + result.Terminal = true + } + calls, err := parseToolCalls(choice.Message.ToolCalls) + if err != nil { + return result, err + } + if len(calls) > 0 { + result.Output = true + } + legacy, err := validLegacyFunctionCall(choice.Message.FunctionCall) + if err != nil { + return result, err + } + result.Output = result.Output || legacy + } + return result, nil +} + +type streamedToolCall struct { + name string + arguments strings.Builder +} + +type OpenAIStreamState struct { + Result + tools map[string]*streamedToolCall +} + +func NewOpenAIStreamState() *OpenAIStreamState { + return &OpenAIStreamState{tools: make(map[string]*streamedToolCall)} +} + +func (s *OpenAIStreamState) Observe(response *dto.ChatCompletionsStreamResponse) error { + if s == nil || response == nil { + return fmt.Errorf("response is nil") + } + for _, choice := range response.Choices { + delta := choice.Delta + if delta.GetContentString() != "" || delta.GetReasoningContent() != "" || + (delta.Refusal != nil && *delta.Refusal != "") || rawPresent(delta.Audio) { + s.Output = true + } + if choice.FinishReason != nil && *choice.FinishReason != "" { + s.Terminal = true + if strings.EqualFold(*choice.FinishReason, "content_filter") { + s.Filtered = true + } + } + if legacy, err := validLegacyFunctionCall(delta.FunctionCall); err != nil { + return err + } else if legacy { + s.Output = true + } + for position, call := range delta.ToolCalls { + index := position + if call.Index != nil { + index = *call.Index + } + key := fmt.Sprintf("%d:%d", choice.Index, index) + tool := s.tools[key] + if tool == nil { + tool = &streamedToolCall{} + s.tools[key] = tool + } + if call.Function.Name != "" { + tool.name = call.Function.Name + } + tool.arguments.WriteString(call.Function.Arguments) + } + } + return nil +} + +func (s *OpenAIStreamState) ObserveCompletion(response *dto.CompletionsStreamResponse) error { + if s == nil || response == nil { + return fmt.Errorf("response is nil") + } + for _, choice := range response.Choices { + if choice.Text != "" { + s.Output = true + } + if choice.FinishReason != "" { + s.Terminal = true + if strings.EqualFold(choice.FinishReason, "content_filter") { + s.Filtered = true + } + } + } + return nil +} + +func (s *OpenAIStreamState) Validate(done bool) error { + if s == nil { + return fmt.Errorf("stream state is nil") + } + for key, tool := range s.tools { + if strings.TrimSpace(tool.name) == "" { + return fmt.Errorf("tool call %s has no function name", key) + } + if !validArguments(tool.arguments.String()) { + return fmt.Errorf("tool call %s has incomplete JSON arguments", key) + } + s.Output = true + } + if !s.Valid() { + return fmt.Errorf("stream contained no semantic output") + } + if !s.Terminal && !done { + return fmt.Errorf("stream ended before a terminal event") + } + return nil +} + +func (s *OpenAIStreamState) ToolCount() int { + if s == nil { + return 0 + } + return len(s.tools) +} + +func ClaudeResponse(response *dto.ClaudeResponse) (Result, error) { + var result Result + if response == nil { + return result, fmt.Errorf("response is nil") + } + if response.Completion != "" { + result.Output = true + } + if response.StopReason != "" { + result.Terminal = true + if strings.Contains(strings.ToLower(response.StopReason), "refusal") { + result.Filtered = true + } + } + for i, block := range response.Content { + switch block.Type { + case "text": + result.Output = result.Output || block.GetText() != "" + case "thinking", "redacted_thinking": + result.Output = true + case "tool_use": + if strings.TrimSpace(block.Name) == "" { + return result, fmt.Errorf("tool_use block %d has no name", i) + } + result.Output = true + case "server_tool_use", "web_search_tool_result": + result.Output = true + default: + if block.Type != "" { + result.Output = true + } + } + } + return result, nil +} + +type ClaudeStreamState struct { + Result + tools map[int]*streamedToolCall +} + +func NewClaudeStreamState() *ClaudeStreamState { + return &ClaudeStreamState{tools: make(map[int]*streamedToolCall)} +} + +func (s *ClaudeStreamState) Observe(response *dto.ClaudeResponse) error { + if s == nil || response == nil { + return fmt.Errorf("response is nil") + } + switch response.Type { + case "content_block_start": + if response.ContentBlock == nil { + return fmt.Errorf("content_block_start has no content block") + } + block := response.ContentBlock + switch block.Type { + case "tool_use": + if strings.TrimSpace(block.Name) == "" { + return fmt.Errorf("tool_use block has no name") + } + index := response.GetIndex() + s.tools[index] = &streamedToolCall{name: block.Name} + case "text": + s.Output = s.Output || block.GetText() != "" + case "thinking", "redacted_thinking", "server_tool_use", "web_search_tool_result": + s.Output = true + } + case "content_block_delta": + if response.Delta == nil { + return fmt.Errorf("content_block_delta has no delta") + } + if response.Delta.Text != nil && *response.Delta.Text != "" { + s.Output = true + } + if response.Delta.Thinking != nil && *response.Delta.Thinking != "" { + s.Output = true + } + if response.Delta.Type == "input_json_delta" { + index := response.GetIndex() + tool := s.tools[index] + if tool == nil { + return fmt.Errorf("tool arguments received before tool_use block %d", index) + } + if response.Delta.PartialJson != nil { + tool.arguments.WriteString(*response.Delta.PartialJson) + } + } + case "message_delta": + if response.Delta != nil && response.Delta.StopReason != nil { + s.Terminal = true + if strings.Contains(strings.ToLower(*response.Delta.StopReason), "refusal") { + s.Filtered = true + } + } + case "message_stop": + s.Terminal = true + } + return nil +} + +func (s *ClaudeStreamState) Validate() error { + if s == nil { + return fmt.Errorf("stream state is nil") + } + for index, tool := range s.tools { + if strings.TrimSpace(tool.name) == "" { + return fmt.Errorf("tool call %d has no name", index) + } + if !validArguments(tool.arguments.String()) { + return fmt.Errorf("tool call %d has incomplete JSON arguments", index) + } + s.Output = true + } + if !s.Valid() { + return fmt.Errorf("stream contained no semantic output") + } + if !s.Terminal { + return fmt.Errorf("stream ended before message_stop") + } + return nil +} + +func GeminiResponse(response *dto.GeminiChatResponse) (Result, error) { + var result Result + if response == nil { + return result, fmt.Errorf("response is nil") + } + if response.PromptFeedback != nil && response.PromptFeedback.BlockReason != nil { + result.Filtered = true + result.Terminal = true + } + for candidateIndex, candidate := range response.Candidates { + if candidate.FinishReason != nil && *candidate.FinishReason != "" { + result.Terminal = true + switch strings.ToUpper(*candidate.FinishReason) { + case "SAFETY", "RECITATION", "BLOCKLIST", "PROHIBITED_CONTENT", "SPII": + result.Filtered = true + } + } + for partIndex, part := range candidate.Content.Parts { + switch { + case part.FunctionCall != nil: + if strings.TrimSpace(part.FunctionCall.FunctionName) == "" { + return result, fmt.Errorf("candidate %d part %d function call has no name", candidateIndex, partIndex) + } + result.Output = true + case part.Text != "": + result.Output = true + case part.InlineData != nil && part.InlineData.Data != "": + result.Output = true + case part.FileData != nil && part.FileData.FileUri != "": + result.Output = true + case part.ExecutableCode != nil && part.ExecutableCode.Code != "": + result.Output = true + case part.CodeExecutionResult != nil && (part.CodeExecutionResult.Output != "" || part.CodeExecutionResult.Outcome != ""): + result.Output = true + } + } + } + return result, nil +} + +func ResponsesResponse(response *dto.OpenAIResponsesResponse) (Result, error) { + var result Result + if response == nil { + return result, fmt.Errorf("response is nil") + } + var status string + if len(response.Status) > 0 { + _ = common.Unmarshal(response.Status, &status) + } + switch status { + case "completed": + result.Terminal = true + case "failed", "cancelled": + return result, fmt.Errorf("responses API returned status %s", status) + case "incomplete": + result.Terminal = true + if response.IncompleteDetails != nil && response.IncompleteDetails.Reason == "content_filter" { + result.Filtered = true + } + } + for i, output := range response.Output { + switch output.Type { + case "function_call", "custom_tool_call": + if strings.TrimSpace(output.Name) == "" { + return result, fmt.Errorf("output item %d tool call has no name", i) + } + if !validArguments(output.ArgumentsString()) { + return result, fmt.Errorf("output item %d tool call has invalid JSON arguments", i) + } + result.Output = true + case "image_generation_call", "computer_call", "web_search_call", "file_search_call", "code_interpreter_call": + result.Output = true + default: + for _, content := range output.Content { + if content.Text != "" || content.Refusal != "" { + result.Output = true + } + } + } + } + return result, nil +} + +type ResponsesStreamState struct { + Result + tools map[string]*streamedToolCall +} + +func NewResponsesStreamState() *ResponsesStreamState { + return &ResponsesStreamState{tools: make(map[string]*streamedToolCall)} +} + +func responsesToolKey(event *dto.ResponsesStreamResponse) string { + if event.ItemID != "" { + return event.ItemID + } + if event.Item != nil && event.Item.ID != "" { + return event.Item.ID + } + if event.OutputIndex != nil { + return fmt.Sprintf("output:%d", *event.OutputIndex) + } + return "tool:0" +} + +func responsesToolKeys(event *dto.ResponsesStreamResponse) []string { + keys := make([]string, 0, 3) + if event.ItemID != "" { + keys = append(keys, event.ItemID) + } + if event.Item != nil && event.Item.ID != "" { + keys = append(keys, event.Item.ID) + } + if event.OutputIndex != nil { + keys = append(keys, fmt.Sprintf("output:%d", *event.OutputIndex)) + } + if len(keys) == 0 { + keys = append(keys, responsesToolKey(event)) + } + return keys +} + +func (s *ResponsesStreamState) toolFor(event *dto.ResponsesStreamResponse) *streamedToolCall { + keys := responsesToolKeys(event) + var tool *streamedToolCall + for _, key := range keys { + if s.tools[key] != nil { + tool = s.tools[key] + break + } + } + if tool == nil { + tool = &streamedToolCall{} + } + for _, key := range keys { + s.tools[key] = tool + } + return tool +} + +func (s *ResponsesStreamState) Observe(event *dto.ResponsesStreamResponse) error { + if s == nil || event == nil { + return fmt.Errorf("response event is nil") + } + switch event.Type { + case "response.failed", "response.error": + return fmt.Errorf("responses stream returned %s", event.Type) + case "response.completed", "response.incomplete": + s.Terminal = true + if event.Response != nil { + result, err := ResponsesResponse(event.Response) + if err != nil { + return err + } + s.Output = s.Output || result.Output + s.Filtered = s.Filtered || result.Filtered + } + case "response.output_text.delta", "response.reasoning_summary_text.delta", "response.reasoning_text.delta": + if event.Delta != "" { + s.Output = true + } + case "response.output_item.added", "response.output_item.done": + if event.Item == nil { + return nil + } + switch event.Item.Type { + case "function_call", "custom_tool_call": + if strings.TrimSpace(event.Item.Name) == "" { + return fmt.Errorf("stream tool call has no name") + } + tool := s.toolFor(event) + tool.name = event.Item.Name + arguments := event.Item.ArgumentsString() + if arguments != "" && tool.arguments.Len() == 0 { + tool.arguments.WriteString(arguments) + } + case "image_generation_call", "computer_call", "web_search_call", "file_search_call", "code_interpreter_call": + s.Output = true + default: + for _, content := range event.Item.Content { + if content.Text != "" || content.Refusal != "" { + s.Output = true + } + } + } + case "response.function_call_arguments.delta", "response.custom_tool_call_input.delta": + tool := s.toolFor(event) + tool.arguments.WriteString(event.Delta) + case "response.function_call_arguments.done", "response.custom_tool_call_input.done": + tool := s.toolFor(event) + if tool.arguments.Len() == 0 { + tool.arguments.WriteString(event.Delta) + } + } + return nil +} + +func (s *ResponsesStreamState) Validate() error { + if s == nil { + return fmt.Errorf("stream state is nil") + } + seen := make(map[*streamedToolCall]bool) + for key, tool := range s.tools { + if seen[tool] { + continue + } + seen[tool] = true + if strings.TrimSpace(tool.name) == "" { + return fmt.Errorf("tool call %s has no name", key) + } + if !validArguments(tool.arguments.String()) { + return fmt.Errorf("tool call %s has incomplete JSON arguments", key) + } + s.Output = true + } + if !s.Valid() { + return fmt.Errorf("stream contained no semantic output") + } + if !s.Terminal { + return fmt.Errorf("stream ended before response.completed or response.incomplete") + } + return nil +} diff --git a/relay/responsevalidator/validator_test.go b/relay/responsevalidator/validator_test.go new file mode 100644 index 000000000000..e0c3d2888c71 --- /dev/null +++ b/relay/responsevalidator/validator_test.go @@ -0,0 +1,163 @@ +package responsevalidator + +import ( + "testing" + + "github.com/QuantumNous/new-api/common" + "github.com/QuantumNous/new-api/dto" + "github.com/stretchr/testify/require" +) + +func decode[T any](t *testing.T, payload string) *T { + t.Helper() + var value T + require.NoError(t, common.Unmarshal([]byte(payload), &value)) + return &value +} + +func TestOpenAIChatResponseSemanticOutput(t *testing.T) { + t.Parallel() + + tests := []struct { + name string + payload string + valid bool + wantErr bool + }{ + {"empty choices", `{"choices":[]}`, false, false}, + {"empty assistant shell", `{"choices":[{"message":{"role":"assistant"},"finish_reason":"stop"}]}`, false, false}, + {"text", `{"choices":[{"message":{"content":"ok"},"finish_reason":"stop"}]}`, true, false}, + {"completion text", `{"choices":[{"text":"ok","finish_reason":"stop"}]}`, true, false}, + {"refusal", `{"choices":[{"message":{"refusal":"no"},"finish_reason":"stop"}]}`, true, false}, + {"filtered", `{"choices":[{"message":{},"finish_reason":"content_filter"}]}`, true, false}, + {"tool only", `{"choices":[{"message":{"tool_calls":[{"type":"function","function":{"name":"lookup","arguments":"{}"}}]},"finish_reason":"tool_calls"}]}`, true, false}, + {"tool missing name", `{"choices":[{"message":{"tool_calls":[{"type":"function","function":{"arguments":"{}"}}]},"finish_reason":"tool_calls"}]}`, false, true}, + {"tool malformed arguments", `{"choices":[{"message":{"tool_calls":[{"type":"function","function":{"name":"lookup","arguments":"{"}}]},"finish_reason":"tool_calls"}]}`, false, true}, + } + + for _, test := range tests { + test := test + t.Run(test.name, func(t *testing.T) { + t.Parallel() + result, err := OpenAIChatResponse(decode[dto.OpenAITextResponse](t, test.payload)) + if test.wantErr { + require.Error(t, err) + return + } + require.NoError(t, err) + require.Equal(t, test.valid, result.Valid()) + }) + } +} + +func TestOpenAIStreamStateToolOnlyAndTerminalRules(t *testing.T) { + t.Parallel() + + state := NewOpenAIStreamState() + require.NoError(t, state.Observe(decode[dto.ChatCompletionsStreamResponse](t, + `{"choices":[{"index":0,"delta":{"role":"assistant"}}]}`))) + require.Error(t, state.Validate(true)) + + state = NewOpenAIStreamState() + require.NoError(t, state.Observe(decode[dto.ChatCompletionsStreamResponse](t, + `{"choices":[{"index":0,"delta":{"tool_calls":[{"index":0,"function":{"name":"lookup","arguments":"{"}}]}}]}`))) + require.NoError(t, state.Observe(decode[dto.ChatCompletionsStreamResponse](t, + `{"choices":[{"index":0,"delta":{"tool_calls":[{"index":0,"function":{"arguments":"}"}}]},"finish_reason":"tool_calls"}]}`))) + require.NoError(t, state.Validate(false)) + require.Equal(t, 1, state.ToolCount()) + + incomplete := NewOpenAIStreamState() + require.NoError(t, incomplete.Observe(decode[dto.ChatCompletionsStreamResponse](t, + `{"choices":[{"index":0,"delta":{"tool_calls":[{"index":0,"function":{"name":"lookup","arguments":"{"}}]},"finish_reason":"tool_calls"}]}`))) + require.Error(t, incomplete.Validate(false)) + + unterminated := NewOpenAIStreamState() + require.NoError(t, unterminated.Observe(decode[dto.ChatCompletionsStreamResponse](t, + `{"choices":[{"index":0,"delta":{"content":"partial"}}]}`))) + require.Error(t, unterminated.Validate(false)) + require.NoError(t, unterminated.Validate(true)) +} + +func TestClaudeSemanticOutput(t *testing.T) { + t.Parallel() + + result, err := ClaudeResponse(decode[dto.ClaudeResponse](t, + `{"type":"message","stop_reason":"tool_use","content":[]}`)) + require.NoError(t, err) + require.False(t, result.Valid()) + + result, err = ClaudeResponse(decode[dto.ClaudeResponse](t, + `{"type":"message","stop_reason":"tool_use","content":[{"type":"tool_use","name":"lookup","input":{}}]}`)) + require.NoError(t, err) + require.True(t, result.Valid()) + + state := NewClaudeStreamState() + require.NoError(t, state.Observe(decode[dto.ClaudeResponse](t, `{"type":"message_start"}`))) + require.NoError(t, state.Observe(decode[dto.ClaudeResponse](t, `{"type":"message_stop"}`))) + require.Error(t, state.Validate()) + + state = NewClaudeStreamState() + require.NoError(t, state.Observe(decode[dto.ClaudeResponse](t, + `{"type":"content_block_start","index":0,"content_block":{"type":"tool_use","name":"lookup"}}`))) + require.NoError(t, state.Observe(decode[dto.ClaudeResponse](t, + `{"type":"content_block_delta","index":0,"delta":{"type":"input_json_delta","partial_json":"{}"}}`))) + require.NoError(t, state.Observe(decode[dto.ClaudeResponse](t, `{"type":"message_stop"}`))) + require.NoError(t, state.Validate()) +} + +func TestGeminiSemanticOutput(t *testing.T) { + t.Parallel() + + result, err := GeminiResponse(decode[dto.GeminiChatResponse](t, + `{"candidates":[{"content":{"role":"model","parts":[]},"finishReason":"STOP"}]}`)) + require.NoError(t, err) + require.False(t, result.Valid()) + + result, err = GeminiResponse(decode[dto.GeminiChatResponse](t, + `{"candidates":[{"content":{"parts":[{"functionCall":{"name":"lookup","args":{}}}]},"finishReason":"STOP"}]}`)) + require.NoError(t, err) + require.True(t, result.Valid()) + + _, err = GeminiResponse(decode[dto.GeminiChatResponse](t, + `{"candidates":[{"content":{"parts":[{"functionCall":{"args":{}}}]},"finishReason":"STOP"}]}`)) + require.Error(t, err) + + result, err = GeminiResponse(decode[dto.GeminiChatResponse](t, + `{"candidates":[{"content":{"parts":[]},"finishReason":"SAFETY"}]}`)) + require.NoError(t, err) + require.True(t, result.Valid()) +} + +func TestResponsesSemanticOutput(t *testing.T) { + t.Parallel() + + result, err := ResponsesResponse(decode[dto.OpenAIResponsesResponse](t, + `{"status":"completed","output":[]}`)) + require.NoError(t, err) + require.False(t, result.Valid()) + + result, err = ResponsesResponse(decode[dto.OpenAIResponsesResponse](t, + `{"status":"completed","output":[{"type":"function_call","name":"lookup","arguments":"{}"}]}`)) + require.NoError(t, err) + require.True(t, result.Valid()) + + _, err = ResponsesResponse(decode[dto.OpenAIResponsesResponse](t, + `{"status":"failed","output":[]}`)) + require.Error(t, err) + + state := NewResponsesStreamState() + require.NoError(t, state.Observe(decode[dto.ResponsesStreamResponse](t, + `{"type":"response.created"}`))) + require.NoError(t, state.Observe(decode[dto.ResponsesStreamResponse](t, + `{"type":"response.completed","response":{"status":"completed","output":[]}}`))) + require.Error(t, state.Validate()) + + state = NewResponsesStreamState() + require.NoError(t, state.Observe(decode[dto.ResponsesStreamResponse](t, + `{"type":"response.output_item.added","output_index":0,"item":{"id":"call_1","type":"function_call","name":"lookup"}}`))) + require.NoError(t, state.Observe(decode[dto.ResponsesStreamResponse](t, + `{"type":"response.function_call_arguments.delta","item_id":"call_1","delta":"{}"}`))) + require.NoError(t, state.Observe(decode[dto.ResponsesStreamResponse](t, + `{"type":"response.completed","response":{"status":"completed","output":[]}}`))) + require.NoError(t, state.Validate()) +}