Files
aichat/stream/ollama.go
T
2026-06-17 11:29:51 +08:00

220 lines
5.8 KiB
Go

package stream
import (
"bufio"
"bytes"
"context"
"encoding/json"
"errors"
"fmt"
"io"
"net/http"
"strings"
"time"
"aichat/llm"
"github.com/volcengine/volcengine-go-sdk/service/arkruntime/model"
)
type ollamaChatRequest struct {
Model string `json:"model"`
Messages []ollamaChatMessage `json:"messages"`
Stream bool `json:"stream"`
Options map[string]int `json:"options,omitempty"`
}
type ollamaChatMessage struct {
Role string `json:"role"`
Content string `json:"content"`
Images []string `json:"images,omitempty"`
}
type ollamaChatResponse struct {
Message struct {
Role string `json:"role"`
Content string `json:"content"`
Thinking string `json:"thinking"`
} `json:"message"`
Done bool `json:"done"`
PromptEvalCount int `json:"prompt_eval_count"`
EvalCount int `json:"eval_count"`
DoneReason string `json:"done_reason"`
}
func StreamOllamaChat(ctx context.Context, profile *llm.Profile, messages []*model.ChatCompletionMessage, promptTokens int, usage *Tracker, emit EmitFunc, onDone func(string)) error {
requestMessages, err := buildOllamaMessages(messages)
if err != nil {
return err
}
baseURL, err := llm.OllamaBaseURL(profile)
if err != nil {
return err
}
body, err := json.Marshal(ollamaChatRequest{
Model: profile.Config.Model,
Messages: requestMessages,
Stream: true,
Options: map[string]int{"num_predict": 4096},
})
if err != nil {
return err
}
req, err := http.NewRequestWithContext(ctx, http.MethodPost, strings.TrimRight(baseURL, "/")+"/api/chat", bytes.NewReader(body))
if err != nil {
return err
}
req.Header.Set("Content-Type", "application/json")
resp, err := http.DefaultClient.Do(req)
if err != nil {
return err
}
defer resp.Body.Close()
if resp.StatusCode < 200 || resp.StatusCode >= 300 {
data, _ := io.ReadAll(io.LimitReader(resp.Body, 4096))
return fmt.Errorf("Ollama 原生接口调用失败: %s %s", resp.Status, strings.TrimSpace(string(data)))
}
emit(Frame{Type: "trace", Tool: "model", Stage: "stream", Status: "running", Message: "Ollama 视觉模型已开始输出"})
parseThinkTags := llm.ShouldParseThinkTags(profile)
thinkParser := &Parser{}
var full strings.Builder
completionTokens := 0
streamStarted := time.Now()
peakTokensPerSecond := 0.0
emitDelta := func(delta string) {
if delta == "" {
return
}
full.WriteString(delta)
completionTokens += EstimateTokenCount(delta)
usage.SetModel(promptTokens, completionTokens)
currentSpeed := TokensPerSecond(completionTokens, streamStarted)
if currentSpeed > peakTokensPerSecond {
peakTokensPerSecond = currentSpeed
}
stats := usage.Snapshot(currentSpeed, peakTokensPerSecond)
emit(Frame{Type: "delta", Text: delta, Stats: &stats})
}
emitContent := func(delta string) {
if delta == "" {
return
}
if !parseThinkTags {
emitDelta(delta)
return
}
visible, reasoning := thinkParser.Accept(delta)
if reasoning != "" {
emit(Frame{Type: "reasoning", Text: reasoning})
}
emitDelta(visible)
}
scanner := bufio.NewScanner(resp.Body)
scanner.Buffer(make([]byte, 0, 64*1024), 10*1024*1024)
for scanner.Scan() {
line := strings.TrimSpace(scanner.Text())
if line == "" {
continue
}
var chunk ollamaChatResponse
if err := json.Unmarshal([]byte(line), &chunk); err != nil {
return fmt.Errorf("解析 Ollama 流失败: %w", err)
}
if chunk.Message.Thinking != "" {
emit(Frame{Type: "reasoning", Text: chunk.Message.Thinking})
}
emitContent(chunk.Message.Content)
if chunk.Done {
if chunk.PromptEvalCount > 0 || chunk.EvalCount > 0 {
usage.SetModel(chunk.PromptEvalCount, chunk.EvalCount)
}
break
}
}
if err := scanner.Err(); err != nil {
return err
}
if parseThinkTags {
visible, reasoning := thinkParser.Flush()
if reasoning != "" {
emit(Frame{Type: "reasoning", Text: reasoning})
}
emitDelta(visible)
}
if onDone != nil {
onDone(full.String())
}
finalStats := usage.Snapshot(TokensPerSecond(completionTokens, streamStarted), peakTokensPerSecond)
emit(Frame{Type: "stats", Stats: &finalStats})
emit(Frame{Type: "trace", Tool: "model", Stage: "stream", Status: "success", Message: "回答生成完成"})
return nil
}
func buildOllamaMessages(messages []*model.ChatCompletionMessage) ([]ollamaChatMessage, error) {
result := make([]ollamaChatMessage, 0, len(messages))
for _, msg := range messages {
if msg == nil {
continue
}
role := string(msg.Role)
if msg.Role == model.ChatMessageRoleTool {
role = string(model.ChatMessageRoleUser)
}
item := ollamaChatMessage{Role: role}
if msg.Content == nil {
if len(msg.ToolCalls) > 0 {
continue
}
result = append(result, item)
continue
}
if msg.Content.StringValue != nil {
item.Content = *msg.Content.StringValue
if msg.Role == model.ChatMessageRoleTool {
item.Content = "工具结果:\n" + item.Content
}
result = append(result, item)
continue
}
for _, part := range msg.Content.ListValue {
if part == nil {
continue
}
switch part.Type {
case model.ChatCompletionMessageContentPartTypeText:
if part.Text != "" {
if item.Content != "" {
item.Content += "\n"
}
item.Content += part.Text
}
case model.ChatCompletionMessageContentPartTypeImageURL:
if part.ImageURL == nil {
continue
}
image, err := ollamaImagePayload(part.ImageURL.URL)
if err != nil {
return nil, err
}
item.Images = append(item.Images, image)
}
}
result = append(result, item)
}
return result, nil
}
func ollamaImagePayload(raw string) (string, error) {
raw = strings.TrimSpace(raw)
if strings.HasPrefix(strings.ToLower(raw), "data:") {
comma := strings.Index(raw, ",")
if comma < 0 {
return "", errors.New("图片 base64 数据格式错误")
}
return strings.TrimSpace(raw[comma+1:]), nil
}
return raw, nil
}