From c19fea25842ffc92ac1b2b35399562ea3bd9e04b Mon Sep 17 00:00:00 2001 From: root Date: Mon, 10 Aug 2026 12:55:43 +0800 Subject: [PATCH] =?UTF-8?q?provider:=20=E9=9F=B3=E9=A2=91=E5=A4=9A?= =?UTF-8?q?=E6=A8=A1=E6=80=81=E5=BA=8F=E5=88=97=E5=8C=96=E4=B8=BA=20OpenAI?= =?UTF-8?q?=E3=80=8Cinput=5Faudio=E3=80=8D=E6=A0=BC=E5=BC=8F?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit HomeAgent 的 ContentBlock.AudioURL 原本输出 audio_url(非 OpenAI 标准块); 现通过 MarshalJSON 在 base64 数据时自动转为 OpenAI input_audio ({data,format}),供支持音频的模型识别。非 base64(url)保留原样透传。 - parseAudioDataURL: data:mime;base64,data → (data,format) - audioFormatFromMIME: wav/mp3/mp4/ogg/flac - 新增单测:base64 转 input_audio、url 保留 audio_url 验证: go test ./... 27 包 0 失败;Windows 交叉编译通过;部署后服务健康 --- internal/agent/api/provider.go | 62 +++++++++++++++++++++++++++++++ internal/agent/api/router_test.go | 47 +++++++++++++++++++++-- 2 files changed, 105 insertions(+), 4 deletions(-) diff --git a/internal/agent/api/provider.go b/internal/agent/api/provider.go index 40d0ed2..4f49cae 100644 --- a/internal/agent/api/provider.go +++ b/internal/agent/api/provider.go @@ -3,6 +3,7 @@ package api import ( "bufio" "context" + "encoding/base64" "encoding/json" "fmt" "io" @@ -32,6 +33,67 @@ type AudioURL struct { URL string `json:"url"` } +// MarshalJSON 把音频块序列化成 OpenAI「input_audio」多模态格式(base64 内嵌), +// 供支持音频的模型识别。audio_url 非 OpenAI 标准块;含 base64 数据时转 +// input_audio,否则回落到原生 audio_url(透传)。 +func (b ContentBlock) MarshalJSON() ([]byte, error) { + if b.Type == "audio_url" && b.AudioURL != nil && b.AudioURL.URL != "" { + if data, format, ok := parseAudioDataURL(b.AudioURL.URL); ok { + return json.Marshal(map[string]interface{}{ + "type": "input_audio", + "input_audio": map[string]string{ + "data": data, + "format": format, + }, + }) + } + } + type alias ContentBlock + return json.Marshal(alias(b)) +} + +// parseAudioDataURL 从 data:;base64, 提取 base64 与 format。 +// 非 base64(如 http url)返回 ok=false。 +func parseAudioDataURL(url string) (data, format string, ok bool) { + const prefix = "data:" + if !strings.HasPrefix(url, prefix) { + return "", "", false + } + rest := url[len(prefix):] + comma := strings.IndexByte(rest, ',') + if comma < 0 { + return "", "", false + } + mime := rest[:comma] + data = rest[comma+1:] + if mime == "" || data == "" { + return "", "", false + } + if _, err := base64.StdEncoding.DecodeString(data); err != nil { + return "", "", false + } + format = audioFormatFromMIME(mime) + return data, format, true +} + +func audioFormatFromMIME(mime string) string { + m := strings.ToLower(strings.TrimSpace(mime)) + switch { + case strings.Contains(m, "wav"): + return "wav" + case strings.Contains(m, "mp3"), strings.Contains(m, "mpeg"): + return "mp3" + case strings.Contains(m, "mp4"), strings.Contains(m, "m4a"): + return "mp4" + case strings.Contains(m, "ogg"), strings.Contains(m, "opus"): + return "ogg" + case strings.Contains(m, "flac"): + return "flac" + default: + return "wav" + } +} + // Message 表示对话消息。当 Blocks 不为空时 content 在 JSON 中序列化为数组(多模态格式)。 type Message struct { Role string `json:"role"` diff --git a/internal/agent/api/router_test.go b/internal/agent/api/router_test.go index 9f1b144..daaf07d 100644 --- a/internal/agent/api/router_test.go +++ b/internal/agent/api/router_test.go @@ -2,14 +2,16 @@ package api import ( "context" + "encoding/json" + "strings" "testing" ) // stubRoutableProvider implements both Provider and RoutableProvider. type stubRoutableProvider struct { - name string - model string - prio int + name string + model string + prio int } func (s *stubRoutableProvider) Name() string { return s.name } @@ -63,4 +65,41 @@ func names(ps []Provider) []string { out[i] = p.Name() } return out -} \ No newline at end of file +} +func TestMessageAudioBlockToInputAudio(t *testing.T) { + const b64 = "QUJDREVG" // base64 of "ABCDEF" + m := Message{ + Role: "user", + Blocks: []ContentBlock{ + {Type: "text", Text: "what is this"}, + {Type: "audio_url", AudioURL: &AudioURL{URL: "data:audio/wav;base64," + b64}}, + }, + } + data, err := json.Marshal(m) + if err != nil { + t.Fatalf("marshal: %v", err) + } + s := string(data) + if !strings.Contains(s, `"type":"input_audio"`) { + t.Fatalf("expected input_audio, got: %s", s) + } + if !strings.Contains(s, `"format":"wav"`) || !strings.Contains(s, b64) { + t.Fatalf("missing data/format: %s", s) + } + if strings.Contains(s, `"type":"audio_url"`) { + t.Fatalf("audio_url should be converted: %s", s) + } +} + +func TestMessageAudioURLPassthroughWhenNotBase64(t *testing.T) { + m := Message{ + Role: "user", + Blocks: []ContentBlock{ + {Type: "audio_url", AudioURL: &AudioURL{URL: "https://cdn.example/a.wav"}}, + }, + } + data, _ := json.Marshal(m) + if !strings.Contains(string(data), `"type":"audio_url"`) { + t.Fatalf("non-base64 audio_url should stay: %s", data) + } +}