Files

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")
}
}