diff --git a/internal/entity/models/minimax.go b/internal/entity/models/minimax.go
index 4dde6a5035..20a08d8d48 100644
--- a/internal/entity/models/minimax.go
+++ b/internal/entity/models/minimax.go
@@ -17,6 +17,7 @@
package models
import (
+ "bufio"
"bytes"
"context"
"encoding/hex"
@@ -52,53 +53,6 @@ func (m *MinimaxModel) Name() string {
return "minimax"
}
-type MinimaxChatResponse struct {
- ID string `json:"id"`
- Object string `json:"object"`
- Created int64 `json:"created"`
- Model string `json:"model"`
- Choices []struct {
- Index int `json:"index"`
- Message struct {
- Role string `json:"role"`
- Content *string `json:"content"`
- ToolCalls []map[string]any `json:"tool_calls"`
- ReasoningContent *string `json:"reasoning_content"`
- } `json:"message"`
- FinishReason string `json:"finish_reason"`
- } `json:"choices"`
- Usage struct {
- PromptTokens int `json:"prompt_tokens"`
- PromptTokensDetails struct {
- CachedTokens int `json:"cached_tokens"`
- } `json:"prompt_tokens_details"`
- CompletionTokens int `json:"completion_tokens"`
- TotalTokens int `json:"total_tokens"`
- TotalCharacters int `json:"total_characters"`
- } `json:"usage"`
- BaseResp struct {
- StatusCode int `json:"status_code"`
- StatusMsg string `json:"status_msg"`
- } `json:"base_resp"`
- Error struct {
- Message string `json:"message"`
- Type string `json:"type"`
- } `json:"error"`
-}
-
-func extractMinimaxChatResponseError(result *MinimaxChatResponse) string {
- if result == nil {
- return ""
- }
- if result.BaseResp.StatusCode != 0 {
- if result.BaseResp.StatusMsg != "" {
- return result.BaseResp.StatusMsg
- }
- return fmt.Sprintf("status_code %d", result.BaseResp.StatusCode)
- }
- return result.Error.Message
-}
-
func validateMinimaxModelName(modelName string) (string, error) {
if strings.TrimSpace(modelName) == "" {
return "", fmt.Errorf("model name is required")
@@ -155,7 +109,6 @@ func (m *MinimaxModel) ChatWithMessages(ctx context.Context, modelName string, m
if err := m.baseModel.APIConfigCheck(apiConfig); err != nil {
return nil, err
}
- apiKey := strings.TrimSpace(*apiConfig.ApiKey)
modelName, err := validateMinimaxModelName(modelName)
if err != nil {
return nil, err
@@ -189,87 +142,23 @@ func (m *MinimaxModel) ChatWithMessages(ctx context.Context, modelName string, m
}
- jsonData, err := json.Marshal(reqBody)
+ body, err := m.baseModel.doRequest(ctx, url, apiConfig, reqBody, nonStreamCallTimeout)
if err != nil {
- return nil, fmt.Errorf("failed to marshal request: %w", err)
+ return nil, err
}
- ctx, cancel := context.WithTimeout(ctx, nonStreamCallTimeout)
- defer cancel()
-
- req, err := http.NewRequestWithContext(ctx, "POST", url, bytes.NewBuffer(jsonData))
- if err != nil {
- return nil, fmt.Errorf("failed to create request: %w", err)
+ // MiniMax can embed errors in a base_resp block with HTTP 200.
+ // Check for these before using the shared handler so the caller
+ // sees the real error message instead of "no choices in response".
+ var result map[string]any
+ if err := json.Unmarshal(body, &result); err != nil {
+ return nil, fmt.Errorf("failed to parse response: %w", err)
+ }
+ if errMsg := extractMinimaxAPIError(result); errMsg != "" {
+ return nil, fmt.Errorf("minimax API error: %s", errMsg)
}
- req.Header.Set("Content-Type", "application/json")
- req.Header.Set("Authorization", fmt.Sprintf("Bearer %s", apiKey))
-
- resp, err := m.baseModel.httpClient.Do(req)
- if err != nil {
- return nil, fmt.Errorf("failed to send request: %w", err)
- }
- defer resp.Body.Close()
-
- body, err := io.ReadAll(resp.Body)
- if err != nil {
- return nil, fmt.Errorf("failed to read response body: %w", err)
- }
-
- if resp.StatusCode != http.StatusOK {
- return nil, fmt.Errorf("minimax API error: status %d: %s", resp.StatusCode, extractMinimaxErrorBody(body))
- }
-
- return parseChatCompletionResponse(body, chatModelConfig, modelUsage, func(body []byte, chatConfig *ChatConfig) (chatResponseParts, error) {
- var result MinimaxChatResponse
- if err := json.Unmarshal(body, &result); err != nil {
- return chatResponseParts{}, fmt.Errorf("failed to unmarshal response: %w", err)
- }
-
- if errMsg := extractMinimaxChatResponseError(&result); errMsg != "" {
- return chatResponseParts{}, fmt.Errorf("minimax API error: %s", errMsg)
- }
- if len(result.Choices) == 0 {
- return chatResponseParts{}, fmt.Errorf("no choices in response")
- }
-
- choice := &result.Choices[0]
- content := ""
- if choice.Message.Content != nil {
- content = *choice.Message.Content
- }
-
- reasonContent := ""
- if chatConfig != nil && chatConfig.Thinking != nil && *chatConfig.Thinking {
- if choice.Message.ReasoningContent != nil {
- reasonContent = *choice.Message.ReasoningContent
- }
- if reasonContent != "" && reasonContent[0] == '\n' {
- reasonContent = reasonContent[1:]
- }
- }
-
- totalTokens := result.Usage.TotalTokens
- if totalTokens == 0 {
- totalTokens = result.Usage.PromptTokens + result.Usage.CompletionTokens
- }
- var usage *TokenUsage
- if totalTokens > 0 {
- usage = &TokenUsage{
- PromptTokens: result.Usage.PromptTokens,
- CompletionTokens: result.Usage.CompletionTokens,
- TotalTokens: totalTokens,
- }
- }
-
- return chatResponseParts{
- RequestID: result.ID,
- Content: &content,
- ReasonContent: &reasonContent,
- ToolCalls: choice.Message.ToolCalls,
- Usage: usage,
- }, nil
- })
+ return HandleNonStreamingResponse(body, modelUsage, chatModelConfig, OpenAIParserConfig)
}
// ChatStreamlyWithSender sends messages and streams response via sender function (best performance, no channel)
@@ -314,6 +203,8 @@ func (m *MinimaxModel) ChatStreamlyWithSender(ctx context.Context, modelName str
}
+ reqBody["stream_options"] = map[string]interface{}{"include_usage": true}
+
jsonData, err := json.Marshal(reqBody)
if err != nil {
return fmt.Errorf("failed to marshal request: %w", err)
@@ -342,75 +233,59 @@ func (m *MinimaxModel) ChatStreamlyWithSender(ctx context.Context, modelName str
return fmt.Errorf("minimax API error: status %d: %s", resp.StatusCode, extractMinimaxErrorBody(body))
}
- // SSE parsing: read line by line
- sawTerminal := false
- accumulatedToolCalls := make(map[int]map[string]any)
- done, err := ParseSSEStream[map[string]interface{}](resp.Body, func(event map[string]interface{}) error {
- tokenUsage, found, usageErr := decodeOpenAICompatibleStreamUsage(event)
- if usageErr != nil {
- return usageErr
- }
- if found {
- applyStreamUsage(modelConfig, modelUsage, tokenUsage)
- }
+ // Pipe the response through a base_resp checker. MiniMax can send
+ // error events (e.g. rate limits) without a choices array, and the
+ // shared handler skips those silently. We surface them so the retry
+ // predicates can match and the caller sees the real reason.
+ pr, pw := io.Pipe()
+ // Close pr when this function returns so an early exit from
+ // HandleStreamingResponse unblocks the producer goroutine below
+ // (its pw.Write fails) and releases resp.Body instead of leaving
+ // the reader blocked on a live pipe.
+ defer pr.Close()
+ streamErr := make(chan error, 1)
+ go func() {
+ defer pw.Close()
+ defer resp.Body.Close()
- choices, ok := event["choices"].([]interface{})
- if !ok || len(choices) == 0 {
- // MiniMax can send an error event (rate limit, etc.)
- // without a choices array. Surface it so the retry
- // predicates can match and the caller sees the real
- // reason instead of a generic "stream ended" error.
- if errMsg := extractMinimaxAPIError(event); errMsg != "" {
- return fmt.Errorf("minimax API error: %s", errMsg)
+ var scanErr error
+ // Ensure streamErr always receives a result, on every exit
+ // path, so the final receive below can never block.
+ defer func() {
+ select {
+ case streamErr <- scanErr:
+ default:
}
- return nil
- }
+ }()
- firstChoice, ok := choices[0].(map[string]interface{})
- if !ok {
- return nil
- }
-
- delta, ok := firstChoice["delta"].(map[string]interface{})
- if !ok {
- return nil
- }
-
- accumulateToolCallDeltas(delta, accumulatedToolCalls)
-
- content, ok := delta["content"].(string)
- if ok && content != "" {
- if err := sender(&content, nil); err != nil {
- return err
+ scanner := bufio.NewScanner(resp.Body)
+ scanner.Buffer(make([]byte, 64*1024), 1024*1024)
+ for scanner.Scan() {
+ line := scanner.Text()
+ if strings.HasPrefix(line, "data:") {
+ data := strings.TrimSpace(line[5:])
+ if data != "" && data != "[DONE]" {
+ var event map[string]any
+ if json.Unmarshal([]byte(data), &event) == nil {
+ if errMsg := extractMinimaxAPIError(event); errMsg != "" {
+ pw.CloseWithError(fmt.Errorf("minimax API error: %s", errMsg))
+ return
+ }
+ }
+ }
+ }
+ if _, err := pw.Write([]byte(line + "\n")); err != nil {
+ scanErr = err
+ return
}
}
+ scanErr = scanner.Err()
+ }()
- reasoningContent, ok := delta["reasoning_content"].(string)
- if ok && reasoningContent != "" {
- if err := sender(nil, &reasoningContent); err != nil {
- return err
- }
- }
-
- finishReason, ok := firstChoice["finish_reason"].(string)
- if ok && finishReason != "" {
- sawTerminal = true
- }
-
- return nil
- })
- if err != nil {
- return fmt.Errorf("failed to scan response body: %w", err)
+ if err := HandleStreamingResponse(pr, modelUsage, modelConfig, OpenAIParserConfig, sender); err != nil {
+ return err
}
- if !done && !sawTerminal {
- return fmt.Errorf("minimax: stream ended before [DONE] or finish_reason")
- }
-
- setSortedToolCallsResult(modelConfig, accumulatedToolCalls)
-
- // Send [DONE] marker for OpenAI compatibility
- endOfStream := "[DONE]"
- return sender(&endOfStream, nil)
+ return <-streamErr
}
// Embed embeds a list of texts into embeddings
diff --git a/internal/entity/models/mistral.go b/internal/entity/models/mistral.go
index 40cf44eaee..217cce4cad 100644
--- a/internal/entity/models/mistral.go
+++ b/internal/entity/models/mistral.go
@@ -33,7 +33,76 @@ import (
"time"
)
-// MistralModel implements ModelDriver for Mistral AI.
+// normalizeMistralStructuredContent rewrites a Mistral magistral response
+// whose message.content is a structured array ([{type:text},{type:thinking,
+// thinking:[{type:text}]}]) into the flat string shape the shared handler
+// expects: content becomes the concatenated text parts and a top-level
+// reasoning_content carries the thinking parts. The rewrite is in-place and
+// only applied when content is actually an array.
+func normalizeMistralStructuredContent(body []byte) []byte {
+ var result map[string]any
+ if err := json.Unmarshal(body, &result); err != nil {
+ return body
+ }
+ choices, ok := result["choices"].([]any)
+ if !ok || len(choices) == 0 {
+ return body
+ }
+ firstChoice, ok := choices[0].(map[string]any)
+ if !ok {
+ return body
+ }
+ messageMap, ok := firstChoice["message"].(map[string]any)
+ if !ok {
+ return body
+ }
+ parts, ok := messageMap["content"].([]any)
+ if !ok {
+ return body
+ }
+
+ var answer, reasoning strings.Builder
+ for _, p := range parts {
+ part, ok := p.(map[string]any)
+ if !ok {
+ continue
+ }
+ switch part["type"] {
+ case "text":
+ if t, ok := part["text"].(string); ok {
+ answer.WriteString(t)
+ }
+ case "thinking":
+ if thinking, ok := part["thinking"].([]any); ok {
+ for _, tp := range thinking {
+ if tpm, ok := tp.(map[string]any); ok {
+ if t, ok := tpm["text"].(string); ok {
+ reasoning.WriteString(t)
+ }
+ }
+ }
+ }
+ }
+ }
+ // Only rewrite if we actually extracted something; otherwise leave
+ // the body untouched so the shared handler surfaces its normal error.
+ if answer.Len() == 0 && reasoning.Len() == 0 {
+ return body
+ }
+
+ messageMap["content"] = answer.String()
+ if existing, ok := messageMap["reasoning_content"].(string); ok && existing != "" {
+ reasoning.WriteString(existing)
+ }
+ if reasoning.Len() > 0 {
+ messageMap["reasoning_content"] = reasoning.String()
+ }
+ out, err := json.Marshal(result)
+ if err != nil {
+ return body
+ }
+ return out
+}
type MistralModel struct {
baseModel BaseModel
@@ -58,33 +127,6 @@ func (m *MistralModel) Name() string {
return "mistral"
}
-type MistralChatResponse struct {
- ID string `json:"id"`
- Object string `json:"object"`
- Created int64 `json:"created"`
- Model string `json:"model"`
- Choices []struct {
- Index int `json:"index"`
- Message struct {
- Role string `json:"role"`
- Content interface{} `json:"content"`
- ToolCalls []map[string]any `json:"tool_calls"`
- ReasoningContent *string `json:"reasoning_content"`
- } `json:"message"`
- FinishReason string `json:"finish_reason"`
- } `json:"choices"`
- Usage struct {
- CompletionTokens int `json:"completion_tokens"`
- NumCachedTokens int `json:"num_cached_tokens"`
- PromptAudioSeconds int `json:"prompt_audio_seconds"`
- PromptTokenDetails struct {
- CachedTokens int `json:"cached_tokens"`
- }
- PromptTokens int `json:"prompt_tokens"`
- TotalTokens int `json:"total_tokens"`
- } `json:"usage"`
-}
-
// ChatWithMessages sends multiple messages with roles and returns the response.
func (m *MistralModel) ChatWithMessages(ctx context.Context, modelName string, messages []Message, apiConfig *APIConfig, chatModelConfig *ChatConfig, modelUsage *common.ModelUsage) (*ChatResponse, error) {
if err := m.baseModel.APIConfigCheck(apiConfig); err != nil {
@@ -99,122 +141,19 @@ func (m *MistralModel) ChatWithMessages(ctx context.Context, modelName string, m
if err != nil {
return nil, err
}
- baseURL = strings.TrimSuffix(baseURL, "/")
url := fmt.Sprintf("%s/%s", baseURL, m.baseModel.URLSuffix.Chat)
reqBody := buildRequestBody(chatModelConfig, modelName, messages, false)
- jsonData, err := json.Marshal(reqBody)
+ body, err := m.baseModel.doRequest(ctx, url, apiConfig, reqBody, nonStreamCallTimeout)
if err != nil {
- return nil, fmt.Errorf("failed to marshal request: %w", err)
+ return nil, err
}
- ctx, cancel := context.WithTimeout(ctx, nonStreamCallTimeout)
- defer cancel()
+ // Mistral magistral returns content as a structured array. Normalize it
+ // to the flat string shape the shared handler understands.
+ body = normalizeMistralStructuredContent(body)
- req, err := http.NewRequestWithContext(ctx, "POST", url, bytes.NewBuffer(jsonData))
- if err != nil {
- return nil, fmt.Errorf("failed to create request: %w", err)
- }
-
- req.Header.Set("Content-Type", "application/json")
- req.Header.Set("Authorization", fmt.Sprintf("Bearer %s", *apiConfig.ApiKey))
-
- resp, err := m.baseModel.httpClient.Do(req)
- if err != nil {
- return nil, fmt.Errorf("failed to send request: %w", err)
- }
- defer resp.Body.Close()
-
- body, err := io.ReadAll(resp.Body)
- if err != nil {
- return nil, fmt.Errorf("failed to read response: %w", err)
- }
-
- if resp.StatusCode != http.StatusOK {
- return nil, fmt.Errorf("API request failed with status %d: %s", resp.StatusCode, string(body))
- }
-
- return parseChatCompletionResponse(body, chatModelConfig, modelUsage, func(body []byte, _ *ChatConfig) (chatResponseParts, error) {
- var result MistralChatResponse
- if err := json.Unmarshal(body, &result); err != nil {
- return chatResponseParts{}, fmt.Errorf("failed to parse response: %w", err)
- }
- if len(result.Choices) == 0 {
- return chatResponseParts{}, fmt.Errorf("no choices in response")
- }
-
- choice := &result.Choices[0]
- content, reasonContent, err := extractMistralContent(choice.Message.Content)
- if err != nil {
- return chatResponseParts{}, err
- }
- if reasonContent == "" && choice.Message.ReasoningContent != nil {
- reasonContent = *choice.Message.ReasoningContent
- }
-
- totalTokens := result.Usage.TotalTokens
- if totalTokens == 0 {
- totalTokens = result.Usage.PromptTokens + result.Usage.CompletionTokens
- }
- var usage *TokenUsage
- if totalTokens > 0 {
- usage = &TokenUsage{
- PromptTokens: result.Usage.PromptTokens,
- CompletionTokens: result.Usage.CompletionTokens,
- TotalTokens: totalTokens,
- }
- }
-
- return chatResponseParts{
- RequestID: result.ID,
- Content: &content,
- ReasonContent: &reasonContent,
- ToolCalls: choice.Message.ToolCalls,
- Usage: usage,
- }, nil
- })
-}
-
-func extractMistralContent(raw interface{}) (string, string, error) {
- switch v := raw.(type) {
- case string:
- return v, "", nil
- case []interface{}:
- var answer, reasoning strings.Builder
- for _, part := range v {
- pm, ok := part.(map[string]interface{})
- if !ok {
- continue
- }
- switch pm["type"] {
- case "text":
- if t, ok := pm["text"].(string); ok {
- answer.WriteString(t)
- }
- case "thinking":
- // thinking is an array of inner text parts; concatenate
- // any inner element with a non-empty text field.
- inner, ok := pm["thinking"].([]interface{})
- if !ok {
- continue
- }
- for _, sub := range inner {
- sm, ok := sub.(map[string]interface{})
- if !ok {
- continue
- }
- if t, ok := sm["text"].(string); ok {
- reasoning.WriteString(t)
- }
- }
- }
- }
- return answer.String(), reasoning.String(), nil
- case nil:
- return "", "", nil
- default:
- return "", "", fmt.Errorf("mistral: unsupported content type %T", raw)
- }
+ return HandleNonStreamingResponse(body, modelUsage, chatModelConfig, OpenAIParserConfig)
}
// ChatStreamlyWithSender sends messages and streams the response
@@ -238,89 +177,15 @@ func (m *MistralModel) ChatStreamlyWithSender(ctx context.Context, modelName str
if err != nil {
return err
}
- baseURL = strings.TrimSuffix(baseURL, "/")
url := fmt.Sprintf("%s/%s", baseURL, m.baseModel.URLSuffix.Chat)
reqBody := buildRequestBody(chatModelConfig, modelName, messages, true)
-
- jsonData, err := json.Marshal(reqBody)
- if err != nil {
- return fmt.Errorf("failed to marshal request: %w", err)
+ reqBody["stream_options"] = map[string]interface{}{
+ "include_usage": true,
}
- req, err := http.NewRequestWithContext(ctx, "POST", url, bytes.NewBuffer(jsonData))
- if err != nil {
- return fmt.Errorf("failed to create request: %w", err)
- }
-
- req.Header.Set("Content-Type", "application/json")
- req.Header.Set("Authorization", fmt.Sprintf("Bearer %s", *apiConfig.ApiKey))
-
- resp, err := m.baseModel.httpClient.Do(req)
- if err != nil {
- return fmt.Errorf("failed to send request: %w", err)
- }
- defer resp.Body.Close()
-
- if resp.StatusCode != http.StatusOK {
- body, _ := io.ReadAll(resp.Body)
- return fmt.Errorf("API request failed with status %d: %s", resp.StatusCode, string(body))
- }
-
- sawTerminal := false
- accumulatedToolCalls := make(map[int]map[string]any)
- done, err := ParseSSEStream[map[string]interface{}](resp.Body, func(event map[string]interface{}) error {
- tokenUsage, found, usageErr := decodeOpenAICompatibleStreamUsage(event)
- if usageErr != nil {
- return usageErr
- }
- if found {
- applyStreamUsage(chatModelConfig, modelUsage, tokenUsage)
- }
-
- choices, ok := event["choices"].([]interface{})
- if !ok || len(choices) == 0 {
- return nil
- }
-
- firstChoice, ok := choices[0].(map[string]interface{})
- if !ok {
- return nil
- }
-
- delta, ok := firstChoice["delta"].(map[string]interface{})
- if !ok {
- return nil
- }
-
- accumulateToolCallDeltas(delta, accumulatedToolCalls)
-
- content, ok := delta["content"].(string)
- if ok && content != "" {
- if err := sender(&content, nil); err != nil {
- return err
- }
- }
-
- finishReason, ok := firstChoice["finish_reason"].(string)
- if ok && finishReason != "" {
- sawTerminal = true
- }
- return nil
+ return m.baseModel.doStreamRequest(ctx, url, apiConfig, reqBody, streamCallTimeout, func(body io.ReadCloser) error {
+ return HandleStreamingResponse(body, modelUsage, chatModelConfig, OpenAIParserConfig, sender)
})
- if err != nil {
- return fmt.Errorf("failed to scan response body: %w", err)
- }
- setSortedToolCallsResult(chatModelConfig, accumulatedToolCalls)
- if !done && !sawTerminal {
- return fmt.Errorf("mistral: stream ended before [DONE] or finish_reason")
- }
-
- endOfStream := "[DONE]"
- if err := sender(&endOfStream, nil); err != nil {
- return err
- }
-
- return nil
}
type mistralEmbeddingData struct {
diff --git a/internal/entity/models/mistral_test.go b/internal/entity/models/mistral_test.go
index d01785cac8..099879c190 100644
--- a/internal/entity/models/mistral_test.go
+++ b/internal/entity/models/mistral_test.go
@@ -914,63 +914,3 @@ func TestMistralChatIgnoresUnknownContentPartTypes(t *testing.T) {
t.Errorf("Answer=%q want %q", *resp.Answer, "Hello")
}
}
-
-// Direct unit coverage of the helper, including the nil and bad-type
-// edge cases that won't surface in the integration tests above.
-func TestExtractMistralContent(t *testing.T) {
- tests := []struct {
- name string
- input interface{}
- wantAns string
- wantReason string
- wantErr bool
- }{
- {"plain string", "hi", "hi", "", false},
- {"empty string", "", "", "", false},
- {"nil", nil, "", "", false},
- {"empty array", []interface{}{}, "", "", false},
- {
- "text only",
- []interface{}{
- map[string]interface{}{"type": "text", "text": "a"},
- map[string]interface{}{"type": "text", "text": "b"},
- },
- "ab", "", false,
- },
- {
- "thinking then text",
- []interface{}{
- map[string]interface{}{
- "type": "thinking",
- "thinking": []interface{}{
- map[string]interface{}{"type": "text", "text": "why "},
- map[string]interface{}{"type": "text", "text": "this"},
- },
- },
- map[string]interface{}{"type": "text", "text": "answer"},
- },
- "answer", "why this", false,
- },
- {"unknown root type", 42, "", "", true},
- }
- for _, tc := range tests {
- t.Run(tc.name, func(t *testing.T) {
- ans, reason, err := extractMistralContent(tc.input)
- if tc.wantErr {
- if err == nil {
- t.Errorf("want error, got nil")
- }
- return
- }
- if err != nil {
- t.Errorf("unexpected err: %v", err)
- }
- if ans != tc.wantAns {
- t.Errorf("answer=%q want %q", ans, tc.wantAns)
- }
- if reason != tc.wantReason {
- t.Errorf("reason=%q want %q", reason, tc.wantReason)
- }
- })
- }
-}
diff --git a/internal/entity/models/modelscope.go b/internal/entity/models/modelscope.go
index c07b743c64..0aceb193dd 100644
--- a/internal/entity/models/modelscope.go
+++ b/internal/entity/models/modelscope.go
@@ -17,7 +17,6 @@
package models
import (
- "bytes"
"context"
"encoding/json"
"fmt"
@@ -25,31 +24,13 @@ import (
"net/http"
"ragflow/internal/common"
"strings"
- "sync"
- "time"
)
-// modelscopeStreamIdleTimeout bounds how long a stream can go without
-var modelscopeStreamIdleTimeout = 60 * time.Second
-
// ModelScopeModel implements ModelDriver for ModelScope chat models.
type ModelScopeModel struct {
baseModel BaseModel
}
-type modelscopeChatChoice struct {
- Message struct {
- Content string `json:"content"`
- ReasoningContent string `json:"reasoning_content"`
- Reasoning string `json:"reasoning"`
- Thinking string `json:"thinking"`
- } `json:"message"`
-}
-
-type modelscopeChatResponse struct {
- Choices []modelscopeChatChoice `json:"choices"`
-}
-
type modelscopeModelListResponse struct {
Data []ModelListItem `json:"data"`
}
@@ -85,28 +66,6 @@ func normalizeModelScopeBaseURL(base string) string {
return trimmed
}
-func modelscopeReasoningFromStrings(reasoningContent string, reasoning string, thinking string) string {
- switch {
- case reasoningContent != "":
- return reasoningContent
- case reasoning != "":
- return reasoning
- case thinking != "":
- return thinking
- default:
- return ""
- }
-}
-
-func modelscopeReasoningFromMap(value map[string]interface{}) string {
- for _, field := range []string{"reasoning_content", "reasoning", "thinking"} {
- if text, ok := value[field].(string); ok && text != "" {
- return text
- }
- }
- return ""
-}
-
// ChatWithMessages sends multiple messages with roles and returns the response.
func (m *ModelScopeModel) ChatWithMessages(ctx context.Context, modelName string, messages []Message, apiConfig *APIConfig, chatModelConfig *ChatConfig, modelUsage *common.ModelUsage) (*ChatResponse, error) {
if err := m.baseModel.APIConfigCheck(apiConfig); err != nil {
@@ -125,56 +84,13 @@ func (m *ModelScopeModel) ChatWithMessages(ctx context.Context, modelName string
url := fmt.Sprintf("%s/%s", baseURL, m.baseModel.URLSuffix.Chat)
reqBody := buildRequestBody(chatModelConfig, modelName, messages, false)
- jsonData, err := json.Marshal(reqBody)
+
+ body, err := m.baseModel.doRequest(ctx, url, apiConfig, reqBody, nonStreamCallTimeout)
if err != nil {
- return nil, fmt.Errorf("failed to marshal request: %w", err)
+ return nil, err
}
- ctx, cancel := context.WithTimeout(ctx, nonStreamCallTimeout)
- defer cancel()
-
- req, err := http.NewRequestWithContext(ctx, "POST", url, bytes.NewBuffer(jsonData))
- if err != nil {
- return nil, fmt.Errorf("failed to create request: %w", err)
- }
- req.Header.Set("Content-Type", "application/json")
- if auth := BearerAuth(apiConfig); auth != "" {
- req.Header.Set("Authorization", auth)
- }
-
- resp, err := m.baseModel.httpClient.Do(req)
- if err != nil {
- return nil, fmt.Errorf("failed to send request: %w", err)
- }
- defer resp.Body.Close()
-
- body, err := io.ReadAll(resp.Body)
- if err != nil {
- return nil, fmt.Errorf("failed to read response: %w", err)
- }
- if resp.StatusCode != http.StatusOK {
- return nil, fmt.Errorf("API request failed with status %d: %s", resp.StatusCode, string(body))
- }
-
- var result modelscopeChatResponse
- if err = json.Unmarshal(body, &result); err != nil {
- return nil, fmt.Errorf("failed to parse response: %w", err)
- }
- if len(result.Choices) == 0 {
- return nil, fmt.Errorf("no choices in response")
- }
-
- content := result.Choices[0].Message.Content
- reasonContent := modelscopeReasoningFromStrings(
- result.Choices[0].Message.ReasoningContent,
- result.Choices[0].Message.Reasoning,
- result.Choices[0].Message.Thinking,
- )
-
- return &ChatResponse{
- Answer: &content,
- ReasonContent: &reasonContent,
- }, nil
+ return HandleNonStreamingResponse(body, modelUsage, chatModelConfig, OpenAIParserConfig)
}
// ChatStreamlyWithSender sends messages and streams response via sender.
@@ -189,8 +105,8 @@ func (m *ModelScopeModel) ChatStreamlyWithSender(ctx context.Context, modelName
if len(messages) == 0 {
return fmt.Errorf("messages is empty")
}
- if chatModelConfig != nil && chatModelConfig.Stream != nil && !*chatModelConfig.Stream {
- return fmt.Errorf("stream must be true in ChatStreamlyWithSender")
+ if err := validateStreamConfig(chatModelConfig); err != nil {
+ return err
}
baseURL, err := m.baseModel.GetBaseURL(apiConfig)
@@ -201,105 +117,11 @@ func (m *ModelScopeModel) ChatStreamlyWithSender(ctx context.Context, modelName
url := fmt.Sprintf("%s/%s", baseURL, m.baseModel.URLSuffix.Chat)
reqBody := buildRequestBody(chatModelConfig, modelName, messages, true)
- jsonData, err := json.Marshal(reqBody)
- if err != nil {
- return fmt.Errorf("failed to marshal request: %w", err)
- }
+ reqBody["stream_options"] = map[string]interface{}{"include_usage": true}
- ctx, cancel := context.WithCancel(context.Background())
- defer cancel()
-
- req, err := http.NewRequestWithContext(ctx, "POST", url, bytes.NewBuffer(jsonData))
- if err != nil {
- return fmt.Errorf("failed to create request: %w", err)
- }
- req.Header.Set("Content-Type", "application/json")
- if auth := BearerAuth(apiConfig); auth != "" {
- req.Header.Set("Authorization", auth)
- }
-
- resp, err := m.baseModel.httpClient.Do(req)
- if err != nil {
- return fmt.Errorf("failed to send request: %w", err)
- }
- defer resp.Body.Close()
- if resp.StatusCode != http.StatusOK {
- body, _ := io.ReadAll(resp.Body)
- return fmt.Errorf("API request failed with status %d: %s", resp.StatusCode, string(body))
- }
-
- lastActive := time.Now()
- var lastActiveMu sync.Mutex
- done := make(chan struct{})
- defer close(done)
- go func() {
- ticker := time.NewTicker(modelscopeStreamIdleTimeout / 4)
- defer ticker.Stop()
- for {
- select {
- case <-done:
- return
- case now := <-ticker.C:
- lastActiveMu.Lock()
- idle := now.Sub(lastActive)
- lastActiveMu.Unlock()
- if idle >= modelscopeStreamIdleTimeout {
- cancel()
- return
- }
- }
- }
- }()
-
- sawTerminal := false
- accumulatedToolCalls := make(map[int]map[string]any)
- streamDone, err := ParseSSEStream[map[string]interface{}](resp.Body, func(event map[string]interface{}) error {
- lastActiveMu.Lock()
- lastActive = time.Now()
- lastActiveMu.Unlock()
-
- choices, ok := event["choices"].([]interface{})
- if !ok || len(choices) == 0 {
- return nil
- }
- firstChoice, ok := choices[0].(map[string]interface{})
- if !ok {
- return nil
- }
-
- if delta, ok := firstChoice["delta"].(map[string]interface{}); ok {
- accumulateToolCallDeltas(delta, accumulatedToolCalls)
- if reasoning := modelscopeReasoningFromMap(delta); reasoning != "" {
- if err := sender(nil, &reasoning); err != nil {
- return err
- }
- }
- if content, ok := delta["content"].(string); ok && content != "" {
- if err := sender(&content, nil); err != nil {
- return err
- }
- }
- }
-
- if finishReason, ok := firstChoice["finish_reason"].(string); ok && finishReason != "" {
- sawTerminal = true
- }
- return nil
+ return m.baseModel.doStreamRequest(ctx, url, apiConfig, reqBody, streamCallTimeout, func(body io.ReadCloser) error {
+ return HandleStreamingResponse(body, modelUsage, chatModelConfig, OpenAIParserConfig, sender)
})
- if err != nil {
- if ctx.Err() != nil {
- return fmt.Errorf("modelscope: stream idle for more than %s, aborted", modelscopeStreamIdleTimeout)
- }
- return fmt.Errorf("failed to scan response body: %w", err)
- }
- if !streamDone && !sawTerminal {
- return fmt.Errorf("modelscope: stream ended before [DONE] or finish_reason")
- }
-
- setSortedToolCallsResult(chatModelConfig, accumulatedToolCalls)
-
- endOfStream := "[DONE]"
- return sender(&endOfStream, nil)
}
func (m *ModelScopeModel) Embed(ctx context.Context, modelName *string, texts []string, apiConfig *APIConfig, embeddingConfig *EmbeddingConfig, modelUsage *common.ModelUsage) ([]EmbeddingData, error) {
diff --git a/internal/entity/models/modelscope_test.go b/internal/entity/models/modelscope_test.go
index bec6da6a10..c3b047979b 100644
--- a/internal/entity/models/modelscope_test.go
+++ b/internal/entity/models/modelscope_test.go
@@ -23,7 +23,6 @@ import (
"net/http/httptest"
"strings"
"testing"
- "time"
)
func newModelScopeForTest(baseURL string) *ModelScopeModel {
@@ -36,15 +35,6 @@ func newModelScopeForTest(baseURL string) *ModelScopeModel {
)
}
-func withModelScopeIdleTimeout(t *testing.T, d time.Duration) {
- t.Helper()
- original := modelscopeStreamIdleTimeout
- modelscopeStreamIdleTimeout = d
- t.Cleanup(func() {
- modelscopeStreamIdleTimeout = original
- })
-}
-
func TestModelScopeName(t *testing.T) {
m := newModelScopeForTest("http://unused")
if got := m.Name(); got != "ModelScope" {
@@ -271,37 +261,6 @@ func TestModelScopeStreamRejectsFalseStreamConfig(t *testing.T) {
}
}
-func TestModelScopeStreamCancelsOnIdle(t *testing.T) {
- withSSRFBypass(t)
- ctx := t.Context()
- withModelScopeIdleTimeout(t, 200*time.Millisecond)
-
- hold := make(chan struct{})
- srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
- w.Header().Set("Content-Type", "text/event-stream")
- w.WriteHeader(http.StatusOK)
- if f, ok := w.(http.Flusher); ok {
- _, _ = io.WriteString(w, `data: {"choices":[{"delta":{"content":"hi"}}]}`+"\n")
- f.Flush()
- }
- select {
- case <-hold:
- case <-r.Context().Done():
- }
- }))
- t.Cleanup(srv.Close)
- t.Cleanup(func() { close(hold) })
-
- m := newModelScopeForTest(srv.URL)
- err := m.ChatStreamlyWithSender(ctx, "Qwen/Qwen2.5-7B-Instruct",
- []Message{{Role: "user", Content: "x"}},
- &APIConfig{}, nil, nil,
- func(*string, *string) error { return nil })
- if err == nil || !strings.Contains(err.Error(), "stream idle") {
- t.Errorf("expected stream-idle error, got %v", err)
- }
-}
-
func TestModelScopeListModelsAndCheckConnection(t *testing.T) {
withSSRFBypass(t)
ctx := t.Context()
diff --git a/internal/entity/models/moonshot.go b/internal/entity/models/moonshot.go
index b38586ff2d..8f921fb193 100644
--- a/internal/entity/models/moonshot.go
+++ b/internal/entity/models/moonshot.go
@@ -85,7 +85,6 @@ func (m *MoonshotModel) ChatWithMessages(ctx context.Context, modelName string,
if err := m.baseModel.APIConfigCheck(apiConfig); err != nil {
return nil, err
}
- apiKey := strings.TrimSpace(*apiConfig.ApiKey)
modelName, err := validateMoonshotModelName(modelName)
if err != nil {
return nil, err
@@ -116,84 +115,12 @@ func (m *MoonshotModel) ChatWithMessages(ctx context.Context, modelName string,
}
- jsonData, err := json.Marshal(reqBody)
+ body, err := m.baseModel.doRequest(ctx, url, apiConfig, reqBody, nonStreamCallTimeout)
if err != nil {
- return nil, fmt.Errorf("failed to marshal request: %w", err)
+ return nil, err
}
- ctx, cancel := context.WithTimeout(ctx, nonStreamCallTimeout)
- defer cancel()
-
- req, err := http.NewRequestWithContext(ctx, "POST", url, bytes.NewBuffer(jsonData))
- if err != nil {
- return nil, fmt.Errorf("failed to create request: %w", err)
- }
-
- req.Header.Set("Content-Type", "application/json")
- req.Header.Set("Authorization", fmt.Sprintf("Bearer %s", apiKey))
-
- resp, err := m.baseModel.httpClient.Do(req)
- if err != nil {
- return nil, fmt.Errorf("failed to send request: %w", err)
- }
- defer resp.Body.Close()
-
- body, err := io.ReadAll(resp.Body)
- if err != nil {
- return nil, fmt.Errorf("failed to read response: %w", err)
- }
-
- if resp.StatusCode != http.StatusOK {
- return nil, fmt.Errorf("API request failed with status %d: %s", resp.StatusCode, string(body))
- }
-
- return parseChatCompletionResponse(body, chatModelConfig, modelUsage, func(body []byte, chatConfig *ChatConfig) (chatResponseParts, error) {
- var result MoonshotChatResponse
- if err := json.Unmarshal(body, &result); err != nil {
- return chatResponseParts{}, fmt.Errorf("failed to parse response: %w", err)
- }
-
- if len(result.Choices) == 0 {
- return chatResponseParts{}, fmt.Errorf("no choices in response")
- }
-
- choice := &result.Choices[0]
- content := ""
- if choice.Message.Content != nil {
- content = *choice.Message.Content
- }
-
- reasonContent := ""
- if chatConfig != nil && chatConfig.Thinking != nil && *chatConfig.Thinking {
- if choice.Message.ReasoningContent != nil {
- reasonContent = *choice.Message.ReasoningContent
- }
- if reasonContent != "" && reasonContent[0] == '\n' {
- reasonContent = reasonContent[1:]
- }
- }
-
- totalTokens := result.Usage.TotalTokens
- if totalTokens == 0 {
- totalTokens = result.Usage.PromptTokens + result.Usage.CompletionTokens
- }
- var usage *TokenUsage
- if totalTokens > 0 {
- usage = &TokenUsage{
- PromptTokens: result.Usage.PromptTokens,
- CompletionTokens: result.Usage.CompletionTokens,
- TotalTokens: totalTokens,
- }
- }
-
- return chatResponseParts{
- RequestID: result.ID,
- Content: &content,
- ReasonContent: &reasonContent,
- ToolCalls: choice.Message.ToolCalls,
- Usage: usage,
- }, nil
- })
+ return HandleNonStreamingResponse(body, modelUsage, chatModelConfig, OpenAIParserConfig)
}
// ChatStreamlyWithSender sends messages and streams response via sender function (best performance, no channel)
@@ -262,75 +189,7 @@ func (m *MoonshotModel) ChatStreamlyWithSender(ctx context.Context, modelName st
return fmt.Errorf("API request failed with status %d: %s", resp.StatusCode, string(body))
}
- // SSE parsing: read line by line
- sawTerminal := false
- if chatModelConfig != nil {
- chatModelConfig.ToolCallsResult = nil
- }
- accumulatedToolCalls := make(map[int]map[string]any)
- done, err := ParseSSEStream[map[string]interface{}](resp.Body, func(event map[string]interface{}) error {
- tokenUsage, found, usageErr := decodeOpenAICompatibleStreamUsage(event)
- if usageErr != nil {
- return usageErr
- }
- if found {
- applyStreamUsage(chatModelConfig, modelUsage, tokenUsage)
- }
-
- choices, ok := event["choices"].([]interface{})
- if !ok || len(choices) == 0 {
- return nil
- }
-
- firstChoice, ok := choices[0].(map[string]interface{})
- if !ok {
- return nil
- }
-
- delta, ok := firstChoice["delta"].(map[string]interface{})
- if !ok {
- return nil
- }
-
- accumulateToolCallDeltas(delta, accumulatedToolCalls)
-
- reasoningContent, ok := delta["reasoning_content"].(string)
- if ok && reasoningContent != "" {
- if err := sender(nil, &reasoningContent); err != nil {
- return err
- }
- }
-
- content, ok := delta["content"].(string)
- if ok && content != "" {
- if err := sender(&content, nil); err != nil {
- return err
- }
- }
-
- finishReason, ok := firstChoice["finish_reason"].(string)
- if ok && finishReason != "" {
- sawTerminal = true
- }
-
- return nil
- })
- if err != nil {
- return fmt.Errorf("failed to scan response body: %w", err)
- }
- if !done && !sawTerminal {
- return fmt.Errorf("moonshot: stream ended before [DONE] or finish_reason")
- }
-
- setSortedToolCallsResult(chatModelConfig, accumulatedToolCalls)
-
- // Send [DONE] marker for OpenAI compatibility
- endOfStream := "[DONE]"
- if err = sender(&endOfStream, nil); err != nil {
- return err
- }
-
- return nil
+ return HandleStreamingResponse(resp.Body, modelUsage, chatModelConfig, OpenAIParserConfig, sender)
}
// Embed embeds a list of texts into embeddings
diff --git a/internal/entity/models/n1n.go b/internal/entity/models/n1n.go
index 6207ce5716..6d80446bb1 100644
--- a/internal/entity/models/n1n.go
+++ b/internal/entity/models/n1n.go
@@ -88,33 +88,12 @@ func newN1NJSONRequest(ctx context.Context, method, endpoint string, payload int
return req, nil
}
-type n1nChatChoice struct {
- Message n1nChatMessage `json:"message"`
- Delta n1nChatDelta `json:"delta"`
- FinishReason string `json:"finish_reason"`
-}
-
-type n1nChatMessage struct {
- Content *string `json:"content"`
- ReasoningContent string `json:"reasoning_content"`
-}
-
-type n1nChatDelta struct {
- Content string `json:"content"`
- ReasoningContent string `json:"reasoning_content"`
-}
-
-type n1nChatResponse struct {
- Choices []n1nChatChoice `json:"choices"`
-}
-
// ChatWithMessages sends a single, non-streaming chat completion
// against n1n.ai's /v1/chat/completions endpoint.
func (n *N1NModel) ChatWithMessages(ctx context.Context, modelName string, messages []Message, apiConfig *APIConfig, chatModelConfig *ChatConfig, modelUsage *common.ModelUsage) (*ChatResponse, error) {
if err := n.baseModel.APIConfigCheck(apiConfig); err != nil {
return nil, err
}
- apiKey := *apiConfig.ApiKey
if strings.TrimSpace(modelName) == "" {
return nil, fmt.Errorf("model name is required")
}
@@ -136,49 +115,12 @@ func (n *N1NModel) ChatWithMessages(ctx context.Context, modelName string, messa
reqBody["thinking"] = map[string]interface{}{"type": thinkingType}
}
- ctx, cancel := context.WithTimeout(ctx, nonStreamCallTimeout)
- defer cancel()
-
- req, err := newN1NJSONRequest(ctx, "POST", endpoint, reqBody, apiKey)
+ body, err := n.baseModel.doRequest(ctx, endpoint, apiConfig, reqBody, nonStreamCallTimeout)
if err != nil {
return nil, err
}
- resp, err := n.baseModel.httpClient.Do(req)
- if err != nil {
- return nil, fmt.Errorf("failed to send request: %w", err)
- }
- defer resp.Body.Close()
-
- body, err := io.ReadAll(resp.Body)
- if err != nil {
- return nil, fmt.Errorf("failed to read response: %w", err)
- }
-
- if resp.StatusCode != http.StatusOK {
- return nil, fmt.Errorf("n1n chat API error: %s, body: %s", resp.Status, string(body))
- }
-
- var parsed n1nChatResponse
- if err := json.Unmarshal(body, &parsed); err != nil {
- return nil, fmt.Errorf("failed to parse response: %w", err)
- }
- if len(parsed.Choices) == 0 {
- return nil, fmt.Errorf("no choices in response")
- }
- if parsed.Choices[0].Message.Content == nil {
- return nil, fmt.Errorf("invalid content format")
- }
-
- content := *parsed.Choices[0].Message.Content
- chatResp := &ChatResponse{
- Answer: &content,
- }
- if parsed.Choices[0].Message.ReasoningContent != "" {
- reasonContent := parsed.Choices[0].Message.ReasoningContent
- chatResp.ReasonContent = &reasonContent
- }
- return chatResp, nil
+ return HandleNonStreamingResponse(body, modelUsage, chatModelConfig, OpenAIParserConfig)
}
// ChatStreamlyWithSender sends a streaming chat completion.
@@ -196,7 +138,6 @@ func (n *N1NModel) ChatStreamlyWithSender(ctx context.Context, modelName string,
if len(messages) == 0 {
return fmt.Errorf("messages is empty")
}
- apiKey := *apiConfig.ApiKey
endpoint, err := n.endpointURL(n1nRegion(apiConfig), n.baseModel.URLSuffix.Chat)
if err != nil {
@@ -208,6 +149,7 @@ func (n *N1NModel) ChatStreamlyWithSender(ctx context.Context, modelName string,
}
reqBody := buildRequestBody(chatModelConfig, modelName, messages, true)
+ reqBody["stream_options"] = map[string]interface{}{"include_usage": true}
if chatModelConfig != nil && chatModelConfig.Thinking != nil {
thinkingType := "disabled"
if *chatModelConfig.Thinking {
@@ -216,86 +158,9 @@ func (n *N1NModel) ChatStreamlyWithSender(ctx context.Context, modelName string,
reqBody["thinking"] = map[string]interface{}{"type": thinkingType}
}
- req, err := newN1NJSONRequest(ctx, "POST", endpoint, reqBody, apiKey)
- if err != nil {
- return err
- }
-
- resp, err := n.baseModel.httpClient.Do(req)
- if err != nil {
- return fmt.Errorf("failed to send request: %w", err)
- }
- defer resp.Body.Close()
-
- if resp.StatusCode != http.StatusOK {
- body, _ := io.ReadAll(resp.Body)
- return fmt.Errorf("n1n chat stream API error: %s, body: %s", resp.Status, string(body))
- }
-
- sawTerminal := false
- accumulatedToolCalls := make(map[int]map[string]interface{})
- done, err := ParseSSEStream[map[string]interface{}](resp.Body, func(event map[string]interface{}) error {
- common.Info(fmt.Sprintf("%v", event))
-
- tokenUsage, found, usageErr := decodeOpenAICompatibleStreamUsage(event)
- if usageErr != nil {
- return usageErr
- }
- if found {
- applyStreamUsage(chatModelConfig, modelUsage, tokenUsage)
- }
-
- choices, ok := event["choices"].([]interface{})
- if !ok || len(choices) == 0 {
- return nil
- }
-
- firstChoice, ok := choices[0].(map[string]interface{})
- if !ok {
- return nil
- }
-
- delta, ok := firstChoice["delta"].(map[string]interface{})
- if !ok {
- return nil
- }
-
- accumulateToolCallDeltas(delta, accumulatedToolCalls)
-
- content, ok := delta["content"].(string)
- if ok && content != "" {
- if err := sender(&content, nil); err != nil {
- return err
- }
- }
-
- reasoningContent, ok := delta["reasoning_content"].(string)
- if ok && reasoningContent != "" {
- if err := sender(nil, &reasoningContent); err != nil {
- return err
- }
- }
-
- finishReason, ok := firstChoice["finish_reason"].(string)
- if ok && finishReason != "" {
- sawTerminal = true
- }
- return nil
+ return n.baseModel.doStreamRequest(ctx, endpoint, apiConfig, reqBody, streamCallTimeout, func(body io.ReadCloser) error {
+ return HandleStreamingResponse(body, modelUsage, chatModelConfig, OpenAIParserConfig, sender)
})
- if err != nil {
- return fmt.Errorf("failed to scan response body: %w", err)
- }
- if !done && !sawTerminal {
- return fmt.Errorf("deepseek: stream ended before [DONE] or finish_reason")
- }
-
- setSortedToolCallsResult(chatModelConfig, accumulatedToolCalls)
-
- endOfStream := "[DONE]"
- if err := sender(&endOfStream, nil); err != nil {
- return err
- }
- return nil
}
type n1nEmbeddingData struct {
diff --git a/internal/entity/models/novita.go b/internal/entity/models/novita.go
index 839b6206d2..6920f4cfd7 100644
--- a/internal/entity/models/novita.go
+++ b/internal/entity/models/novita.go
@@ -52,158 +52,6 @@ func (n *NovitaModel) Name() string {
return "NovitaAI"
}
-type NovitaChatResponse struct {
- ID string `json:"id"`
- Choices []struct {
- FinishReason string `json:"finish_reason"`
- Index int `json:"index"`
- Message struct {
- Content string `json:"content"`
- ReasoningContent string `json:"reasoning_content"`
- Role string `json:"role"`
- ToolCalls []map[string]any `json:"tool_calls"`
- } `json:"message"`
- Logprobs interface{} `json:"logprobs"`
- } `json:"choices"`
- Created int `json:"created"`
- Model string `json:"model"`
- Object string `json:"object"`
- Usage struct {
- CompletionTokens int `json:"completion_tokens"`
- PromptTokens int `json:"prompt_tokens"`
- TotalTokens int `json:"total_tokens"`
- } `json:"usage"`
-}
-
-const (
- novitaThinkOpen = ""
- novitaThinkClose = ""
-)
-
-// splitNovitaThink walks a complete content string and returns the
-// visible portion + the concatenated chain-of-thought from inside
-// any ... blocks. Multiple think blocks are
-// concatenated; tags themselves are stripped. Used by the
-// non-streaming path where the whole content is available at once.
-func splitNovitaThink(raw string) (visible, reasoning string) {
- var v, r strings.Builder
- inside := false
- for {
- var marker string
- if inside {
- marker = novitaThinkClose
- } else {
- marker = novitaThinkOpen
- }
- idx := strings.Index(raw, marker)
- if idx < 0 {
- if inside {
- r.WriteString(raw)
- } else {
- v.WriteString(raw)
- }
- break
- }
- if inside {
- r.WriteString(raw[:idx])
- } else {
- v.WriteString(raw[:idx])
- }
- raw = raw[idx+len(marker):]
- inside = !inside
- }
- return v.String(), r.String()
-}
-
-// novitaThinkExtractor maintains state across streaming chunks so
-// that a ... block spanning multiple SSE events still
-// gets split correctly between content and reasoning. The buffer
-// preserves up to (len(closingMarker)-1) trailing bytes of each
-// chunk in case the next chunk completes a partial tag.
-type novitaThinkExtractor struct {
- buf strings.Builder
- inside bool
-}
-
-// novitaThinkSegment is one routing decision: emit `content` via the
-// sender's first arg, or emit `reasoning` via the sender's second arg.
-// Exactly one of the two fields is non-empty.
-type novitaThinkSegment struct {
- content string
- reasoning string
-}
-
-// Feed appends an incoming chunk and returns any segments that are
-// now safe to emit. Trailing bytes that could be the start of a tag
-// are held back in the buffer until the next call.
-func (e *novitaThinkExtractor) Feed(chunk string) []novitaThinkSegment {
- e.buf.WriteString(chunk)
- s := e.buf.String()
- var out []novitaThinkSegment
- for {
- var marker, otherMarker string
- if e.inside {
- marker = novitaThinkClose
- otherMarker = novitaThinkOpen
- } else {
- marker = novitaThinkOpen
- otherMarker = novitaThinkClose
- }
- idx := strings.Index(s, marker)
- if idx < 0 {
- // No closing/opening marker yet. Emit everything except a
- // possible partial-tag suffix at the very end. Reserve
- // (max marker length - 1) trailing bytes so we don't
- // emit "".
- reserve := max(len(otherMarker)-1, len(marker)-1)
- safe := max(len(s)-reserve, 0)
- // Don't reserve if the trailing bytes can't possibly be
- // the start of a tag (no '<' suffix).
- if safe < len(s) && !strings.Contains(s[safe:], "<") {
- safe = len(s)
- }
- if safe > 0 {
- if e.inside {
- out = append(out, novitaThinkSegment{reasoning: s[:safe]})
- } else {
- out = append(out, novitaThinkSegment{content: s[:safe]})
- }
- s = s[safe:]
- }
- break
- }
- if idx > 0 {
- if e.inside {
- out = append(out, novitaThinkSegment{reasoning: s[:idx]})
- } else {
- out = append(out, novitaThinkSegment{content: s[:idx]})
- }
- }
- s = s[idx+len(marker):]
- e.inside = !e.inside
- }
- e.buf.Reset()
- e.buf.WriteString(s)
- return out
-}
-
-// Flush returns the buffered tail when the stream ends. A stream that
-// ends mid-tag would not normally happen with a well-behaved upstream,
-// but if it does the partial bytes are emitted according to the
-// current mode so nothing is silently lost.
-func (e *novitaThinkExtractor) Flush() *novitaThinkSegment {
- s := e.buf.String()
- e.buf.Reset()
- if s == "" {
- return nil
- }
- if e.inside {
- return &novitaThinkSegment{reasoning: s}
- }
- return &novitaThinkSegment{content: s}
-}
-
// ChatWithMessages sends multiple messages with roles and returns the response.
func (n *NovitaModel) ChatWithMessages(ctx context.Context, modelName string, messages []Message, apiConfig *APIConfig, chatModelConfig *ChatConfig, modelUsage *common.ModelUsage) (*ChatResponse, error) {
if err := n.baseModel.APIConfigCheck(apiConfig); err != nil {
@@ -227,91 +75,20 @@ func (n *NovitaModel) ChatWithMessages(ctx context.Context, modelName string, me
}
}
- jsonData, err := json.Marshal(reqBody)
+ body, err := n.baseModel.doRequest(ctx, url, apiConfig, reqBody, nonStreamCallTimeout)
if err != nil {
- return nil, fmt.Errorf("failed to marshal request: %w", err)
+ return nil, err
}
- ctx, cancel := context.WithTimeout(ctx, nonStreamCallTimeout)
- defer cancel()
-
- req, err := http.NewRequestWithContext(ctx, "POST", url, bytes.NewBuffer(jsonData))
- if err != nil {
- return nil, fmt.Errorf("failed to create request: %w", err)
- }
-
- req.Header.Set("Content-Type", "application/json")
- req.Header.Set("Authorization", fmt.Sprintf("Bearer %s", *apiConfig.ApiKey))
-
- resp, err := n.baseModel.httpClient.Do(req)
- if err != nil {
- return nil, fmt.Errorf("failed to send request: %w", err)
- }
- defer resp.Body.Close()
-
- body, err := io.ReadAll(resp.Body)
- if err != nil {
- return nil, fmt.Errorf("failed to read response: %w", err)
- }
-
- if resp.StatusCode != http.StatusOK {
- return nil, fmt.Errorf("API request failed with status %d: %s", resp.StatusCode, string(body))
- }
-
- return parseChatCompletionResponse(body, chatModelConfig, modelUsage, func(body []byte, chatConfig *ChatConfig) (chatResponseParts, error) {
- var result NovitaChatResponse
- if err := json.Unmarshal(body, &result); err != nil {
- return chatResponseParts{}, fmt.Errorf("failed to parse response: %w", err)
- }
- if len(result.Choices) == 0 {
- return chatResponseParts{}, fmt.Errorf("no choices in response")
- }
-
- choice := &result.Choices[0]
- if choice.Message.Content == "" && len(choice.Message.ToolCalls) == 0 {
- return chatResponseParts{}, fmt.Errorf("invalid content format")
- }
-
- // Novita emits chain-of-thought in two different shapes depending
- // on the model and on enable_thinking:
- // - qwen3-* and other inline-style models: chain-of-thought is
- // embedded inside content as ... tags.
- // - deepseek-v3.1 / glm-4.5 (and any model with separate
- // reasoning enabled): chain-of-thought arrives in a separate
- // `reasoning_content` field, with `content` already cleaned.
- // Handle both so the visible Answer is always tag-free and any
- // reasoning the upstream supplied is preserved.
- visible, reasoning := splitNovitaThink(choice.Message.Content)
- if choice.Message.ReasoningContent != "" {
- if reasoning != "" {
- reasoning += "\n" + choice.Message.ReasoningContent
- } else {
- reasoning = choice.Message.ReasoningContent
- }
- }
-
- return chatResponseParts{
- RequestID: result.ID,
- Content: &visible,
- ReasonContent: &reasoning,
- ToolCalls: choice.Message.ToolCalls,
- Usage: &TokenUsage{
- PromptTokens: result.Usage.PromptTokens,
- CompletionTokens: result.Usage.CompletionTokens,
- TotalTokens: result.Usage.TotalTokens,
- },
- }, nil
- })
+ return HandleNonStreamingResponse(body, modelUsage, chatModelConfig, OpenAIParserConfig)
}
// ChatStreamlyWithSender sends messages and streams the response via
// the sender. Handles both reasoning shapes Novita can emit:
// - delta.reasoning_content (deepseek-v3.1 / glm-4.5 / any model
// with separate reasoning): forwarded as-is to the second arg.
-// - delta.content containing ... (qwen3-* and other
-// inline-style models): a stateful extractor splits tag bytes
-// across SSE chunk boundaries, then routes content/reasoning to
-// the first/second sender arg respectively.
+// - delta.content (qwen3-* and other inline-style models): forwarded
+// to the first arg.
func (n *NovitaModel) ChatStreamlyWithSender(ctx context.Context, modelName string, messages []Message, apiConfig *APIConfig, chatModelConfig *ChatConfig, modelUsage *common.ModelUsage, sender func(*string, *string) error) error {
if err := n.baseModel.APIConfigCheck(apiConfig); err != nil {
return err
@@ -341,117 +118,186 @@ func (n *NovitaModel) ChatStreamlyWithSender(ctx context.Context, modelName stri
}
}
- jsonData, err := json.Marshal(reqBody)
- if err != nil {
- return fmt.Errorf("failed to marshal request: %w", err)
- }
+ reqBody["stream_options"] = map[string]interface{}{"include_usage": true}
- req, err := http.NewRequestWithContext(ctx, "POST", url, bytes.NewBuffer(jsonData))
- if err != nil {
- return fmt.Errorf("failed to create request: %w", err)
- }
- req.Header.Set("Content-Type", "application/json")
- req.Header.Set("Authorization", fmt.Sprintf("Bearer %s", *apiConfig.ApiKey))
+ return n.baseModel.doStreamRequest(ctx, url, apiConfig, reqBody, streamCallTimeout, func(body io.ReadCloser) error {
+ // Novita qwen3 embeds ... inline in
+ // delta.content (tags can span multiple SSE deltas). Split
+ // those blocks so reasoning routes to the sender's second arg.
+ return novitaHandleStream(body, modelUsage, chatModelConfig, sender)
+ })
+}
- resp, err := n.baseModel.httpClient.Do(req)
- if err != nil {
- return fmt.Errorf("failed to send request: %w", err)
- }
- defer resp.Body.Close()
+// novitaThinkSegment is one routing decision: emit `content` via the
+// sender's first arg, or emit `reasoning` via the second. Exactly one of
+// the two fields is non-empty.
+type novitaThinkSegment struct {
+ content string
+ reasoning string
+}
- if resp.StatusCode != http.StatusOK {
- body, _ := io.ReadAll(resp.Body)
- return fmt.Errorf("API request failed with status %d: %s", resp.StatusCode, string(body))
- }
+// novitaThinkSplitter holds state across streaming content chunks so a
+// ... block that spans multiple SSE deltas is still split
+// correctly. Trailing bytes that could be the start of a tag are held
+// back until the next chunk.
+type novitaThinkSplitter struct {
+ buf strings.Builder
+ inside bool
+}
- extractor := &novitaThinkExtractor{}
- sawTerminal := false
- done, err := ParseSSEStream[map[string]interface{}](resp.Body, func(event map[string]interface{}) error {
- tokenUsage, found, usageErr := decodeOpenAICompatibleStreamUsage(event)
- if usageErr != nil {
- return usageErr
+const (
+ novitaThinkOpen = ""
+ novitaThinkClose = ""
+)
+
+func (s *novitaThinkSplitter) feed(chunk string) []novitaThinkSegment {
+ s.buf.WriteString(chunk)
+ str := s.buf.String()
+ var out []novitaThinkSegment
+ for {
+ var marker string
+ if s.inside {
+ marker = novitaThinkClose
+ } else {
+ marker = novitaThinkOpen
}
- if found {
- applyStreamUsage(chatModelConfig, modelUsage, tokenUsage)
+ idx := strings.Index(str, marker)
+ if idx < 0 {
+ // No marker yet. Emit everything except a possible
+ // partial-tag suffix at the very end.
+ reserve := max(len(novitaThinkOpen)-1, len(novitaThinkClose)-1)
+ safe := max(len(str)-reserve, 0)
+ if safe < len(str) && !strings.Contains(str[safe:], "<") {
+ safe = len(str)
+ }
+ if safe > 0 {
+ if s.inside {
+ out = append(out, novitaThinkSegment{reasoning: str[:safe]})
+ } else {
+ out = append(out, novitaThinkSegment{content: str[:safe]})
+ }
+ str = str[safe:]
+ }
+ s.buf.Reset()
+ s.buf.WriteString(str)
+ return out
+ }
+ if s.inside {
+ out = append(out, novitaThinkSegment{reasoning: str[:idx]})
+ } else {
+ out = append(out, novitaThinkSegment{content: str[:idx]})
+ }
+ str = str[idx+len(marker):]
+ s.inside = !s.inside
+ }
+}
+
+func (s *novitaThinkSplitter) flush() *novitaThinkSegment {
+ if s.buf.Len() == 0 {
+ return nil
+ }
+ remaining := s.buf.String()
+ s.buf.Reset()
+ if s.inside {
+ return &novitaThinkSegment{reasoning: remaining}
+ }
+ return &novitaThinkSegment{content: remaining}
+}
+
+// novitaHandleStream processes a Novita streaming chat response, splitting
+// inline ... blocks in delta.content across SSE deltas.
+func novitaHandleStream(
+ body io.Reader,
+ modelUsage *common.ModelUsage,
+ chatConfig *ChatConfig,
+ sender func(*string, *string) error,
+) error {
+ if sender == nil {
+ return fmt.Errorf("sender is required")
+ }
+
+ var sawTerminal bool
+ thinkSplitter := &novitaThinkSplitter{}
+
+ done, err := ParseSSEStream[map[string]any](body, func(event map[string]any) error {
+ tokenUsage, found := extractOpenAIStreamUsage(event)
+ if found && chatConfig != nil {
+ applyStreamUsage(chatConfig, modelUsage, tokenUsage)
}
- choices, ok := event["choices"].([]interface{})
+ if apiErr, ok := event["error"]; ok && apiErr != nil {
+ return fmt.Errorf("upstream stream error: %v", apiErr)
+ }
+
+ choices, ok := event["choices"].([]any)
if !ok || len(choices) == 0 {
return nil
}
- firstChoice, ok := choices[0].(map[string]interface{})
+
+ firstChoice, ok := choices[0].(map[string]any)
if !ok {
return nil
}
- delta, ok := firstChoice["delta"].(map[string]interface{})
+
+ delta, ok := firstChoice["delta"].(map[string]any)
if !ok {
return nil
}
- // deepseek-v3.1 / glm-4.5 (and other models that emit reasoning
- // separately) put chain-of-thought in delta.reasoning_content
- // rather than inside content as .... Surface it
- // before any content from the same chunk so callers piping to
- // a UI render reasoning before the visible answer for that
- // token, matching the wire ordering Novita emits.
- if r, ok := delta["reasoning_content"].(string); ok && r != "" {
- rr := r
- if err := sender(nil, &rr); err != nil {
+
+ if reasoningContent, ok := delta["reasoning_content"].(string); ok && reasoningContent != "" {
+ if err := sender(nil, &reasoningContent); err != nil {
return err
}
}
- if c, ok := delta["content"].(string); ok && c != "" {
- for _, seg := range extractor.Feed(c) {
+
+ if content, ok := delta["content"].(string); ok && content != "" {
+ for _, seg := range thinkSplitter.feed(content) {
+ if seg.content != "" {
+ c := seg.content
+ if err := sender(&c, nil); err != nil {
+ return err
+ }
+ }
if seg.reasoning != "" {
r := seg.reasoning
if err := sender(nil, &r); err != nil {
return err
}
}
- if seg.content != "" {
- cc := seg.content
- if err := sender(&cc, nil); err != nil {
- return err
- }
- }
}
}
- if finish, ok := firstChoice["finish_reason"].(string); ok && finish != "" {
+
+ if finishReason, ok := firstChoice["finish_reason"].(string); ok && finishReason != "" {
sawTerminal = true
}
+
return nil
})
if err != nil {
return fmt.Errorf("failed to scan response body: %w", err)
}
- // Flush any buffered tail (rare, but covers the case where the
- // stream ends right after the last chunk without us seeing the
- // closing tag).
- if seg := extractor.Flush(); seg != nil {
+ if !done && !sawTerminal {
+ return fmt.Errorf("stream ended before [DONE] or finish_reason")
+ }
+
+ if seg := thinkSplitter.flush(); seg != nil {
+ if seg.content != "" {
+ c := seg.content
+ if err := sender(&c, nil); err != nil {
+ return err
+ }
+ }
if seg.reasoning != "" {
r := seg.reasoning
if err := sender(nil, &r); err != nil {
return err
}
}
- if seg.content != "" {
- cc := seg.content
- if err := sender(&cc, nil); err != nil {
- return err
- }
- }
- }
-
- if !done && !sawTerminal {
- return fmt.Errorf("novita: stream ended before [DONE] or finish_reason")
}
endOfStream := "[DONE]"
- if err := sender(&endOfStream, nil); err != nil {
- return err
- }
-
- return nil
+ return sender(&endOfStream, nil)
}
// ListModels returns the list of model ids visible to the API key.
diff --git a/internal/entity/models/novita_test.go b/internal/entity/models/novita_test.go
index 9bcf17c736..49038895d7 100644
--- a/internal/entity/models/novita_test.go
+++ b/internal/entity/models/novita_test.go
@@ -101,151 +101,6 @@ func newNovitaSSEServer(t *testing.T, expectedPath, ssePayload string) *httptest
}))
}
-// ---- think-tag split helpers ----
-
-func TestSplitNovitaThinkPureText(t *testing.T) {
- v, r := splitNovitaThink("hello world")
- if v != "hello world" || r != "" {
- t.Errorf("got (%q,%q)", v, r)
- }
-}
-
-func TestSplitNovitaThinkSingleBlock(t *testing.T) {
- v, r := splitNovitaThink("15% = 0.15. 0.15*80 = 12.The answer is 12.")
- if v != "The answer is 12." {
- t.Errorf("visible=%q", v)
- }
- if r != "15% = 0.15. 0.15*80 = 12." {
- t.Errorf("reasoning=%q", r)
- }
-}
-
-func TestSplitNovitaThinkLeadingText(t *testing.T) {
- v, r := splitNovitaThink("intro thoughttail")
- if v != "intro tail" {
- t.Errorf("visible=%q", v)
- }
- if r != "thought" {
- t.Errorf("reasoning=%q", r)
- }
-}
-
-func TestSplitNovitaThinkMultipleBlocks(t *testing.T) {
- v, r := splitNovitaThink("Apart1Bpart2")
- if v != "part1part2" {
- t.Errorf("visible=%q", v)
- }
- if r != "AB" {
- t.Errorf("reasoning=%q", r)
- }
-}
-
-func TestSplitNovitaThinkUnclosedTag(t *testing.T) {
- // Unclosed -> everything after the open tag is reasoning;
- // content stops at the open tag. This matches a real upstream that
- // got cut off mid-reasoning by max_tokens.
- v, r := splitNovitaThink("visible still thinking when tokens ran out")
- if v != "visible " {
- t.Errorf("visible=%q", v)
- }
- if r != "still thinking when tokens ran out" {
- t.Errorf("reasoning=%q", r)
- }
-}
-
-// ---- streaming extractor ----
-
-// Helper to push multiple chunks through and concatenate all output by
-// kind. Each chunk goes into Feed; the output is what's safe to emit.
-func feedAll(e *novitaThinkExtractor, chunks []string) (content, reasoning string) {
- var cb, rb strings.Builder
- for _, c := range chunks {
- for _, seg := range e.Feed(c) {
- cb.WriteString(seg.content)
- rb.WriteString(seg.reasoning)
- }
- }
- if seg := e.Flush(); seg != nil {
- cb.WriteString(seg.content)
- rb.WriteString(seg.reasoning)
- }
- return cb.String(), rb.String()
-}
-
-func TestNovitaThinkExtractorSingleChunk(t *testing.T) {
- e := &novitaThinkExtractor{}
- c, r := feedAll(e, []string{"hello thought world"})
- if c != "hello world" {
- t.Errorf("content=%q", c)
- }
- if r != "thought" {
- t.Errorf("reasoning=%q", r)
- }
-}
-
-func TestNovitaThinkExtractorTagSpansChunks(t *testing.T) {
- // "" split across two SSE deltas: ""
- e := &novitaThinkExtractor{}
- c, r := feedAll(e, []string{"hello thoughttail"})
- if c != "hello tail" {
- t.Errorf("content=%q", c)
- }
- if r != "thought" {
- t.Errorf("reasoning=%q", r)
- }
-}
-
-func TestNovitaThinkExtractorClosingTagSpansChunks(t *testing.T) {
- // "" split across two deltas
- e := &novitaThinkExtractor{}
- c, r := feedAll(e, []string{"reasoningvisible"})
- if c != "visible" {
- t.Errorf("content=%q", c)
- }
- if r != "reasoning" {
- t.Errorf("reasoning=%q", r)
- }
-}
-
-func TestNovitaThinkExtractorTokenBoundaries(t *testing.T) {
- // Simulate the kind of chunking we saw on the wire for qwen3 — many
- // small chunks, sometimes splitting tag bytes.
- e := &novitaThinkExtractor{}
- c, r := feedAll(e, []string{
- "<", "think>", "Ok", "ay, ", "compute. ", "think>", "12", "."})
- if c != "12." {
- t.Errorf("content=%q", c)
- }
- if r != "Okay, compute. " {
- t.Errorf("reasoning=%q", r)
- }
-}
-
-func TestNovitaThinkExtractorNoTags(t *testing.T) {
- e := &novitaThinkExtractor{}
- c, r := feedAll(e, []string{"plain ", "content ", "all ", "the way"})
- if c != "plain content all the way" {
- t.Errorf("content=%q", c)
- }
- if r != "" {
- t.Errorf("reasoning=%q", r)
- }
-}
-
-func TestNovitaThinkExtractorLessThanIsNotTagStart(t *testing.T) {
- // "<10" or " 0 {
- chatResponse.Usage = &TokenUsage{
- PromptTokens: pt, CompletionTokens: ct, TotalTokens: tt,
- }
- }
-
- return chatResponse, nil
+ return HandleNonStreamingResponse(body, modelUsage, chatModelConfig, OpenAIParserConfig)
}
// ChatStreamlyWithSender sends messages and streams the response
@@ -242,102 +155,7 @@ func (o *OpenAIModel) ChatStreamlyWithSender(ctx context.Context, modelName stri
return fmt.Errorf("API request failed with status %d: %s", resp.StatusCode, string(body))
}
- sawTerminal := false
- accumulatedToolCalls := make(map[int]map[string]interface{})
- // Capture the authoritative usage block from the final streaming
- // chunk (when provider honours stream_options.include_usage=true).
- // The last chunk in the stream carries the "usage" key alongside
- // empty choices; we overwrite on every chunk so the final frame
- // wins, matching Python's chat_model.py usage_from_response loop.
- var streamUsage *TokenUsage
- scanner := bufio.NewScanner(resp.Body)
- for scanner.Scan() {
- line := scanner.Text()
-
- // SSE data line starts with "data:"
- if !strings.HasPrefix(line, "data:") {
- continue
- }
-
- // Extract JSON after "data:"
- data := strings.TrimSpace(line[5:])
-
- // [DONE] marks the end of the stream
- if data == "[DONE]" {
- sawTerminal = true
- break
- }
-
- // Parse the JSON event
- var event map[string]interface{}
- if err = json.Unmarshal([]byte(data), &event); err != nil {
- continue
- }
-
- // Extract usage from this chunk. When stream_options.include_usage
- // is true, the final chunk carries the full usage breakdown at the
- // top level of the event alongside (possibly empty) choices.
- if pt, ct, tt := extractUsageFromMap(event); tt > 0 {
- streamUsage = &TokenUsage{PromptTokens: pt, CompletionTokens: ct, TotalTokens: tt}
- }
-
- choices, ok := event["choices"].([]interface{})
- if !ok || len(choices) == 0 {
- continue
- }
-
- firstChoice, ok := choices[0].(map[string]interface{})
- if !ok {
- continue
- }
-
- delta, ok := firstChoice["delta"].(map[string]interface{})
- if !ok {
- continue
- }
-
- accumulateToolCallDeltas(delta, accumulatedToolCalls)
-
- reasoningContent, ok := delta["reasoning_content"].(string)
- if ok && reasoningContent != "" {
- if err := sender(nil, &reasoningContent); err != nil {
- return err
- }
- }
-
- content, ok := delta["content"].(string)
- if ok && content != "" {
- if err := sender(&content, nil); err != nil {
- return err
- }
- }
-
- finishReason, ok := firstChoice["finish_reason"].(string)
- if ok && finishReason != "" {
- sawTerminal = true
- }
- }
- if err := scanner.Err(); err != nil {
- return fmt.Errorf("failed to scan response body: %w", err)
- }
- if !sawTerminal {
- return fmt.Errorf("openai: stream ended before [DONE] or finish_reason")
- }
-
- setSortedToolCallsResult(chatModelConfig, accumulatedToolCalls)
-
- // Populate UsageResult with the authoritative usage from the stream.
- if streamUsage != nil && chatModelConfig != nil {
- chatModelConfig.UsageResult = streamUsage
- }
-
- // Send the [DONE] marker for OpenAI compatibility
- endOfStream := "[DONE]"
- if err := sender(&endOfStream, nil); err != nil {
- return err
- }
-
- return nil
+ return HandleStreamingResponse(resp.Body, modelUsage, chatModelConfig, OpenAIParserConfig, sender)
}
type openaiEmbeddingResponse struct {
diff --git a/internal/entity/models/openai_api_compatible.go b/internal/entity/models/openai_api_compatible.go
index 5a0ee13c9a..dc54518260 100644
--- a/internal/entity/models/openai_api_compatible.go
+++ b/internal/entity/models/openai_api_compatible.go
@@ -79,6 +79,87 @@ func (m *OpenAIAPICompatibleModel) ListModels(ctx context.Context, apiConfig *AP
return filtered, nil
}
+// ChatWithMessages sends multiple messages with roles and returns response
+func (m *OpenAIAPICompatibleModel) ChatWithMessages(ctx context.Context, modelName string, messages []Message, apiConfig *APIConfig, chatModelConfig *ChatConfig, modelUsage *common.ModelUsage) (*ChatResponse, error) {
+ if err := m.baseModel.APIConfigCheck(apiConfig); err != nil {
+ return nil, err
+ }
+
+ if len(messages) == 0 {
+ return nil, fmt.Errorf("messages is empty")
+ }
+
+ resolvedBaseURL, err := m.baseModel.GetBaseURL(apiConfig)
+ if err != nil {
+ return nil, err
+ }
+ url := fmt.Sprintf("%s/%s", resolvedBaseURL, m.baseModel.URLSuffix.Chat)
+
+ // Build request body
+ reqBody := buildRequestBody(chatModelConfig, modelName, messages, false)
+
+ if chatModelConfig != nil && chatModelConfig.Thinking != nil {
+ if *chatModelConfig.Thinking {
+ reqBody["thinking"] = map[string]interface{}{
+ "type": "enabled",
+ }
+ } else {
+ reqBody["thinking"] = map[string]interface{}{
+ "type": "disabled",
+ }
+ }
+ }
+
+ body, err := m.baseModel.doRequest(ctx, url, apiConfig, reqBody, nonStreamCallTimeout)
+ if err != nil {
+ return nil, err
+ }
+
+ return HandleNonStreamingResponse(body, modelUsage, chatModelConfig, OpenAIParserConfig)
+}
+
+// ChatStreamlyWithSender sends messages and streams response via sender function
+func (m *OpenAIAPICompatibleModel) ChatStreamlyWithSender(ctx context.Context, modelName string, messages []Message, apiConfig *APIConfig, chatModelConfig *ChatConfig, modelUsage *common.ModelUsage, sender func(*string, *string) error) error {
+ if err := m.baseModel.APIConfigCheck(apiConfig); err != nil {
+ return err
+ }
+
+ if len(messages) == 0 {
+ return fmt.Errorf("messages is empty")
+ }
+ if sender == nil {
+ return fmt.Errorf("sender is required")
+ }
+
+ resolvedBaseURL, err := m.baseModel.GetBaseURL(apiConfig)
+ if err != nil {
+ return err
+ }
+ url := fmt.Sprintf("%s/%s", resolvedBaseURL, m.baseModel.URLSuffix.Chat)
+
+ // Build request body with streaming enabled
+ reqBody := buildRequestBody(chatModelConfig, modelName, messages, true)
+ reqBody["stream_options"] = map[string]interface{}{
+ "include_usage": true,
+ }
+
+ if chatModelConfig != nil && chatModelConfig.Thinking != nil {
+ if *chatModelConfig.Thinking {
+ reqBody["thinking"] = map[string]interface{}{
+ "type": "enabled",
+ }
+ } else {
+ reqBody["thinking"] = map[string]interface{}{
+ "type": "disabled",
+ }
+ }
+ }
+
+ return m.baseModel.doStreamRequest(ctx, url, apiConfig, reqBody, streamCallTimeout, func(body io.ReadCloser) error {
+ return HandleStreamingResponse(body, modelUsage, chatModelConfig, OpenAIParserConfig, sender)
+ })
+}
+
// Hint keywords for model type inference, matching Python's
// OpenAIAPICompatible class-level hint constants.
var (
diff --git a/internal/entity/models/openrouter.go b/internal/entity/models/openrouter.go
index e573467dfa..ef5c0ae4d5 100644
--- a/internal/entity/models/openrouter.go
+++ b/internal/entity/models/openrouter.go
@@ -54,37 +54,6 @@ func (o *OpenRouterModel) Name() string {
return "openrouter"
}
-// OpenRouterChatResponse mirrors OpenRouter's chat-completions response.
-type OpenRouterChatResponse struct {
- ID string `json:"id"`
- Object string `json:"object"`
- Created int64 `json:"created"`
- Model string `json:"model"`
- Choices []struct {
- FinishReason string `json:"finish_reason"`
- Index int `json:"index"`
- Logprobs any `json:"logprobs"`
- Message struct {
- Content string `json:"content"`
- Reasoning string `json:"reasoning"`
- Role string `json:"role"`
- ToolCalls []map[string]any `json:"tool_calls"`
- } `json:"message"`
- } `json:"choices"`
- SystemFingerprint string `json:"system_fingerprint"`
- Usage struct {
- CompletionTokens int `json:"completion_tokens"`
- PromptTokens int `json:"prompt_tokens"`
- TotalTokens int `json:"total_tokens"`
- CompletionTokensDetails struct {
- ReasoningTokens int `json:"reasoning_tokens"`
- } `json:"completion_tokens_details"`
- PromptTokensDetails struct {
- CachedTokens int `json:"cached_tokens"`
- } `json:"prompt_tokens_details"`
- } `json:"usage"`
-}
-
func (o *OpenRouterModel) ChatWithMessages(ctx context.Context, modelName string, messages []Message, apiConfig *APIConfig, chatModelConfig *ChatConfig, modelUsage *common.ModelUsage) (*ChatResponse, error) {
if err := o.baseModel.APIConfigCheck(apiConfig); err != nil {
return nil, err
@@ -108,67 +77,12 @@ func (o *OpenRouterModel) ChatWithMessages(ctx context.Context, modelName string
}
}
- jsonData, err := json.Marshal(reqBody)
+ body, err := o.baseModel.doRequest(ctx, url, apiConfig, reqBody, nonStreamCallTimeout)
if err != nil {
- return nil, fmt.Errorf("failed to marshal request: %w", err)
+ return nil, err
}
- ctx, cancel := context.WithTimeout(ctx, nonStreamCallTimeout)
- defer cancel()
-
- req, err := http.NewRequestWithContext(ctx, "POST", url, bytes.NewBuffer(jsonData))
- if err != nil {
- return nil, fmt.Errorf("failed to create request: %w", err)
- }
-
- req.Header.Add("Content-Type", "application/json")
- req.Header.Add("Authorization", fmt.Sprintf("Bearer %s", *apiConfig.ApiKey))
-
- resp, err := o.baseModel.httpClient.Do(req)
- if err != nil {
- return nil, fmt.Errorf("failed to send request: %w", err)
- }
- defer resp.Body.Close()
-
- body, err := io.ReadAll(resp.Body)
- if err != nil {
- return nil, fmt.Errorf("failed to read response body: %w", err)
- }
-
- if resp.StatusCode != http.StatusOK {
- return nil, fmt.Errorf("failed to send request: %d %s", resp.StatusCode, string(body))
- }
-
- return parseChatCompletionResponse(body, chatModelConfig, modelUsage, func(body []byte, chatConfig *ChatConfig) (chatResponseParts, error) {
- var result OpenRouterChatResponse
- if err := json.Unmarshal(body, &result); err != nil {
- return chatResponseParts{}, fmt.Errorf("failed to unmarshal response: %w", err)
- }
- if len(result.Choices) == 0 {
- return chatResponseParts{}, fmt.Errorf("no choices in response")
- }
-
- choice := result.Choices[0]
- if choice.Message.Content == "" && len(choice.Message.ToolCalls) == 0 {
- return chatResponseParts{}, fmt.Errorf("response contains neither content nor tool calls")
- }
- reasonContent := ""
- if chatConfig != nil && chatConfig.Thinking != nil && *chatConfig.Thinking {
- reasonContent = strings.TrimPrefix(choice.Message.Reasoning, "\n")
- }
-
- return chatResponseParts{
- RequestID: result.ID,
- Content: &choice.Message.Content,
- ReasonContent: &reasonContent,
- ToolCalls: choice.Message.ToolCalls,
- Usage: &TokenUsage{
- PromptTokens: result.Usage.PromptTokens,
- CompletionTokens: result.Usage.CompletionTokens,
- TotalTokens: result.Usage.TotalTokens,
- },
- }, nil
- })
+ return HandleNonStreamingResponse(body, modelUsage, chatModelConfig, OpenAIParserConfig)
}
func (o *OpenRouterModel) ChatStreamlyWithSender(ctx context.Context, modelName string, messages []Message, apiConfig *APIConfig, modelConfig *ChatConfig, modelUsage *common.ModelUsage, sender func(*string, *string) error) error {
@@ -206,95 +120,11 @@ func (o *OpenRouterModel) ChatStreamlyWithSender(ctx context.Context, modelName
}
}
- jsonData, err := json.Marshal(reqBody)
- if err != nil {
- return fmt.Errorf("failed to marshal request: %w", err)
- }
-
- ctx, cancel := context.WithTimeout(ctx, streamCallTimeout)
- defer cancel()
-
- req, err := http.NewRequestWithContext(ctx, "POST", url, bytes.NewBuffer(jsonData))
- if err != nil {
- return fmt.Errorf("failed to create request: %w", err)
- }
-
- req.Header.Set("Content-Type", "application/json")
- req.Header.Set("Authorization", fmt.Sprintf("Bearer %s", *apiConfig.ApiKey))
-
- resp, err := o.baseModel.httpClient.Do(req)
- if err != nil {
- return fmt.Errorf("failed to send request: %w", err)
- }
- defer resp.Body.Close()
-
- if resp.StatusCode != http.StatusOK {
- body, _ := io.ReadAll(resp.Body)
- return fmt.Errorf("invalid status code: %d, body: %s", resp.StatusCode, string(body))
- }
-
- sawTerminal := false
- accumulatedToolCalls := make(map[int]map[string]any)
- done, err := ParseSSEStream[map[string]any](resp.Body, func(event map[string]any) error {
- common.Info(fmt.Sprintf("%v", event))
-
- tokenUsage, found, usageErr := decodeOpenAICompatibleStreamUsage(event)
- if usageErr != nil {
- return usageErr
- }
- if found {
- applyStreamUsage(modelConfig, modelUsage, tokenUsage)
- }
- choices, ok := event["choices"].([]any)
- if !ok || len(choices) == 0 {
- return nil
- }
-
- choice, ok := choices[0].(map[string]any)
- if !ok {
- return nil
- }
- if finishReason, ok := choice["finish_reason"].(string); ok && finishReason != "" {
- sawTerminal = true
- }
- delta, ok := choice["delta"].(map[string]any)
- if !ok {
- return nil
- }
-
- accumulateToolCallDeltas(delta, accumulatedToolCalls)
-
- reasoningContent, ok := delta["reasoning"].(string)
- if ok && reasoningContent != "" {
- if err := sender(nil, &reasoningContent); err != nil {
- return err
- }
- }
-
- content, ok := delta["content"].(string)
- if ok && content != "" {
- if err := sender(&content, nil); err != nil {
- return err
- }
- }
-
- return nil
+ return o.baseModel.doStreamRequest(ctx, url, apiConfig, reqBody, streamCallTimeout, func(body io.ReadCloser) error {
+ // OpenRouter emits reasoning under delta.reasoning (not
+ // delta.reasoning_content), so it uses its own ParserConfig.
+ return HandleStreamingResponse(body, modelUsage, modelConfig, OpenRouterParserConfig, sender)
})
- if err != nil {
- return fmt.Errorf("failed to scan response body: %w", err)
- }
- if !done && !sawTerminal {
- return fmt.Errorf("openrouter: stream ended before [DONE] or finish_reason")
- }
- setSortedToolCallsResult(modelConfig, accumulatedToolCalls)
-
- // Send [DONE] marker for OpenAI compatibility
- endOfStream := "[DONE]"
- if err = sender(&endOfStream, nil); err != nil {
- return err
- }
-
- return nil
}
// OpenRouterEmbeddingResponse mirrors OpenRouter's embeddings response.
diff --git a/internal/entity/models/parser_config.go b/internal/entity/models/parser_config.go
index 28f7530212..6512c3e914 100644
--- a/internal/entity/models/parser_config.go
+++ b/internal/entity/models/parser_config.go
@@ -16,8 +16,8 @@
package models
-// ParserConfig maps a protocol to its usage parsers. Drivers select a
-// ParserConfig instead of implementing usage extraction individually.
+// ParserConfig maps a protocol to its usage and reasoning parsers. Drivers
+// select a ParserConfig instead of implementing extraction individually.
type ParserConfig struct {
// Protocol is the protocol identifier (e.g. "openai", "claude").
Protocol string
@@ -25,12 +25,45 @@ type ParserConfig struct {
ResponseParser func(map[string]any) (*TokenUsage, bool)
// StreamParser extracts usage from one streaming event.
StreamParser func(map[string]any) (*TokenUsage, bool)
+ // ExtractStreamReasoning extracts the reasoning text from a parsed
+ // delta map (the delta field of one streaming event), if any.
+ // Defaults to reading delta.reasoning_content.
+ ExtractStreamReasoning func(delta map[string]any) string
+}
+
+// extractDefaultStreamReasoning reads the reasoning text from a parsed
+// delta (delta.reasoning_content).
+func extractDefaultStreamReasoning(delta map[string]any) string {
+ if r, ok := delta["reasoning_content"].(string); ok {
+ return r
+ }
+ return ""
}
// OpenAIParserConfig is the ParserConfig for OpenAI-compatible APIs
-// (NVIDIA NIM, DeepSeek, Aliyun, Moonshot, xAI, OpenRouter, ...).
+// (NVIDIA NIM, DeepSeek, Aliyun, Moonshot, xAI, ...).
var OpenAIParserConfig = &ParserConfig{
- Protocol: "openai",
- ResponseParser: extractOpenAIUsage,
- StreamParser: extractOpenAIStreamUsage,
+ Protocol: "openai",
+ ResponseParser: extractOpenAIUsage,
+ StreamParser: extractOpenAIStreamUsage,
+ ExtractStreamReasoning: extractDefaultStreamReasoning,
+}
+
+// extractOpenRouterStreamReasoning reads the reasoning text from an
+// OpenRouter streaming event. OpenRouter uses delta.reasoning (not
+// delta.reasoning_content) for its reasoning content.
+func extractOpenRouterStreamReasoning(delta map[string]any) string {
+ if r, ok := delta["reasoning"].(string); ok {
+ return r
+ }
+ return ""
+}
+
+// OpenRouterParserConfig is the ParserConfig for OpenRouter, which emits
+// reasoning under delta.reasoning instead of delta.reasoning_content.
+var OpenRouterParserConfig = &ParserConfig{
+ Protocol: "openai",
+ ResponseParser: extractOpenAIUsage,
+ StreamParser: extractOpenAIStreamUsage,
+ ExtractStreamReasoning: extractOpenRouterStreamReasoning,
}
diff --git a/internal/entity/models/response_handler.go b/internal/entity/models/response_handler.go
index d52d01a70a..e7d8cefadc 100644
--- a/internal/entity/models/response_handler.go
+++ b/internal/entity/models/response_handler.go
@@ -115,8 +115,15 @@ func HandleStreamingResponse(
accumulateToolCallDeltas(delta, accumulatedToolCalls)
- if reasoningContent, ok := delta["reasoning_content"].(string); ok && reasoningContent != "" {
- if err := sender(nil, &reasoningContent); err != nil {
+ // Extract reasoning via the protocol hook so each provider can
+ // name its reasoning field differently (reasoning_content,
+ // reasoning, ...) without the shared handler knowing which.
+ extractReasoning := cfg.ExtractStreamReasoning
+ if extractReasoning == nil {
+ extractReasoning = extractDefaultStreamReasoning
+ }
+ if reasoning := extractReasoning(delta); reasoning != "" {
+ if err := sender(nil, &reasoning); err != nil {
return err
}
}