162 lines
5.7 KiB
Go
162 lines
5.7 KiB
Go
package realtime
|
|
|
|
import (
|
|
"ai-operator/internal/ai"
|
|
"ai-operator/internal/config"
|
|
"ai-operator/internal/media"
|
|
"context"
|
|
"encoding/json"
|
|
"github.com/gorilla/websocket"
|
|
"net/http"
|
|
"net/http/httptest"
|
|
"strings"
|
|
"testing"
|
|
"time"
|
|
)
|
|
|
|
func TestURLBuildersParsers(t *testing.T) {
|
|
cfg := config.OpenAIConfig{RealtimeURL: "wss://api.openai.com/v1/realtime", RealtimeModel: "gpt realtime", RealtimeVoice: "marin", RealtimeInputSampleRate: 24000, RealtimeOutputSampleRate: 24000, RealtimeTurnDetection: "server_vad", RealtimeReasoningEffort: "low"}
|
|
u, _, err := BuildURL(cfg)
|
|
if err != nil || !strings.Contains(u, "model=gpt+realtime") {
|
|
t.Fatalf("url=%s err=%v", u, err)
|
|
}
|
|
b, err := BuildSessionUpdate(cfg, "test")
|
|
if err != nil || strings.Contains(string(b), "sk-") {
|
|
t.Fatal(string(b))
|
|
}
|
|
if !strings.Contains(string(b), `"voice":"marin"`) {
|
|
t.Fatalf("session.update missing configured voice: %s", b)
|
|
}
|
|
for _, fn := range []func() ([]byte, error){func() ([]byte, error) { return BuildAudioAppend([]byte{1, 2}) }, BuildCommit, BuildResponseCreate, BuildResponseCancel, BuildInputAudioClear, func() ([]byte, error) {
|
|
return BuildToolResult(ai.ToolResult{ToolCallID: "tc", Result: map[string]string{"ok": "true"}})
|
|
}} {
|
|
x, err := fn()
|
|
if err != nil || !json.Valid(x) {
|
|
t.Fatalf("bad event %s %v", x, err)
|
|
}
|
|
}
|
|
ev, _, err := ParseServerEvent([]byte(`{"type":"response.output_audio.delta","delta":"AQI="}`))
|
|
if err != nil || ev.Type != ai.VoiceEventAssistantAudioDelta || len(ev.Audio) != 2 {
|
|
t.Fatalf("ev=%+v err=%v", ev, err)
|
|
}
|
|
ev, _, _ = ParseServerEvent([]byte(`{"type":"input_audio_buffer.speech_started"}`))
|
|
if ev.Type != ai.VoiceEventInterruption {
|
|
t.Fatal(ev.Type)
|
|
}
|
|
ev, _, _ = ParseServerEvent([]byte(`{"type":"response.done","response":{"output":[{"type":"function_call","call_id":"c","name":"set_language","arguments":"{\"language\":\"ru\"}"}]}}`))
|
|
if ev.ToolCall == nil || ev.ToolCall.Name != "set_language" {
|
|
t.Fatalf("tool=%+v", ev.ToolCall)
|
|
}
|
|
if _, _, err := ParseServerEvent([]byte(`bad`)); err == nil {
|
|
t.Fatal("want err")
|
|
}
|
|
}
|
|
func TestProviderFakeWebSocket(t *testing.T) {
|
|
up := websocket.Upgrader{}
|
|
var auth string
|
|
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
|
auth = r.Header.Get("Authorization")
|
|
c, err := up.Upgrade(w, r, nil)
|
|
if err != nil {
|
|
return
|
|
}
|
|
defer c.Close()
|
|
_ = c.WriteJSON(map[string]any{"type": "session.created"})
|
|
_, msg, _ := c.ReadMessage()
|
|
if !strings.Contains(string(msg), "session.update") {
|
|
t.Errorf("msg=%s", msg)
|
|
}
|
|
_ = c.WriteJSON(map[string]any{"type": "session.updated"})
|
|
_, _, _ = c.ReadMessage()
|
|
_ = c.WriteJSON(map[string]any{"type": "response.output_audio.delta", "delta": "AQI="})
|
|
time.Sleep(20 * time.Millisecond)
|
|
}))
|
|
defer srv.Close()
|
|
cfg := config.Config{OpenAI: config.OpenAIConfig{APIKey: "sk-test", RealtimeURL: "ws" + strings.TrimPrefix(srv.URL, "http"), RealtimeModel: "m", RealtimeConnectTimeout: time.Second, RealtimeMaxAudioChunkBytes: 100, RealtimeMaxInputAudioBytes: 1000, RealtimeInputSampleRate: 24000, RealtimeOutputSampleRate: 24000, RealtimeTurnDetection: "server_vad", RealtimeReasoningEffort: "low"}}
|
|
p := NewProvider(cfg, nil)
|
|
ctx, cancel := context.WithTimeout(context.Background(), time.Second)
|
|
defer cancel()
|
|
if err := p.StartSession(ctx, ai.VoiceSessionConfig{CallID: "c1"}); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
if auth != "Bearer sk-test" {
|
|
t.Fatal("auth missing")
|
|
}
|
|
if err := p.SendAudio(ctx, media.AudioChunk{Data: []byte{0, 0}}); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
select {
|
|
case ev := <-p.Events():
|
|
_ = ev
|
|
case <-time.After(time.Second):
|
|
t.Fatal("no event")
|
|
}
|
|
_ = p.Close(ctx)
|
|
}
|
|
func TestProviderCanStartAfterClose(t *testing.T) {
|
|
up := websocket.Upgrader{}
|
|
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
|
c, err := up.Upgrade(w, r, nil)
|
|
if err != nil {
|
|
return
|
|
}
|
|
defer c.Close()
|
|
_ = c.WriteJSON(map[string]any{"type": "session.created"})
|
|
_, msg, _ := c.ReadMessage()
|
|
if !strings.Contains(string(msg), "session.update") {
|
|
t.Errorf("msg=%s", msg)
|
|
}
|
|
_ = c.WriteJSON(map[string]any{"type": "session.updated"})
|
|
_ = c.WriteJSON(map[string]any{"type": "response.output_audio.delta", "delta": "AQI="})
|
|
for {
|
|
if _, _, err := c.ReadMessage(); err != nil {
|
|
return
|
|
}
|
|
}
|
|
}))
|
|
defer srv.Close()
|
|
cfg := config.Config{OpenAI: config.OpenAIConfig{APIKey: "sk-test", RealtimeURL: "ws" + strings.TrimPrefix(srv.URL, "http"), RealtimeModel: "m", RealtimeConnectTimeout: time.Second, RealtimeMaxAudioChunkBytes: 100, RealtimeMaxInputAudioBytes: 1000, RealtimeInputSampleRate: 24000, RealtimeOutputSampleRate: 24000, RealtimeTurnDetection: "server_vad", RealtimeReasoningEffort: "low"}}
|
|
p := NewProvider(cfg, nil)
|
|
for _, callID := range []string{"c1", "c2"} {
|
|
ctx, cancel := context.WithTimeout(context.Background(), time.Second)
|
|
if err := p.StartSession(ctx, ai.VoiceSessionConfig{CallID: callID}); err != nil {
|
|
cancel()
|
|
t.Fatal(err)
|
|
}
|
|
select {
|
|
case ev := <-p.Events():
|
|
if ev.CallID != callID {
|
|
cancel()
|
|
t.Fatalf("call_id=%s want %s", ev.CallID, callID)
|
|
}
|
|
case <-ctx.Done():
|
|
cancel()
|
|
t.Fatal("no event")
|
|
}
|
|
if err := p.Close(ctx); err != nil {
|
|
cancel()
|
|
t.Fatal(err)
|
|
}
|
|
for {
|
|
select {
|
|
case _, ok := <-p.Events():
|
|
if !ok {
|
|
cancel()
|
|
goto next
|
|
}
|
|
case <-ctx.Done():
|
|
cancel()
|
|
t.Fatal("events channel did not close")
|
|
}
|
|
}
|
|
next:
|
|
}
|
|
}
|
|
func TestCostGuard(t *testing.T) {
|
|
p := NewProvider(config.Config{OpenAI: config.OpenAIConfig{RealtimeMaxInputAudioBytes: 1}}, nil)
|
|
p.conn = &websocket.Conn{}
|
|
if err := p.SendAudio(context.Background(), media.AudioChunk{Data: []byte{0, 0}}); err == nil {
|
|
t.Fatal("want cost err")
|
|
}
|
|
}
|