Files
HomeAgent/internal/agent/api/router_test.go
root c19fea2584 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 交叉编译通过;部署后服务健康
2026-08-10 12:55:43 +08:00

106 lines
3.0 KiB
Go

package api
import (
"context"
"encoding/json"
"strings"
"testing"
)
// stubRoutableProvider implements both Provider and RoutableProvider.
type stubRoutableProvider struct {
name string
model string
prio int
}
func (s *stubRoutableProvider) Name() string { return s.name }
func (s *stubRoutableProvider) Model() string { return s.model }
func (s *stubRoutableProvider) Priority() int { return s.prio }
func (s *stubRoutableProvider) MaxContextTokens() int { return 4096 }
func (s *stubRoutableProvider) Chat(context.Context, *CompletionRequest) (*CompletionResponse, error) {
return &CompletionResponse{Content: s.name}, nil
}
func (s *stubRoutableProvider) ChatStream(context.Context, *CompletionRequest) (<-chan StreamChunk, error) {
ch := make(chan StreamChunk, 1)
ch <- StreamChunk{Done: true}
return ch, nil
}
func TestProviderManagerPriorityOrder(t *testing.T) {
m := NewProviderManager()
m.Register("low", &stubRoutableProvider{name: "low", prio: 10})
m.Register("high", &stubRoutableProvider{name: "high", prio: 90})
m.Register("mid", &stubRoutableProvider{name: "mid", prio: 50})
got := m.OrderedProviders()
wantOrder := []string{"high", "mid", "low"}
for i, p := range got {
if p.Name() != wantOrder[i] {
t.Fatalf("order[%d] = %s, want %s (full=%v)", i, p.Name(), wantOrder[i], names(got))
}
}
}
func TestProviderManagerResolveForModel(t *testing.T) {
m := NewProviderManager()
m.Register("a", &stubRoutableProvider{name: "a", model: "gpt-5"})
m.Register("b", &stubRoutableProvider{name: "b", model: "deepseek"})
got := m.ResolveForModel("deepseek")
if len(got) == 0 || got[0].Name() != "b" {
t.Fatalf("resolve deepseek: got %v", names(got))
}
// 未知模型回落 AUTO 链(仍按优先级)
got2 := m.ResolveForModel("unknown")
if len(got2) == 0 {
t.Fatal("unknown model should fall back to auto chain")
}
}
func names(ps []Provider) []string {
out := make([]string, len(ps))
for i, p := range ps {
out[i] = p.Name()
}
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)
}
}