Files
model-gateway/common/util/streaming.go
T

256 lines
6.0 KiB
Go

package util
import (
"context"
"encoding/base64"
"encoding/json"
"fmt"
"model-gateway/service/gateway"
"sort"
"strings"
"github.com/gogf/gf/v2/encoding/gjson"
)
// ParseStreamResponse 流式响应解析
func ParseStreamResponse(ctx context.Context, rawBytes []byte, streamConfig map[string]any) (map[string]any, error) {
enabled, _ := streamConfig["enabled"].(bool)
if !enabled {
return gjson.New(string(rawBytes)).Map(), nil
}
outputType, _ := streamConfig["output_type"].(string)
streamToClient, _ := streamConfig["stream_to_client"].(bool)
events, _ := streamConfig["events"].([]any)
if len(events) == 0 {
return gjson.New(string(rawBytes)).Map(), nil
}
// 如果业务需要流式返回,直接透传原始数据
if streamToClient {
return map[string]any{"stream_data": rawBytes}, nil
}
result := make(map[string]any)
lines := strings.Split(string(rawBytes), "\n")
for _, evt := range events {
e, _ := evt.(map[string]any)
evtType, _ := e["type"].(string)
switch evtType {
case "concat":
processConcat(lines, e, result)
case "base64_concat":
processBase64Concat(lines, e, result)
case "collect":
processCollect(lines, e, result)
case "final":
processFinal(lines, e, result)
}
}
switch outputType {
case "audio":
if audioBytes, ok := result["audio"].([]byte); ok {
oss, err := gateway.UploadByTask(ctx, audioBytes, "mp3")
if err != nil {
return nil, err
}
result["content"] = oss.FileAddressPrefix + oss.FileURL
}
case "text":
if v, ok := result["content"]; ok {
return map[string]any{"content": v, "usage": result["usage"]}, nil
}
case "image":
if v, ok := result["urls"]; ok {
return map[string]any{"content": v, "usage": result["usage"]}, nil
}
}
return result, nil
}
// processConcat 文本拼接
func processConcat(lines []string, event map[string]any, result map[string]any) {
match, _ := event["match"].(string)
aggregateTo, _ := event["aggregate_to"].(string)
fields, _ := event["fields"].(map[string]any)
var parts []string
for _, line := range lines {
line = strings.TrimSpace(line)
if line == "" || line == "[DONE]" {
continue
}
if strings.HasPrefix(line, "event:") {
continue
}
if strings.HasPrefix(line, "data:") {
line = strings.TrimPrefix(line, "data:")
line = strings.TrimSpace(line)
}
var chunk map[string]any
if err := json.Unmarshal([]byte(line), &chunk); err != nil {
continue
}
chunkType, _ := chunk["type"].(string)
if match != "" && !strings.Contains(chunkType, match) {
continue
}
for _, chunkPath := range fields {
val := gjson.New(chunk).Get(chunkPath.(string)).String()
if val != "" {
parts = append(parts, val)
}
}
}
result[aggregateTo] = strings.Join(parts, "")
}
// processBase64Concat base64 拼接
func processBase64Concat(lines []string, event map[string]any, result map[string]any) {
aggregateTo, _ := event["aggregate_to"].(string)
fields, _ := event["fields"].(map[string]any)
var builder strings.Builder
for _, line := range lines {
line = strings.TrimSpace(line)
if line == "" || line == "[DONE]" {
continue
}
var chunk map[string]any
if err := json.Unmarshal([]byte(line), &chunk); err != nil {
continue
}
for _, chunkPath := range fields {
if data := gjson.New(chunk).Get(chunkPath.(string)).String(); data != "" {
builder.WriteString(data)
}
}
}
cleanBase64 := strings.Map(func(r rune) rune {
if r == ' ' || r == '\n' || r == '\r' || r == '\t' {
return -1
}
return r
}, builder.String())
audioBytes, err := base64.StdEncoding.DecodeString(cleanBase64)
if err != nil {
audioBytes, _ = base64.RawStdEncoding.DecodeString(cleanBase64)
}
result[aggregateTo] = audioBytes
}
// processCollect 数组收集
func processCollect(lines []string, event map[string]any, result map[string]any) {
match, _ := event["match"].(string)
aggregateTo, _ := event["aggregate_to"].(string)
orderBy, _ := event["order_by"].(string)
fields, _ := event["fields"].(map[string]any)
var items []map[string]any
for _, line := range lines {
line = strings.TrimSpace(line)
if line == "" || line == "[DONE]" {
continue
}
var chunk map[string]any
if err := json.Unmarshal([]byte(line), &chunk); err != nil {
continue
}
chunkType, _ := chunk["type"].(string)
if match != "" && !strings.Contains(chunkType, match) {
continue
}
item := make(map[string]any)
for localKey, chunkPath := range fields {
item[localKey] = gjson.New(chunk).Get(chunkPath.(string)).Val()
}
items = append(items, item)
}
if orderBy != "" {
sort.Slice(items, func(i, j int) bool {
return fmt.Sprint(items[i][orderBy]) < fmt.Sprint(items[j][orderBy])
})
}
// 如果只有一个字段,直接存值数组
if len(fields) == 1 {
var vals []any
for _, item := range items {
for _, v := range item {
vals = append(vals, v)
}
}
result[aggregateTo] = vals
} else {
result[aggregateTo] = items
}
}
// processFinal 取最后一条匹配的数据
func processFinal(lines []string, event map[string]any, result map[string]any) {
match, _ := event["match"].(string)
aggregateTo, _ := event["aggregate_to"].(string)
fields, _ := event["fields"].(map[string]any)
var lastMatch map[string]any
for _, line := range lines {
line = strings.TrimSpace(line)
if line == "" || line == "[DONE]" {
continue
}
if strings.HasPrefix(line, "event:") {
continue
}
if strings.HasPrefix(line, "data:") {
line = strings.TrimPrefix(line, "data:")
line = strings.TrimSpace(line)
}
var chunk map[string]any
if err := json.Unmarshal([]byte(line), &chunk); err != nil {
continue
}
chunkType, _ := chunk["type"].(string)
if match != "" && !strings.Contains(chunkType, match) {
continue
}
lastMatch = chunk
}
if lastMatch == nil {
return
}
data := make(map[string]any)
for localKey, chunkPath := range fields {
val := gjson.New(lastMatch).Get(chunkPath.(string)).Val()
if val != nil {
data[localKey] = val
}
}
if len(data) > 0 {
result[aggregateTo] = data
}
}