mirror of
https://gitcode.com/JianFeeeee/HomeAgent.git
synced 2026-09-23 02:18:06 +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 (
|
import (
|
||||||
"bufio"
|
"bufio"
|
||||||
"context"
|
"context"
|
||||||
|
"encoding/base64"
|
||||||
"encoding/json"
|
"encoding/json"
|
||||||
"fmt"
|
"fmt"
|
||||||
"io"
|
"io"
|
||||||
@ -32,6 +33,67 @@ type AudioURL struct {
|
|||||||
URL string `json:"url"`
|
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 中序列化为数组(多模态格式)。
|
// Message 表示对话消息。当 Blocks 不为空时 content 在 JSON 中序列化为数组(多模态格式)。
|
||||||
type Message struct {
|
type Message struct {
|
||||||
Role string `json:"role"`
|
Role string `json:"role"`
|
||||||
|
|||||||
@ -2,14 +2,16 @@ package api
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
|
"encoding/json"
|
||||||
|
"strings"
|
||||||
"testing"
|
"testing"
|
||||||
)
|
)
|
||||||
|
|
||||||
// stubRoutableProvider implements both Provider and RoutableProvider.
|
// stubRoutableProvider implements both Provider and RoutableProvider.
|
||||||
type stubRoutableProvider struct {
|
type stubRoutableProvider struct {
|
||||||
name string
|
name string
|
||||||
model string
|
model string
|
||||||
prio int
|
prio int
|
||||||
}
|
}
|
||||||
|
|
||||||
func (s *stubRoutableProvider) Name() string { return s.name }
|
func (s *stubRoutableProvider) Name() string { return s.name }
|
||||||
@ -63,4 +65,41 @@ func names(ps []Provider) []string {
|
|||||||
out[i] = p.Name()
|
out[i] = p.Name()
|
||||||
}
|
}
|
||||||
return out
|
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