mirror of
https://gitcode.com/JianFeeeee/HomeAgent.git
synced 2026-09-21 17:38:10 +00:00
provider: 音频多模态序列化为 OpenAI「input_audio」格式
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 交叉编译通过;部署后服务健康
This commit is contained in:
@ -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:<mime>;base64,<data> 提取 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"`
|
||||
|
||||
@ -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 }
|
||||
@ -64,3 +66,40 @@ func names(ps []Provider) []string {
|
||||
}
|
||||
return out
|
||||
}
|
||||
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)
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user