256 lines
6.0 KiB
Go
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
|
|
}
|
|
}
|