220 lines
5.8 KiB
Go
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
|
|
}
|