diff --git a/internal/entity/models/ollama.go b/internal/entity/models/ollama.go index c01584221b..86f3fd5e08 100644 --- a/internal/entity/models/ollama.go +++ b/internal/entity/models/ollama.go @@ -17,6 +17,7 @@ package models import ( + "bufio" "bytes" "context" "encoding/json" @@ -175,7 +176,7 @@ func (o *OllamaModel) ChatWithMessages(ctx context.Context, modelName string, me return nil, err } - return HandleNonStreamingResponse(body, modelUsage, chatModelConfig, OpenAIParserConfig) + return handleOllamaNonStreamingResponse(body, modelUsage, chatModelConfig) } func (o *OllamaModel) ChatStreamlyWithSender(ctx context.Context, modelName string, messages []Message, apiConfig *APIConfig, modelConfig *ChatConfig, modelUsage *common.ModelUsage, sender func(*string, *string) error) error { @@ -196,23 +197,129 @@ func (o *OllamaModel) ChatStreamlyWithSender(ctx context.Context, modelName stri // Build request body with streaming enabled reqBody := buildOllamaRequestBody(modelConfig, modelName, messages, true) - if modelConfig.Effort != nil && *modelConfig.Effort != "" { - if strings.HasPrefix(strings.ToLower(modelName), "gpt-oss") { - reqBody["think"] = *modelConfig.Effort - } - } else if modelConfig.Thinking != nil { - if *modelConfig.Thinking { - reqBody["think"] = true + if modelConfig != nil { + if modelConfig.Effort != nil && *modelConfig.Effort != "" { + if strings.HasPrefix(strings.ToLower(modelName), "gpt-oss") { + reqBody["think"] = *modelConfig.Effort + } + } else if modelConfig.Thinking != nil { + if *modelConfig.Thinking { + reqBody["think"] = true + } } } reqBody["stream_options"] = map[string]interface{}{"include_usage": true} return o.baseModel.doStreamRequest(ctx, url, apiConfig, reqBody, streamCallTimeout, func(body io.ReadCloser) error { - return HandleStreamingResponse(body, modelUsage, modelConfig, OpenAIParserConfig, sender) + return handleOllamaStreamingResponse(body, modelUsage, modelConfig, sender) }) } +type ollamaChatMessage struct { + Content string `json:"content"` + Thinking string `json:"thinking"` +} + +type ollamaChatResponse struct { + Model string `json:"model"` + Message ollamaChatMessage `json:"message"` + Done bool `json:"done"` + PromptEvalCount int `json:"prompt_eval_count"` + EvalCount int `json:"eval_count"` + TotalDuration int64 `json:"total_duration"` + Error any `json:"error"` +} + +func handleOllamaNonStreamingResponse(body []byte, modelUsage *common.ModelUsage, chatConfig *ChatConfig) (*ChatResponse, error) { + var result ollamaChatResponse + if err := json.Unmarshal(body, &result); err != nil { + return nil, fmt.Errorf("failed to parse response: %w", err) + } + if result.Error != nil { + return nil, fmt.Errorf("upstream error: %v", result.Error) + } + content := result.Message.Content + reasonContent := result.Message.Thinking + usage := ollamaTokenUsage(result) + if usage != nil { + recordResponseUsage(modelUsage, "", usage, "chat") + if chatConfig != nil { + chatConfig.UsageResult = usage + } + } + if content == "" && reasonContent == "" { + return nil, fmt.Errorf("no message in response") + } + return &ChatResponse{ + Answer: &content, + ReasonContent: &reasonContent, + Usage: usage, + }, nil +} + +func handleOllamaStreamingResponse(body io.Reader, modelUsage *common.ModelUsage, chatConfig *ChatConfig, sender func(*string, *string) error) error { + if sender == nil { + return fmt.Errorf("sender is required") + } + var usage *TokenUsage + sawDone := false + scanner := bufio.NewScanner(body) + scanner.Buffer(make([]byte, 64*1024), 1024*1024) + for scanner.Scan() { + line := strings.TrimSpace(scanner.Text()) + if line == "" { + continue + } + var event ollamaChatResponse + if err := json.Unmarshal([]byte(line), &event); err != nil { + return fmt.Errorf("invalid Ollama stream event: %w", err) + } + if event.Error != nil { + return fmt.Errorf("upstream stream error: %v", event.Error) + } + if event.Message.Thinking != "" { + if err := sender(nil, &event.Message.Thinking); err != nil { + return err + } + } + if event.Message.Content != "" { + if err := sender(&event.Message.Content, nil); err != nil { + return err + } + } + if event.Done { + sawDone = true + usage = ollamaTokenUsage(event) + } + } + if err := scanner.Err(); err != nil { + return fmt.Errorf("failed to scan response body: %w", err) + } + if !sawDone { + return fmt.Errorf("stream ended before done") + } + if usage != nil { + recordResponseUsage(modelUsage, "", usage, "chat") + if chatConfig != nil { + chatConfig.UsageResult = usage + } + } + endOfStream := "[DONE]" + return sender(&endOfStream, nil) +} + +func ollamaTokenUsage(result ollamaChatResponse) *TokenUsage { + if result.PromptEvalCount == 0 && result.EvalCount == 0 { + return nil + } + return &TokenUsage{ + PromptTokens: result.PromptEvalCount, + CompletionTokens: result.EvalCount, + TotalTokens: result.PromptEvalCount + result.EvalCount, + } +} + func (o *OllamaModel) Embed(ctx context.Context, modelName *string, request EmbedRequest, apiConfig *APIConfig, embeddingConfig *EmbeddingConfig, modelUsage *common.ModelUsage) ([]EmbeddingData, error) { if err := o.baseModel.APIConfigCheck(apiConfig); err != nil { return nil, err