203 lines
7 KiB
Go
203 lines
7 KiB
Go
package live
|
|
|
|
import (
|
|
"encoding/json"
|
|
"net/http"
|
|
"net/http/httptest"
|
|
"net/url"
|
|
"strings"
|
|
"testing"
|
|
"time"
|
|
|
|
"github.com/gin-gonic/gin"
|
|
"github.com/gorilla/websocket"
|
|
"github.com/router-for-me/CLIProxyAPI/v7/sdk/cliproxy/auth"
|
|
)
|
|
|
|
func TestHandleDirectWebsocketRejectsClientSecretModelMismatch(t *testing.T) {
|
|
gin.SetMode(gin.TestMode)
|
|
handler := NewHandler(auth.NewManager(nil, nil, nil), nil)
|
|
router := gin.New()
|
|
router.GET("/v1/realtime", func(c *gin.Context) {
|
|
c.Set(ClientSecretSessionContextKey, json.RawMessage(`{"type":"realtime","model":"gpt-live-1-codex"}`))
|
|
c.Set(ClientSecretPrincipalContextKey, "sess_123")
|
|
c.Next()
|
|
}, handler.HandleRealtimeWebsocket)
|
|
request := httptest.NewRequest(http.MethodGet, "/v1/realtime?model=another-live-model", nil)
|
|
request.Header.Set("Connection", "Upgrade")
|
|
request.Header.Set("Upgrade", "websocket")
|
|
recorder := httptest.NewRecorder()
|
|
router.ServeHTTP(recorder, request)
|
|
if recorder.Code != http.StatusForbidden {
|
|
t.Fatalf("status = %d, want %d; body=%s", recorder.Code, http.StatusForbidden, recorder.Body.String())
|
|
}
|
|
}
|
|
|
|
func TestHandleDirectWebsocketAppliesClientSecretSession(t *testing.T) {
|
|
gin.SetMode(gin.TestMode)
|
|
upstreamUpdate := make(chan []byte, 1)
|
|
upstreamServer := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, request *http.Request) {
|
|
upgrader := websocket.Upgrader{CheckOrigin: func(*http.Request) bool { return true }}
|
|
connection, errUpgrade := upgrader.Upgrade(writer, request, nil)
|
|
if errUpgrade != nil {
|
|
return
|
|
}
|
|
defer func() { _ = connection.Close() }()
|
|
_, payload, errRead := connection.ReadMessage()
|
|
if errRead != nil {
|
|
return
|
|
}
|
|
upstreamUpdate <- append([]byte(nil), payload...)
|
|
_ = connection.WriteMessage(websocket.TextMessage, []byte(`{"type":"session.created"}`))
|
|
}))
|
|
defer upstreamServer.Close()
|
|
|
|
manager := auth.NewManager(nil, nil, nil)
|
|
manager.RegisterExecutor(&captureExecutor{})
|
|
registerCredential(t, manager, &auth.Auth{
|
|
ID: "codex-oauth",
|
|
Provider: "codex",
|
|
Status: auth.StatusActive,
|
|
Metadata: map[string]any{"access_token": "oauth-token"},
|
|
})
|
|
handler := NewHandler(manager, nil)
|
|
handler.sidebandAPIBaseURL = "ws" + strings.TrimPrefix(upstreamServer.URL, "http") + "/v1"
|
|
router := gin.New()
|
|
router.GET("/v1/realtime", func(c *gin.Context) {
|
|
c.Set(ClientSecretSessionContextKey, json.RawMessage(`{"type":"realtime","model":"gpt-live-1-codex","instructions":"help"}`))
|
|
c.Set(ClientSecretPrincipalContextKey, "sess_123")
|
|
c.Next()
|
|
}, handler.HandleRealtimeWebsocket)
|
|
downstreamServer := httptest.NewServer(router)
|
|
defer downstreamServer.Close()
|
|
|
|
wsURL := "ws" + strings.TrimPrefix(downstreamServer.URL, "http") + "/v1/realtime?model=gpt-realtime"
|
|
connection, _, errDial := websocket.DefaultDialer.Dial(wsURL, nil)
|
|
if errDial != nil {
|
|
t.Fatalf("dial downstream websocket: %v", errDial)
|
|
}
|
|
defer func() { _ = connection.Close() }()
|
|
_, _, _ = connection.ReadMessage()
|
|
|
|
select {
|
|
case update := <-upstreamUpdate:
|
|
var event struct {
|
|
Type string `json:"type"`
|
|
Session struct {
|
|
Model string `json:"model"`
|
|
Instructions string `json:"instructions"`
|
|
} `json:"session"`
|
|
}
|
|
if errUnmarshal := json.Unmarshal(update, &event); errUnmarshal != nil {
|
|
t.Fatalf("unmarshal session update: %v", errUnmarshal)
|
|
}
|
|
if event.Type != "session.update" || event.Session.Model != "" || event.Session.Instructions != "help" {
|
|
t.Fatalf("session update = %+v", event)
|
|
}
|
|
case <-time.After(time.Second):
|
|
t.Fatal("session update not captured")
|
|
}
|
|
}
|
|
|
|
func TestHandleDirectWebsocketRelaysStandardRealtimeFrames(t *testing.T) {
|
|
gin.SetMode(gin.TestMode)
|
|
|
|
upstreamRequest := make(chan *http.Request, 1)
|
|
upstreamMessage := make(chan string, 1)
|
|
upstreamServer := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, request *http.Request) {
|
|
upgrader := websocket.Upgrader{CheckOrigin: func(*http.Request) bool { return true }}
|
|
connection, errUpgrade := upgrader.Upgrade(writer, request, nil)
|
|
if errUpgrade != nil {
|
|
return
|
|
}
|
|
defer func() { _ = connection.Close() }()
|
|
upstreamRequest <- request.Clone(request.Context())
|
|
if errWrite := connection.WriteMessage(websocket.TextMessage, []byte(`{"type":"session.created"}`)); errWrite != nil {
|
|
return
|
|
}
|
|
messageType, payload, errRead := connection.ReadMessage()
|
|
if errRead != nil {
|
|
return
|
|
}
|
|
upstreamMessage <- string(payload)
|
|
_ = connection.WriteMessage(messageType, append([]byte("echo:"), payload...))
|
|
}))
|
|
defer upstreamServer.Close()
|
|
|
|
manager := auth.NewManager(nil, nil, nil)
|
|
manager.RegisterExecutor(&captureExecutor{})
|
|
registerCredential(t, manager, &auth.Auth{
|
|
ID: "codex-oauth",
|
|
Provider: "codex",
|
|
Status: auth.StatusActive,
|
|
Metadata: map[string]any{
|
|
"access_token": "oauth-token",
|
|
"account_id": "account-123",
|
|
},
|
|
})
|
|
handler := NewHandler(manager, nil)
|
|
handler.sidebandAPIBaseURL = "ws" + strings.TrimPrefix(upstreamServer.URL, "http") + "/v1"
|
|
|
|
router := gin.New()
|
|
router.GET("/v1/realtime", handler.HandleRealtimeWebsocket)
|
|
downstreamServer := httptest.NewServer(router)
|
|
defer downstreamServer.Close()
|
|
|
|
wsURL := "ws" + strings.TrimPrefix(downstreamServer.URL, "http") + "/v1/realtime?model=gpt-realtime"
|
|
downstreamHeaders := make(http.Header)
|
|
downstreamHeaders.Set("OpenAI-Alpha", "quicksilver=v2")
|
|
connection, _, errDial := websocket.DefaultDialer.Dial(wsURL, downstreamHeaders)
|
|
if errDial != nil {
|
|
t.Fatalf("dial downstream websocket: %v", errDial)
|
|
}
|
|
defer func() { _ = connection.Close() }()
|
|
|
|
_, created, errRead := connection.ReadMessage()
|
|
if errRead != nil {
|
|
t.Fatalf("read session.created: %v", errRead)
|
|
}
|
|
if string(created) != `{"type":"session.created"}` {
|
|
t.Fatalf("created event = %s", created)
|
|
}
|
|
const event = `{"type":"response.create"}`
|
|
if errWrite := connection.WriteMessage(websocket.TextMessage, []byte(event)); errWrite != nil {
|
|
t.Fatalf("write downstream event: %v", errWrite)
|
|
}
|
|
_, echoed, errRead := connection.ReadMessage()
|
|
if errRead != nil {
|
|
t.Fatalf("read echoed event: %v", errRead)
|
|
}
|
|
if string(echoed) != "echo:"+event {
|
|
t.Fatalf("echoed event = %s", echoed)
|
|
}
|
|
|
|
select {
|
|
case request := <-upstreamRequest:
|
|
if request.Header.Get("Authorization") != "Bearer oauth-token" {
|
|
t.Fatalf("Authorization = %q", request.Header.Get("Authorization"))
|
|
}
|
|
if request.Header.Get("Chatgpt-Account-Id") != "account-123" {
|
|
t.Fatalf("Chatgpt-Account-Id = %q", request.Header.Get("Chatgpt-Account-Id"))
|
|
}
|
|
if request.Header.Get("OpenAI-Alpha") != "" {
|
|
t.Fatalf("OpenAI-Alpha must not be forwarded, got %q", request.Header.Get("OpenAI-Alpha"))
|
|
}
|
|
query, errParse := url.ParseQuery(request.URL.RawQuery)
|
|
if errParse != nil {
|
|
t.Fatalf("parse upstream query: %v", errParse)
|
|
}
|
|
if query.Get("model") != "gpt-realtime" || query.Has("intent") {
|
|
t.Fatalf("upstream query = %v", query)
|
|
}
|
|
case <-time.After(time.Second):
|
|
t.Fatal("upstream request not captured")
|
|
}
|
|
select {
|
|
case payload := <-upstreamMessage:
|
|
if payload != event {
|
|
t.Fatalf("upstream event = %s", payload)
|
|
}
|
|
case <-time.After(time.Second):
|
|
t.Fatal("upstream event not captured")
|
|
}
|
|
}
|