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 } }