diff --git a/examples/realtime-openai-go/README.md b/examples/realtime-openai-go/README.md new file mode 100644 index 000000000..aae8fe183 --- /dev/null +++ b/examples/realtime-openai-go/README.md @@ -0,0 +1,89 @@ +# OpenAI Go SDK Realtime Voice Example + +This example sends spoken audio to CLIProxyAPI and saves the model's spoken reply as a WAV file. + +It uses the official [`github.com/openai/openai-go/v3`](https://github.com/openai/openai-go) SDK to create a short-lived Realtime client secret. The official Go SDK currently exposes the Realtime REST resources but does not provide a WebSocket connection helper, so `github.com/gorilla/websocket` is used for the standard Realtime audio events. + +## Prerequisites + +1. Start CLIProxyAPI with at least one working ChatGPT/Codex OAuth credential. +2. Configure a proxy API key in `config.yaml`. +3. Use Go 1.26 or newer. +4. Prepare a PCM WAV file with these exact properties: + - 24,000 Hz sample rate + - 16-bit signed PCM + - mono + - little-endian + +Convert an existing recording with FFmpeg: + +```bash +ffmpeg -i recording.m4a -ar 24000 -ac 1 -c:a pcm_s16le question.wav +``` + +## Run + +```bash +cd examples/realtime-openai-go + +OPENAI_BASE_URL="http://127.0.0.1:8317/v1" \ +OPENAI_API_KEY="your-proxy-api-key" \ +OPENAI_REALTIME_MODEL="gpt-realtime-2.1" \ +OPENAI_REALTIME_INPUT_WAV="question.wav" \ +OPENAI_REALTIME_OUTPUT_WAV="response.wav" \ +go run . +``` + +Expected output: + +```text +Loaded question.wav (2.4s, 115200 PCM bytes) +Connected to ws://127.0.0.1:8317/v1/realtime?model=gpt-realtime-2.1 using model gpt-realtime-2.1 and voice marin +Sent 2.4s of speech audio +Assistant transcript: The connection is working correctly. +Saved spoken response to response.wav (1.8s, 86400 PCM bytes) +``` + +Play the response: + +```bash +# macOS +afplay response.wav + +# Linux +aplay response.wav + +# Cross-platform with FFmpeg +ffplay -autoexit response.wav +``` + +## Environment variables + +| Variable | Required | Default | Description | +| --- | --- | --- | --- | +| `OPENAI_API_KEY` | Yes | — | API key configured for CLIProxyAPI. | +| `OPENAI_REALTIME_INPUT_WAV` | Yes | — | Input speech WAV file. It must be 24kHz, 16-bit, mono PCM. | +| `OPENAI_REALTIME_OUTPUT_WAV` | No | `response.wav` | Destination for the spoken response. | +| `OPENAI_BASE_URL` | No | `http://127.0.0.1:8317/v1` | CLIProxyAPI OpenAI-compatible base URL. `/v1` is added when the URL has no path. | +| `OPENAI_REALTIME_MODEL` | No | `gpt-realtime-2.1` | Standard Realtime model name. CLIProxyAPI uses it for the upstream standard WebSocket while selecting a compatible Codex OAuth credential internally. | +| `OPENAI_REALTIME_VOICE` | No | `marin` | Realtime output voice. Other common values include `cedar`, `alloy`, `ash`, `coral`, and `echo`. | +| `OPENAI_REALTIME_INSTRUCTIONS` | No | Short spoken response instruction | Session instructions attached to the client secret. | +| `OPENAI_REALTIME_DEBUG` | No | `false` | Print every received Realtime server event. | + +## Audio flow + +1. The official OpenAI Go SDK calls `POST /v1/realtime/client_secrets` with an audio session configured for 24kHz PCM input and output. +2. The returned local `ek_...` credential authenticates the `/v1/realtime` WebSocket. +3. Input WAV samples are sent in 200ms `input_audio_buffer.append` chunks. +4. The client sends `input_audio_buffer.commit` and `response.create`. +5. Base64 `response.output_audio.delta` events are decoded and written to the output WAV. + +The client secret returned by CLIProxyAPI is local to that proxy instance and is not valid against `api.openai.com`. + +## Test + +```bash +go test -race ./... +``` + +The test starts an in-process HTTP/WebSocket server and verifies client-secret configuration, input audio streaming, output audio decoding, and WAV generation. diff --git a/examples/realtime-openai-go/go.mod b/examples/realtime-openai-go/go.mod new file mode 100644 index 000000000..6de15fc0c --- /dev/null +++ b/examples/realtime-openai-go/go.mod @@ -0,0 +1,15 @@ +module github.com/router-for-me/CLIProxyAPI/v7/examples/realtime-openai-go + +go 1.26.0 + +require ( + github.com/gorilla/websocket v1.5.3 + github.com/openai/openai-go/v3 v3.50.0 +) + +require ( + github.com/tidwall/gjson v1.19.0 // indirect + github.com/tidwall/match v1.1.1 // indirect + github.com/tidwall/pretty v1.2.1 // indirect + github.com/tidwall/sjson v1.2.5 // indirect +) diff --git a/examples/realtime-openai-go/go.sum b/examples/realtime-openai-go/go.sum new file mode 100644 index 000000000..8df405b41 --- /dev/null +++ b/examples/realtime-openai-go/go.sum @@ -0,0 +1,14 @@ +github.com/gorilla/websocket v1.5.3 h1:saDtZ6Pbx/0u+bgYQ3q96pZgCzfhKXGPqt7kZ72aNNg= +github.com/gorilla/websocket v1.5.3/go.mod h1:YR8l580nyteQvAITg2hZ9XVh4b55+EU/adAjf1fMHhE= +github.com/openai/openai-go/v3 v3.50.0 h1:CXn+C8a10oQiI5CMyMbCiykhITVhVxhdHX8j3CfLa2U= +github.com/openai/openai-go/v3 v3.50.0/go.mod h1:Ogjo0gDct+Jm7yCqaCjLGQGygeV8xNfNHV1/yKvCji0= +github.com/tidwall/gjson v1.14.2/go.mod h1:/wbyibRr2FHMks5tjHJ5F8dMZh3AcwJEMf5vlfC0lxk= +github.com/tidwall/gjson v1.19.0 h1:xwxm7n691Uf3u5OFjzngavjGTh55KX5q/9w9xHW88JU= +github.com/tidwall/gjson v1.19.0/go.mod h1:V37/opeE/JbLUOfH0QTXiNez2l0RUjYUhpT4szFQAfc= +github.com/tidwall/match v1.1.1 h1:+Ho715JplO36QYgwN9PGYNhgZvoUSc9X2c80KVTi+GA= +github.com/tidwall/match v1.1.1/go.mod h1:eRSPERbgtNPcGhD8UCthc6PmLEQXEWd3PRB5JTxsfmM= +github.com/tidwall/pretty v1.2.0/go.mod h1:ITEVvHYasfjBbM0u2Pg8T2nJnzm8xPwvNhhsoaGGjNU= +github.com/tidwall/pretty v1.2.1 h1:qjsOFOWWQl+N3RsoF5/ssm1pHmJJwhjlSbZ51I6wMl4= +github.com/tidwall/pretty v1.2.1/go.mod h1:ITEVvHYasfjBbM0u2Pg8T2nJnzm8xPwvNhhsoaGGjNU= +github.com/tidwall/sjson v1.2.5 h1:kLy8mja+1c9jlljvWTlSazM7cKDRfJuR/bOJhcY5NcY= +github.com/tidwall/sjson v1.2.5/go.mod h1:Fvgq9kS/6ociJEDnK0Fk1cpYF4FIW6ZF7LAe+6jwd28= diff --git a/examples/realtime-openai-go/main.go b/examples/realtime-openai-go/main.go new file mode 100644 index 000000000..a16f26e27 --- /dev/null +++ b/examples/realtime-openai-go/main.go @@ -0,0 +1,341 @@ +package main + +import ( + "bytes" + "context" + "encoding/base64" + "encoding/json" + "errors" + "fmt" + "io" + "net/http" + "net/url" + "os" + "os/signal" + "strings" + "sync" + "syscall" + "time" + + "github.com/gorilla/websocket" + "github.com/openai/openai-go/v3" + "github.com/openai/openai-go/v3/option" + "github.com/openai/openai-go/v3/realtime" +) + +const ( + defaultBaseURL = "http://127.0.0.1:8317/v1" + defaultModel = "gpt-realtime-2.1" + defaultInstructions = "Listen to the user's speech and reply with a short spoken response." + defaultOutputWAV = "response.wav" + defaultVoice = "marin" + audioSampleRate = 24000 + audioBytesPerSample = 2 + audioChunkDuration = 200 * time.Millisecond +) + +type appConfig struct { + baseURL string + apiKey string + model string + inputWAV string + outputWAV string + instructions string + voice string + debug bool +} + +type realtimeServerEvent struct { + Type string `json:"type"` + Delta string `json:"delta"` + Error *struct { + Message string `json:"message"` + Type string `json:"type"` + Code string `json:"code"` + } `json:"error,omitempty"` + Response *struct { + Status string `json:"status"` + } `json:"response,omitempty"` +} + +func main() { + ctx, stop := signal.NotifyContext(context.Background(), os.Interrupt, syscall.SIGTERM) + defer stop() + + cfg, errConfig := loadConfig() + if errConfig != nil { + fmt.Fprintf(os.Stderr, "configuration error: %v\n", errConfig) + os.Exit(1) + } + if errRun := run(ctx, cfg, os.Stdout); errRun != nil { + fmt.Fprintf(os.Stderr, "realtime example failed: %v\n", errRun) + os.Exit(1) + } +} + +func loadConfig() (appConfig, error) { + baseURL, errBaseURL := normalizeBaseURL(envOrDefault("OPENAI_BASE_URL", defaultBaseURL)) + if errBaseURL != nil { + return appConfig{}, errBaseURL + } + apiKey := strings.TrimSpace(os.Getenv("OPENAI_API_KEY")) + if apiKey == "" { + return appConfig{}, errors.New("OPENAI_API_KEY is required") + } + inputWAV := strings.TrimSpace(os.Getenv("OPENAI_REALTIME_INPUT_WAV")) + if inputWAV == "" { + return appConfig{}, errors.New("OPENAI_REALTIME_INPUT_WAV is required") + } + return appConfig{ + baseURL: baseURL, + apiKey: apiKey, + model: envOrDefault("OPENAI_REALTIME_MODEL", defaultModel), + inputWAV: inputWAV, + outputWAV: envOrDefault("OPENAI_REALTIME_OUTPUT_WAV", defaultOutputWAV), + instructions: envOrDefault("OPENAI_REALTIME_INSTRUCTIONS", defaultInstructions), + voice: envOrDefault("OPENAI_REALTIME_VOICE", defaultVoice), + debug: strings.EqualFold(strings.TrimSpace(os.Getenv("OPENAI_REALTIME_DEBUG")), "true"), + }, nil +} + +func run(ctx context.Context, cfg appConfig, output io.Writer) error { + inputPCM, errInput := readPCM16WAV(cfg.inputWAV) + if errInput != nil { + return fmt.Errorf("read input WAV: %w", errInput) + } + inputDuration := time.Duration(len(inputPCM)) * time.Second / (audioSampleRate * audioBytesPerSample) + fmt.Fprintf(output, "Loaded %s (%s, %d PCM bytes)\n", cfg.inputWAV, inputDuration.Round(time.Millisecond), len(inputPCM)) + + client := openai.NewClient( + option.WithAPIKey(cfg.apiKey), + option.WithBaseURL(cfg.baseURL), + ) + pcmFormat := realtime.RealtimeAudioFormatsUnionParam{ + OfAudioPCM: &realtime.RealtimeAudioFormatsAudioPCMParam{ + Rate: audioSampleRate, + Type: "audio/pcm", + }, + } + credentialCtx, cancelCredential := context.WithTimeout(ctx, 30*time.Second) + secret, errSecret := client.Realtime.ClientSecrets.New(credentialCtx, realtime.ClientSecretNewParams{ + ExpiresAfter: realtime.ClientSecretNewParamsExpiresAfter{ + Anchor: "created_at", + Seconds: openai.Int(600), + }, + Session: realtime.ClientSecretNewParamsSessionUnion{ + OfRealtime: &realtime.RealtimeSessionCreateRequestParam{ + Model: realtime.RealtimeSessionCreateRequestModel(cfg.model), + Instructions: openai.String(cfg.instructions), + OutputModalities: []string{"audio"}, + Audio: realtime.RealtimeAudioConfigParam{ + Input: realtime.RealtimeAudioConfigInputParam{ + Format: pcmFormat, + }, + Output: realtime.RealtimeAudioConfigOutputParam{ + Format: pcmFormat, + Voice: realtime.RealtimeAudioConfigOutputVoiceUnionParam{ + OfString: openai.String(cfg.voice), + }, + }, + }, + }, + }, + }, option.WithJSONSet("session.audio.input.turn_detection", nil)) + cancelCredential() + if errSecret != nil { + return fmt.Errorf("create Realtime client secret with official SDK: %w", errSecret) + } + if secret == nil || strings.TrimSpace(secret.Value) == "" { + return errors.New("official SDK returned an empty Realtime client secret") + } + + websocketURL, errWebsocketURL := realtimeWebsocketURL(cfg.baseURL, cfg.model) + if errWebsocketURL != nil { + return errWebsocketURL + } + headers := make(http.Header) + headers.Set("Authorization", "Bearer "+secret.Value) + connection, response, errDial := websocket.DefaultDialer.DialContext(ctx, websocketURL, headers) + if errDial != nil { + return websocketHandshakeError(response, errDial) + } + var closeOnce sync.Once + closeConnection := func() { + closeOnce.Do(func() { + if errClose := connection.Close(); errClose != nil && !websocket.IsCloseError(errClose, websocket.CloseNormalClosure, websocket.CloseGoingAway) { + fmt.Fprintf(output, "warning: close websocket: %v\n", errClose) + } + }) + } + defer closeConnection() + + connectionDone := make(chan struct{}) + defer close(connectionDone) + go func() { + select { + case <-ctx.Done(): + closeConnection() + case <-connectionDone: + } + }() + + fmt.Fprintf(output, "Connected to %s using model %s and voice %s\n", websocketURL, cfg.model, cfg.voice) + if errSend := sendInputAudio(connection, inputPCM); errSend != nil { + return errSend + } + fmt.Fprintf(output, "Sent %s of speech audio\n", inputDuration.Round(time.Millisecond)) + + var responsePCM bytes.Buffer + fmt.Fprint(output, "Assistant transcript: ") + if errRead := readRealtimeResponse(ctx, connection, output, &responsePCM, cfg.debug); errRead != nil { + return errRead + } + if responsePCM.Len() == 0 { + return errors.New("Realtime response completed without audio") + } + if errWrite := writePCM16WAV(cfg.outputWAV, responsePCM.Bytes()); errWrite != nil { + return fmt.Errorf("write output WAV: %w", errWrite) + } + responseDuration := time.Duration(responsePCM.Len()) * time.Second / (audioSampleRate * audioBytesPerSample) + fmt.Fprintf(output, "Saved spoken response to %s (%s, %d PCM bytes)\n", cfg.outputWAV, responseDuration.Round(time.Millisecond), responsePCM.Len()) + return nil +} + +func sendInputAudio(connection *websocket.Conn, pcm []byte) error { + chunkSize := int(int64(audioSampleRate*audioBytesPerSample) * int64(audioChunkDuration) / int64(time.Second)) + for offset := 0; offset < len(pcm); offset += chunkSize { + end := min(offset+chunkSize, len(pcm)) + if errWrite := connection.WriteJSON(map[string]any{ + "type": "input_audio_buffer.append", + "audio": base64.StdEncoding.EncodeToString(pcm[offset:end]), + }); errWrite != nil { + return fmt.Errorf("append input audio: %w", errWrite) + } + } + if errWrite := connection.WriteJSON(map[string]any{"type": "input_audio_buffer.commit"}); errWrite != nil { + return fmt.Errorf("commit input audio: %w", errWrite) + } + if errWrite := connection.WriteJSON(map[string]any{ + "type": "response.create", + "response": map[string]any{ + "output_modalities": []string{"audio"}, + }, + }); errWrite != nil { + return fmt.Errorf("request spoken Realtime response: %w", errWrite) + } + return nil +} + +func readRealtimeResponse(ctx context.Context, connection *websocket.Conn, output io.Writer, audioOutput *bytes.Buffer, debug bool) error { + for { + _, payload, errRead := connection.ReadMessage() + if errRead != nil { + if errContext := ctx.Err(); errContext != nil { + return errContext + } + if websocket.IsCloseError(errRead, websocket.CloseNormalClosure, websocket.CloseGoingAway) { + return errors.New("Realtime WebSocket closed before response.done") + } + return fmt.Errorf("read Realtime event: %w", errRead) + } + var event realtimeServerEvent + if errUnmarshal := json.Unmarshal(payload, &event); errUnmarshal != nil { + return fmt.Errorf("decode Realtime event: %w", errUnmarshal) + } + if debug { + fmt.Fprintf(output, "\n[event] %s\n", payload) + } + switch event.Type { + case "response.output_audio.delta", "response.audio.delta": + audio, errDecode := base64.StdEncoding.DecodeString(event.Delta) + if errDecode != nil { + return fmt.Errorf("decode response audio delta: %w", errDecode) + } + if audioOutput.Len()+len(audio) > maxOutputPCMBytes { + return fmt.Errorf("response PCM data exceeds %d bytes", maxOutputPCMBytes) + } + if _, errWrite := audioOutput.Write(audio); errWrite != nil { + return fmt.Errorf("buffer response audio: %w", errWrite) + } + case "response.output_audio_transcript.delta", "response.audio_transcript.delta": + fmt.Fprint(output, event.Delta) + case "response.done": + fmt.Fprintln(output) + if event.Response != nil && event.Response.Status != "" && event.Response.Status != "completed" { + return fmt.Errorf("Realtime response finished with status %s", event.Response.Status) + } + return nil + case "error": + if event.Error == nil { + return errors.New("Realtime API returned an unspecified error") + } + return fmt.Errorf("Realtime API error %s/%s: %s", event.Error.Type, event.Error.Code, event.Error.Message) + } + } +} + +func normalizeBaseURL(rawURL string) (string, error) { + parsed, errParse := url.Parse(strings.TrimSpace(rawURL)) + if errParse != nil { + return "", fmt.Errorf("parse OPENAI_BASE_URL: %w", errParse) + } + if parsed.Scheme != "http" && parsed.Scheme != "https" { + return "", errors.New("OPENAI_BASE_URL must use http or https") + } + if parsed.Host == "" { + return "", errors.New("OPENAI_BASE_URL must include a host") + } + parsed.RawQuery = "" + parsed.Fragment = "" + parsed.Path = strings.TrimRight(parsed.Path, "/") + if parsed.Path == "" { + parsed.Path = "/v1" + } + return parsed.String(), nil +} + +func realtimeWebsocketURL(baseURL, model string) (string, error) { + parsed, errParse := url.Parse(baseURL) + if errParse != nil { + return "", fmt.Errorf("parse Realtime base URL: %w", errParse) + } + switch parsed.Scheme { + case "http": + parsed.Scheme = "ws" + case "https": + parsed.Scheme = "wss" + default: + return "", errors.New("Realtime base URL must use http or https") + } + parsed.Path = strings.TrimRight(parsed.Path, "/") + "/realtime" + query := parsed.Query() + query.Set("model", model) + parsed.RawQuery = query.Encode() + return parsed.String(), nil +} + +func websocketHandshakeError(response *http.Response, errDial error) error { + if response == nil { + return fmt.Errorf("connect Realtime WebSocket: %w", errDial) + } + body, errRead := io.ReadAll(io.LimitReader(response.Body, 64<<10)) + errClose := response.Body.Close() + if errRead != nil { + return fmt.Errorf("connect Realtime WebSocket: HTTP %d; read response: %v; dial: %w", response.StatusCode, errRead, errDial) + } + if errClose != nil { + return fmt.Errorf("connect Realtime WebSocket: HTTP %d; close response: %v; dial: %w", response.StatusCode, errClose, errDial) + } + message := strings.TrimSpace(string(body)) + if message == "" { + message = http.StatusText(response.StatusCode) + } + return fmt.Errorf("connect Realtime WebSocket: HTTP %d: %s: %w", response.StatusCode, message, errDial) +} + +func envOrDefault(name, fallback string) string { + if value := strings.TrimSpace(os.Getenv(name)); value != "" { + return value + } + return fallback +} diff --git a/examples/realtime-openai-go/main_test.go b/examples/realtime-openai-go/main_test.go new file mode 100644 index 000000000..aabaf12ea --- /dev/null +++ b/examples/realtime-openai-go/main_test.go @@ -0,0 +1,226 @@ +package main + +import ( + "bytes" + "context" + "encoding/base64" + "encoding/json" + "net/http" + "net/http/httptest" + "os" + "path/filepath" + "strings" + "testing" + + "github.com/gorilla/websocket" +) + +func TestRunSendsAndReceivesSpeechAudio(t *testing.T) { + tmpDir := t.TempDir() + inputPath := filepath.Join(tmpDir, "input.wav") + outputPath := filepath.Join(tmpDir, "response.wav") + inputPCM := make([]byte, 9602) + for index := range inputPCM { + inputPCM[index] = byte(index % 251) + } + if errWrite := writePCM16WAV(inputPath, inputPCM); errWrite != nil { + t.Fatalf("write input WAV: %v", errWrite) + } + responsePCM := []byte{10, 20, 30, 40, 50, 60, 70, 80} + + websocketEvents := make(chan []string, 1) + capturedInput := make(chan []byte, 1) + upgrader := websocket.Upgrader{CheckOrigin: func(*http.Request) bool { return true }} + server := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, request *http.Request) { + switch request.URL.Path { + case "/v1/realtime/client_secrets": + if request.Method != http.MethodPost || request.Header.Get("Authorization") != "Bearer proxy-key" { + http.Error(writer, "invalid client secret request", http.StatusUnauthorized) + return + } + var body map[string]any + if errDecode := json.NewDecoder(request.Body).Decode(&body); errDecode != nil || !validAudioSession(body) { + http.Error(writer, "invalid audio session", http.StatusBadRequest) + return + } + writer.Header().Set("Content-Type", "application/json") + _, _ = writer.Write([]byte(`{ + "value":"ek_test", + "expires_at":4102444800, + "session":{"id":"sess_test","object":"realtime.session","type":"realtime","model":"gpt-realtime"} + }`)) + case "/v1/realtime": + if request.Header.Get("Authorization") != "Bearer ek_test" || request.URL.Query().Get("model") != defaultModel { + http.Error(writer, "invalid websocket request", http.StatusUnauthorized) + return + } + connection, errUpgrade := upgrader.Upgrade(writer, request, nil) + if errUpgrade != nil { + return + } + defer func() { + if errClose := connection.Close(); errClose != nil { + t.Logf("close test websocket: %v", errClose) + } + }() + + types := make([]string, 0, 4) + var receivedPCM bytes.Buffer + for { + _, payload, errRead := connection.ReadMessage() + if errRead != nil { + return + } + var event struct { + Type string `json:"type"` + Audio string `json:"audio"` + } + if errUnmarshal := json.Unmarshal(payload, &event); errUnmarshal != nil { + return + } + types = append(types, event.Type) + if event.Type == "input_audio_buffer.append" { + audio, errDecode := base64.StdEncoding.DecodeString(event.Audio) + if errDecode != nil { + return + } + _, _ = receivedPCM.Write(audio) + } + if event.Type == "response.create" { + break + } + } + websocketEvents <- types + capturedInput <- append([]byte(nil), receivedPCM.Bytes()...) + midpoint := len(responsePCM) / 2 + for _, audio := range [][]byte{responsePCM[:midpoint], responsePCM[midpoint:]} { + if errWrite := connection.WriteJSON(map[string]any{ + "type": "response.output_audio.delta", + "delta": base64.StdEncoding.EncodeToString(audio), + }); errWrite != nil { + return + } + } + if errWrite := connection.WriteJSON(map[string]any{"type": "response.output_audio_transcript.delta", "delta": "Voice response"}); errWrite != nil { + return + } + _ = connection.WriteJSON(map[string]any{"type": "response.done", "response": map[string]any{"status": "completed"}}) + default: + http.NotFound(writer, request) + } + })) + defer server.Close() + + baseURL, errBaseURL := normalizeBaseURL(server.URL + "/v1/") + if errBaseURL != nil { + t.Fatalf("normalizeBaseURL() error = %v", errBaseURL) + } + var output bytes.Buffer + errRun := run(context.Background(), appConfig{ + baseURL: baseURL, + apiKey: "proxy-key", + model: defaultModel, + inputWAV: inputPath, + outputWAV: outputPath, + instructions: defaultInstructions, + voice: defaultVoice, + }, &output) + if errRun != nil { + t.Fatalf("run() error = %v", errRun) + } + if !strings.Contains(output.String(), "Sent") || !strings.Contains(output.String(), "Assistant transcript: Voice response") || !strings.Contains(output.String(), "Saved spoken response") { + t.Fatalf("output = %q", output.String()) + } + select { + case events := <-websocketEvents: + want := []string{"input_audio_buffer.append", "input_audio_buffer.append", "input_audio_buffer.commit", "response.create"} + if strings.Join(events, ",") != strings.Join(want, ",") { + t.Fatalf("client events = %v, want %v", events, want) + } + default: + t.Fatal("websocket events were not captured") + } + select { + case audio := <-capturedInput: + if !bytes.Equal(audio, inputPCM) { + t.Fatalf("input PCM mismatch: got %d bytes, want %d", len(audio), len(inputPCM)) + } + default: + t.Fatal("input audio was not captured") + } + actualResponsePCM, errRead := readPCM16WAV(outputPath) + if errRead != nil { + t.Fatalf("read output WAV: %v", errRead) + } + if !bytes.Equal(actualResponsePCM, responsePCM) { + t.Fatalf("response PCM = %v, want %v", actualResponsePCM, responsePCM) + } +} + +func validAudioSession(body map[string]any) bool { + session, ok := body["session"].(map[string]any) + if !ok || session["type"] != "realtime" || session["model"] != defaultModel { + return false + } + modalities, ok := session["output_modalities"].([]any) + if !ok || len(modalities) != 1 || modalities[0] != "audio" { + return false + } + audio, ok := session["audio"].(map[string]any) + if !ok { + return false + } + input, inputOK := audio["input"].(map[string]any) + output, outputOK := audio["output"].(map[string]any) + if !inputOK || !outputOK { + return false + } + inputFormat, inputFormatOK := input["format"].(map[string]any) + outputFormat, outputFormatOK := output["format"].(map[string]any) + if !inputFormatOK || !outputFormatOK { + return false + } + _, turnDetectionPresent := input["turn_detection"] + return inputFormat["type"] == "audio/pcm" && inputFormat["rate"] == float64(audioSampleRate) && + outputFormat["type"] == "audio/pcm" && outputFormat["rate"] == float64(audioSampleRate) && + output["voice"] == defaultVoice && turnDetectionPresent && input["turn_detection"] == nil +} + +func TestNormalizeBaseURLAddsV1(t *testing.T) { + baseURL, errNormalize := normalizeBaseURL("http://127.0.0.1:8317/") + if errNormalize != nil { + t.Fatalf("normalizeBaseURL() error = %v", errNormalize) + } + if baseURL != "http://127.0.0.1:8317/v1" { + t.Fatalf("baseURL = %q", baseURL) + } + websocketURL, errWebsocketURL := realtimeWebsocketURL(baseURL, defaultModel) + if errWebsocketURL != nil { + t.Fatalf("realtimeWebsocketURL() error = %v", errWebsocketURL) + } + wantWebsocketURL := "ws://127.0.0.1:8317/v1/realtime?model=" + defaultModel + if websocketURL != wantWebsocketURL { + t.Fatalf("websocketURL = %q", websocketURL) + } +} + +func TestReadPCM16WAVRejectsWrongSampleRate(t *testing.T) { + path := filepath.Join(t.TempDir(), "wrong-rate.wav") + if errWrite := writePCM16WAV(path, []byte{1, 2, 3, 4}); errWrite != nil { + t.Fatalf("writePCM16WAV() error = %v", errWrite) + } + payload, errRead := os.ReadFile(path) + if errRead != nil { + t.Fatalf("read WAV: %v", errRead) + } + payload[24] = 0x80 + payload[25] = 0xbb + payload[26] = 0x00 + payload[27] = 0x00 + if errWrite := os.WriteFile(path, payload, 0o644); errWrite != nil { + t.Fatalf("rewrite WAV: %v", errWrite) + } + if _, errRead = readPCM16WAV(path); errRead == nil || !strings.Contains(errRead.Error(), "24000") { + t.Fatalf("readPCM16WAV() error = %v", errRead) + } +} diff --git a/examples/realtime-openai-go/wav.go b/examples/realtime-openai-go/wav.go new file mode 100644 index 000000000..d94d95b90 --- /dev/null +++ b/examples/realtime-openai-go/wav.go @@ -0,0 +1,147 @@ +package main + +import ( + "bytes" + "encoding/binary" + "errors" + "fmt" + "os" +) + +const ( + maxInputPCMBytes = 15 << 20 + maxOutputPCMBytes = 64 << 20 + wavHeaderSize = 44 +) + +func readPCM16WAV(path string) ([]byte, error) { + fileInfo, errStat := os.Stat(path) + if errStat != nil { + return nil, errStat + } + if fileInfo.Size() > maxInputPCMBytes+(1<<20) { + return nil, fmt.Errorf("WAV file is too large: %d bytes", fileInfo.Size()) + } + payload, errRead := os.ReadFile(path) + if errRead != nil { + return nil, errRead + } + if len(payload) < 12 || string(payload[:4]) != "RIFF" || string(payload[8:12]) != "WAVE" { + return nil, errors.New("input is not a RIFF/WAVE file") + } + + var formatFound bool + var audioFormat uint16 + var channels uint16 + var sampleRate uint32 + var bitsPerSample uint16 + var pcm bytes.Buffer + for offset := 12; offset+8 <= len(payload); { + chunkID := string(payload[offset : offset+4]) + chunkSize := int(binary.LittleEndian.Uint32(payload[offset+4 : offset+8])) + chunkStart := offset + 8 + chunkEnd := chunkStart + chunkSize + if chunkSize < 0 || chunkEnd < chunkStart || chunkEnd > len(payload) { + return nil, fmt.Errorf("invalid WAV %q chunk size", chunkID) + } + switch chunkID { + case "fmt ": + if chunkSize < 16 { + return nil, errors.New("WAV fmt chunk is too short") + } + audioFormat = binary.LittleEndian.Uint16(payload[chunkStart : chunkStart+2]) + channels = binary.LittleEndian.Uint16(payload[chunkStart+2 : chunkStart+4]) + sampleRate = binary.LittleEndian.Uint32(payload[chunkStart+4 : chunkStart+8]) + bitsPerSample = binary.LittleEndian.Uint16(payload[chunkStart+14 : chunkStart+16]) + formatFound = true + case "data": + if pcm.Len()+chunkSize > maxInputPCMBytes { + return nil, fmt.Errorf("WAV PCM data exceeds %d bytes", maxInputPCMBytes) + } + _, _ = pcm.Write(payload[chunkStart:chunkEnd]) + } + offset = chunkEnd + if chunkSize%2 != 0 { + offset++ + } + } + if !formatFound { + return nil, errors.New("WAV fmt chunk is missing") + } + if audioFormat != 1 { + return nil, fmt.Errorf("WAV audio format must be PCM (1), got %d", audioFormat) + } + if channels != 1 { + return nil, fmt.Errorf("WAV must be mono, got %d channels", channels) + } + if sampleRate != audioSampleRate { + return nil, fmt.Errorf("WAV sample rate must be %d Hz, got %d Hz", audioSampleRate, sampleRate) + } + if bitsPerSample != 16 { + return nil, fmt.Errorf("WAV must use 16-bit samples, got %d bits", bitsPerSample) + } + if pcm.Len() == 0 { + return nil, errors.New("WAV data chunk is empty or missing") + } + if pcm.Len()%audioBytesPerSample != 0 { + return nil, errors.New("WAV PCM data contains an incomplete sample") + } + return append([]byte(nil), pcm.Bytes()...), nil +} + +func writePCM16WAV(path string, pcm []byte) error { + if len(pcm) == 0 { + return errors.New("cannot write an empty WAV response") + } + if len(pcm) > maxOutputPCMBytes { + return fmt.Errorf("response PCM data exceeds %d bytes", maxOutputPCMBytes) + } + if len(pcm)%audioBytesPerSample != 0 { + return errors.New("response PCM data contains an incomplete sample") + } + + var payload bytes.Buffer + payload.Grow(wavHeaderSize + len(pcm)) + writeString := func(value string) error { + _, errWrite := payload.WriteString(value) + return errWrite + } + writeValue := func(value any) error { + return binary.Write(&payload, binary.LittleEndian, value) + } + if errWrite := writeString("RIFF"); errWrite != nil { + return errWrite + } + if errWrite := writeValue(uint32(36 + len(pcm))); errWrite != nil { + return errWrite + } + if errWrite := writeString("WAVEfmt "); errWrite != nil { + return errWrite + } + for _, value := range []any{ + uint32(16), + uint16(1), + uint16(1), + uint32(audioSampleRate), + uint32(audioSampleRate * audioBytesPerSample), + uint16(audioBytesPerSample), + uint16(16), + } { + if errWrite := writeValue(value); errWrite != nil { + return errWrite + } + } + if errWrite := writeString("data"); errWrite != nil { + return errWrite + } + if errWrite := writeValue(uint32(len(pcm))); errWrite != nil { + return errWrite + } + if _, errWrite := payload.Write(pcm); errWrite != nil { + return errWrite + } + if errWrite := os.WriteFile(path, payload.Bytes(), 0o644); errWrite != nil { + return errWrite + } + return nil +} diff --git a/internal/api/server_middleware.go b/internal/api/server_middleware.go index 8511d4238..447f280cf 100644 --- a/internal/api/server_middleware.go +++ b/internal/api/server_middleware.go @@ -5,6 +5,7 @@ import ( "strings" "github.com/gin-gonic/gin" + codexlive "github.com/router-for-me/CLIProxyAPI/v7/internal/client/codex/live" "github.com/router-for-me/CLIProxyAPI/v7/internal/config" "github.com/router-for-me/CLIProxyAPI/v7/internal/home" "github.com/router-for-me/CLIProxyAPI/v7/internal/logging" @@ -23,6 +24,10 @@ var corsExposedResponseHeaders = []string{ "X-CPA-HOME-BUILD-DATE", "X-SERVER-VERSION", "X-SERVER-BUILD-DATE", + "Location", + "Retry-After", + "X-Request-Id", + "OpenAI-Request-Id", } var corsExposedResponseHeadersJoined = strings.Join(corsExposedResponseHeaders, ", ") @@ -144,6 +149,14 @@ func corsMiddleware() gin.HandlerFunc { // using the configured authentication providers. When no providers are available, // it allows all requests (legacy behaviour). func AuthMiddleware(manager *sdkaccess.Manager) gin.HandlerFunc { + return accessAuthMiddleware(manager, false) +} + +func realtimeStandardAuthMiddleware(manager *sdkaccess.Manager) gin.HandlerFunc { + return accessAuthMiddleware(manager, true) +} + +func accessAuthMiddleware(manager *sdkaccess.Manager, realtimeError bool) gin.HandlerFunc { return func(c *gin.Context) { if manager == nil { c.Next() @@ -167,6 +180,54 @@ func AuthMiddleware(manager *sdkaccess.Manager) gin.HandlerFunc { if statusCode >= http.StatusInternalServerError { log.Errorf("authentication middleware error: %v", err) } + if realtimeError { + errorType := "authentication_error" + code := "invalid_api_key" + if statusCode >= http.StatusInternalServerError { + errorType = "server_error" + code = "authentication_service_error" + } + c.AbortWithStatusJSON(statusCode, gin.H{"error": gin.H{ + "message": err.Message, + "type": errorType, + "param": nil, + "code": code, + }}) + return + } c.AbortWithStatusJSON(statusCode, gin.H{"error": err.Message}) } } + +func realtimeAuthMiddleware(manager *sdkaccess.Manager, handler *codexlive.Handler) gin.HandlerFunc { + fallback := realtimeStandardAuthMiddleware(manager) + return func(c *gin.Context) { + authorization, matched, errAuthenticate := handler.AuthenticateClientSecret(c.Request) + if !matched { + fallback(c) + return + } + if errAuthenticate != nil { + c.AbortWithStatusJSON(http.StatusUnauthorized, gin.H{"error": gin.H{ + "message": errAuthenticate.Error(), + "type": "invalid_request_error", + "param": nil, + "code": "invalid_realtime_client_secret", + }}) + return + } + principal := authorization.IssuerPrincipal + if principal == "" { + principal = authorization.Principal + } + provider := authorization.IssuerProvider + if provider == "" { + provider = "realtime-client-secret" + } + c.Set("userApiKey", principal) + c.Set("accessProvider", provider) + c.Set(codexlive.ClientSecretSessionContextKey, authorization.Session) + c.Set(codexlive.ClientSecretPrincipalContextKey, authorization.Principal) + c.Next() + } +} diff --git a/internal/api/server_routes.go b/internal/api/server_routes.go index 19af6260a..d57af9c70 100644 --- a/internal/api/server_routes.go +++ b/internal/api/server_routes.go @@ -80,11 +80,25 @@ func (s *Server) setupRoutes() { v1.POST("/alpha/search", s.codexAlphaSearch) v1.POST("/live", s.codexLiveHandler.Handle) v1.GET("/live/:call_id", s.codexLiveHandler.HandleSideband) - v1.POST("/realtime/calls", s.codexLiveHandler.Handle) - v1.GET("/realtime/calls/:call_id", s.codexLiveHandler.HandleSideband) - v1.GET("/realtime", s.codexLiveHandler.HandleSideband) } + realtimeAuth := realtimeAuthMiddleware(s.accessManager, s.codexLiveHandler) + standardAuth := realtimeStandardAuthMiddleware(s.accessManager) + s.engine.GET("/v1/realtime", realtimeAuth, s.codexLiveHandler.HandleRealtimeWebsocket) + s.engine.POST("/v1/realtime", realtimeAuth, s.codexLiveHandler.Handle) + s.engine.POST("/v1/realtime/calls", realtimeAuth, s.codexLiveHandler.Handle) + s.engine.GET("/v1/realtime/calls/:call_id", realtimeAuth, s.codexLiveHandler.HandleSideband) + s.engine.POST("/v1/realtime/client_secrets", standardAuth, s.codexLiveHandler.CreateClientSecret) + s.engine.POST("/v1/realtime/sessions", standardAuth, s.codexLiveHandler.CreateLegacySession) + s.engine.POST("/v1/realtime/transcription_sessions", standardAuth, s.codexLiveHandler.HandleTranscriptionSession) + s.engine.GET("/v1/realtime/translations", realtimeAuth, s.codexLiveHandler.HandleTranslation) + s.engine.POST("/v1/realtime/translations", realtimeAuth, s.codexLiveHandler.HandleTranslation) + s.engine.POST("/v1/realtime/translations/client_secrets", standardAuth, s.codexLiveHandler.HandleTranslation) + s.engine.POST("/v1/realtime/calls/:call_id/hangup", standardAuth, s.codexLiveHandler.HandleHangup) + s.engine.POST("/v1/realtime/calls/:call_id/accept", standardAuth, s.codexLiveHandler.HandleSIPControl) + s.engine.POST("/v1/realtime/calls/:call_id/reject", standardAuth, s.codexLiveHandler.HandleSIPControl) + s.engine.POST("/v1/realtime/calls/:call_id/refer", standardAuth, s.codexLiveHandler.HandleSIPControl) + openaiV1 := s.engine.Group("/openai/v1") openaiV1.Use(AuthMiddleware(s.accessManager)) { diff --git a/internal/api/server_test.go b/internal/api/server_test.go index e8a885cb2..65a9ac49b 100644 --- a/internal/api/server_test.go +++ b/internal/api/server_test.go @@ -528,6 +528,84 @@ func TestCodexLiveRoutesRequireAuthAndAreRegistered(t *testing.T) { } } +func TestRealtimeStandardRoutesAndClientSecretAuth(t *testing.T) { + server := newTestServer(t) + + unauthorizedSecret := httptest.NewRequest(http.MethodPost, "/v1/realtime/client_secrets", strings.NewReader(`{"session":{"type":"realtime","model":"gpt-realtime"}}`)) + unauthorizedSecretRecorder := httptest.NewRecorder() + server.engine.ServeHTTP(unauthorizedSecretRecorder, unauthorizedSecret) + if unauthorizedSecretRecorder.Code != http.StatusUnauthorized { + t.Fatalf("client_secrets unauthorized status = %d, want %d", unauthorizedSecretRecorder.Code, http.StatusUnauthorized) + } + var unauthorizedResponse struct { + Error struct { + Type string `json:"type"` + Code string `json:"code"` + } `json:"error"` + } + if errUnmarshal := json.Unmarshal(unauthorizedSecretRecorder.Body.Bytes(), &unauthorizedResponse); errUnmarshal != nil { + t.Fatalf("unmarshal unauthorized response: %v", errUnmarshal) + } + if unauthorizedResponse.Error.Type != "authentication_error" || unauthorizedResponse.Error.Code != "invalid_api_key" { + t.Fatalf("unauthorized error = %+v", unauthorizedResponse.Error) + } + + secretRequest := httptest.NewRequest(http.MethodPost, "/v1/realtime/client_secrets", strings.NewReader(`{"session":{"type":"realtime","model":"gpt-realtime"}}`)) + secretRequest.Header.Set("Authorization", "Bearer test-key") + secretRecorder := httptest.NewRecorder() + server.engine.ServeHTTP(secretRecorder, secretRequest) + if secretRecorder.Code != http.StatusOK { + t.Fatalf("client_secrets status = %d, want %d; body=%s", secretRecorder.Code, http.StatusOK, secretRecorder.Body.String()) + } + var secretResponse struct { + Value string `json:"value"` + } + if errUnmarshal := json.Unmarshal(secretRecorder.Body.Bytes(), &secretResponse); errUnmarshal != nil { + t.Fatalf("unmarshal client secret: %v", errUnmarshal) + } + if secretResponse.Value == "" { + t.Fatal("client secret is empty") + } + + callRequest := httptest.NewRequest(http.MethodPost, "/v1/realtime/calls", strings.NewReader("v=0\r\n")) + callRequest.Header.Set("Authorization", "Bearer "+secretResponse.Value) + callRequest.Header.Set("Content-Type", "application/sdp") + callRecorder := httptest.NewRecorder() + server.engine.ServeHTTP(callRecorder, callRequest) + if callRecorder.Code != http.StatusServiceUnavailable { + t.Fatalf("ephemeral call status = %d, want %d; body=%s", callRecorder.Code, http.StatusServiceUnavailable, callRecorder.Body.String()) + } + + for _, testCase := range []struct { + method string + path string + status int + }{ + {method: http.MethodGet, path: "/v1/realtime?model=gpt-realtime", status: http.StatusUpgradeRequired}, + {method: http.MethodPost, path: "/v1/realtime", status: http.StatusServiceUnavailable}, + {method: http.MethodPost, path: "/v1/realtime/sessions", status: http.StatusOK}, + {method: http.MethodPost, path: "/v1/realtime/transcription_sessions", status: http.StatusNotImplemented}, + {method: http.MethodGet, path: "/v1/realtime/translations", status: http.StatusNotImplemented}, + {method: http.MethodPost, path: "/v1/realtime/translations", status: http.StatusNotImplemented}, + {method: http.MethodPost, path: "/v1/realtime/translations/client_secrets", status: http.StatusNotImplemented}, + {method: http.MethodPost, path: "/v1/realtime/calls/call-123/accept", status: http.StatusNotImplemented}, + {method: http.MethodPost, path: "/v1/realtime/calls/call-123/reject", status: http.StatusNotImplemented}, + {method: http.MethodPost, path: "/v1/realtime/calls/call-123/refer", status: http.StatusNotImplemented}, + {method: http.MethodPost, path: "/v1/realtime/calls/call-123/hangup", status: http.StatusNotFound}, + } { + request := httptest.NewRequest(testCase.method, testCase.path, nil) + request.Header.Set("Authorization", "Bearer test-key") + recorder := httptest.NewRecorder() + server.engine.ServeHTTP(recorder, request) + if recorder.Code != testCase.status { + t.Errorf("%s %s status = %d, want %d; body=%s", testCase.method, testCase.path, recorder.Code, testCase.status, recorder.Body.String()) + } + if testCase.method == http.MethodGet && testCase.path == "/v1/realtime?model=gpt-realtime" && recorder.Header().Get("Upgrade") != "websocket" { + t.Errorf("Upgrade header = %q, want websocket", recorder.Header().Get("Upgrade")) + } + } +} + func TestCodexAlphaSearchForwardsRequest(t *testing.T) { server := newTestServer(t) executor := &codexSearchCaptureExecutor{} diff --git a/internal/client/codex/live/capabilities.go b/internal/client/codex/live/capabilities.go new file mode 100644 index 000000000..6a4e44ed4 --- /dev/null +++ b/internal/client/codex/live/capabilities.go @@ -0,0 +1,215 @@ +package live + +import ( + "context" + "io" + "net/http" + "net/url" + "strings" + + "github.com/gin-gonic/gin" + "github.com/router-for-me/CLIProxyAPI/v7/internal/clienterror" + "github.com/router-for-me/CLIProxyAPI/v7/internal/logging" + "github.com/router-for-me/CLIProxyAPI/v7/internal/runtime/executor/helps" + "github.com/router-for-me/CLIProxyAPI/v7/sdk/cliproxy/auth" + coreexecutor "github.com/router-for-me/CLIProxyAPI/v7/sdk/cliproxy/executor" + log "github.com/sirupsen/logrus" +) + +// HandleTranslation reports that the Codex OAuth upstream has no translation session capability. +func (h *Handler) HandleTranslation(c *gin.Context) { + writeCapabilityNotSupported(c, "Realtime translation sessions") +} + +// HandleTranscriptionSession reports that the Codex OAuth upstream has no transcription-only capability. +func (h *Handler) HandleTranscriptionSession(c *gin.Context) { + writeCapabilityNotSupported(c, "Realtime transcription-only sessions") +} + +// HandleSIPControl reports that the Codex OAuth upstream has no SIP dialog capability. +func (h *Handler) HandleSIPControl(c *gin.Context) { + action := "control" + if c != nil && c.Request != nil && c.Request.URL != nil { + parts := strings.Split(strings.Trim(c.Request.URL.Path, "/"), "/") + if len(parts) > 0 && strings.TrimSpace(parts[len(parts)-1]) != "" { + action = parts[len(parts)-1] + } + } + writeCapabilityNotSupported(c, "Realtime SIP "+action) +} + +// HandleHangup forwards hangup for a locally created WebRTC call using its pinned OAuth credential. +func (h *Handler) HandleHangup(c *gin.Context) { + if h == nil || h.authManager == nil || h.sessions == nil { + writeRealtimeError(c, http.StatusServiceUnavailable, "Codex live session service unavailable", "server_error", "realtime_session_unavailable") + return + } + callID := strings.TrimSpace(c.Param("call_id")) + if !callIDPattern.MatchString(callID) { + writeRealtimeError(c, http.StatusBadRequest, "Invalid Realtime call ID", "invalid_request_error", "invalid_call_id") + return + } + session, ok := h.sessions.peek(callID) + if !ok { + writeRealtimeError(c, http.StatusNotFound, "Realtime call not found", "invalid_request_error", "realtime_call_not_found") + return + } + + if ownerPrincipal, ownerProvider := requestOwner(c); session.ownerPrincipal != "" && (ownerPrincipal != session.ownerPrincipal || ownerProvider != session.ownerProvider) { + writeRealtimeError(c, http.StatusForbidden, "Realtime call belongs to another API principal", "invalid_request_error", "realtime_call_scope_mismatch") + return + } + + ctx := context.WithValue(c.Request.Context(), "gin", c) + var activeSelection *auth.HomeDispatchSelection + var temporarySelection bool + var selected *auth.Auth + if session.homeSelection != nil && session.homeSelection.Active() { + activeSelection = session.homeSelection + selected = activeSelection.CloneAuth() + } else { + selectionOpts := coreexecutor.Options{ + Headers: liveSelectionHeaders(c), + Metadata: map[string]any{ + coreexecutor.PinnedAuthMetadataKey: session.authID, + coreexecutor.ExecutionSessionMetadataKey: callID, + }, + } + selection, selectedAuth, errSelect := h.selectOAuth(ctx, session.model, selectionOpts) + if errSelect != nil { + writeSelectionError(c, errSelect) + return + } + activeSelection = selection + selected = selectedAuth + temporarySelection = selection != nil + } + var selectionRelease func() + if activeSelection != nil { + attemptCtx, releaseAttempt, errAttempt := activeSelection.AttemptContext(ctx) + if errAttempt != nil { + if temporarySelection { + activeSelection.End("attempt_bind_failed") + } + writeRealtimeError(c, http.StatusServiceUnavailable, errAttempt.Error(), "server_error", "realtime_upstream_unavailable") + return + } + ctx = attemptCtx + selectionRelease = releaseAttempt + } + defer func() { + if selectionRelease != nil { + selectionRelease() + } + if temporarySelection && activeSelection != nil { + activeSelection.End("request_closed") + } + }() + if selected == nil { + writeRealtimeError(c, http.StatusServiceUnavailable, "Codex auth unavailable", "server_error", "codex_auth_unavailable") + return + } + logging.SetGinCPATraceID(c, selected.EnsureIndex()) + + body, errRead := readBody(c.Request.Body) + if errRead != nil { + writeRealtimeError(c, http.StatusBadRequest, errRead.Error(), "invalid_request_error", "invalid_request") + return + } + upstreamURL := h.realtimeHTTPBaseURL() + "/realtime/calls/" + url.PathEscape(callID) + "/hangup" + baseHeaders := protocolHeaders(c.Request.Header) + if contentType := strings.TrimSpace(c.GetHeader("Content-Type")); contentType != "" { + baseHeaders.Set("Content-Type", contentType) + } + runtimeConfig := h.currentConfig() + performRequest := func(current *auth.Auth) (*http.Response, error) { + headers := baseHeaders.Clone() + setAccountHeader(headers, current) + request, errRequest := h.authManager.NewHttpRequest(ctx, current, http.MethodPost, upstreamURL, body, headers) + if errRequest != nil { + return nil, errRequest + } + authType, authValue := current.AccountInfo() + helps.RecordAPIRequest(ctx, runtimeConfig, helps.UpstreamRequestLog{ + URL: upstreamURL, + Method: http.MethodPost, + Headers: headersForLogging(request.Header), + Body: body, + Provider: "codex", + AuthID: current.ID, + AuthLabel: current.Label, + AuthType: authType, + AuthValue: authValue, + }) + return h.authManager.HttpRequest(ctx, current, request) + } + response, errRequest := performRequest(selected) + if errRequest != nil { + helps.RecordAPIResponseError(ctx, runtimeConfig, errRequest) + writeRealtimeError(c, clienterror.HTTPStatusFromErrorOr(errRequest, http.StatusBadGateway), errRequest.Error(), "api_error", "realtime_upstream_unavailable") + return + } + if activeSelection != nil && response.StatusCode == http.StatusUnauthorized { + h.authManager.ReportHomeUnauthorized(ctx, selected, "codex", session.model) + _, _ = io.Copy(io.Discard, io.LimitReader(response.Body, 1<<20)) + if errClose := response.Body.Close(); errClose != nil { + log.Errorf("codex realtime hangup: close unauthorized response body error: %v", errClose) + } + refreshed, didRefresh, errRefresh := h.authManager.RefreshHomeSelectionAfterUnauthorized(ctx, activeSelection, selected) + if errRefresh != nil { + writeSelectionError(c, errRefresh) + return + } + if !didRefresh || refreshed == nil { + writeRealtimeError(c, http.StatusUnauthorized, "Codex credential unauthorized", "authentication_error", "realtime_upstream_unauthorized") + return + } + selected = refreshed + logging.SetGinCPATraceID(c, selected.EnsureIndex()) + response, errRequest = performRequest(selected) + if errRequest != nil { + helps.RecordAPIResponseError(ctx, runtimeConfig, errRequest) + writeRealtimeError(c, clienterror.HTTPStatusFromErrorOr(errRequest, http.StatusBadGateway), errRequest.Error(), "api_error", "realtime_upstream_unavailable") + return + } + if response.StatusCode == http.StatusUnauthorized { + h.authManager.ReportHomeUnauthorized(ctx, selected, "codex", session.model) + } + } + defer func() { + if errClose := response.Body.Close(); errClose != nil { + log.Errorf("codex realtime hangup: close response body error: %v", errClose) + } + }() + responseBody, errResponse := readLimitedBody(response.Body) + if errResponse != nil { + helps.RecordAPIResponseError(ctx, runtimeConfig, errResponse) + writeRealtimeError(c, http.StatusBadGateway, "Failed to read Realtime hangup response", "api_error", "realtime_upstream_unavailable") + return + } + helps.RecordAPIResponseMetadata(ctx, runtimeConfig, response.StatusCode, callResponseHeaders(response.Header)) + helps.AppendAPIResponseChunk(ctx, runtimeConfig, responseBody) + if response.StatusCode >= http.StatusOK && response.StatusCode < http.StatusMultipleChoices { + if selectionRelease != nil { + selectionRelease() + selectionRelease = nil + } + h.sessions.complete(session, "client_hangup") + } + if contentType := response.Header.Get("Content-Type"); contentType != "" { + c.Header("Content-Type", contentType) + } + copyRealtimeHandshakeHeaders(c.Writer.Header(), response.Header) + c.Status(response.StatusCode) + if _, errWrite := c.Writer.Write(responseBody); errWrite != nil { + log.WithError(errWrite).Warn("codex realtime hangup: write response body failed") + } +} + +func (h *Handler) realtimeHTTPBaseURL() string { + return strings.TrimRight(websocketHTTPURL(h.sidebandAPIBaseURL), "/") +} + +func writeCapabilityNotSupported(c *gin.Context, capability string) { + writeRealtimeError(c, http.StatusNotImplemented, capability+" are not supported by the ChatGPT/Codex OAuth upstream", "not_supported_error", "realtime_capability_not_supported") +} diff --git a/internal/client/codex/live/capabilities_test.go b/internal/client/codex/live/capabilities_test.go new file mode 100644 index 000000000..680ed0ef7 --- /dev/null +++ b/internal/client/codex/live/capabilities_test.go @@ -0,0 +1,97 @@ +package live + +import ( + "io" + "net/http" + "net/http/httptest" + "strings" + "testing" + + "github.com/gin-gonic/gin" + "github.com/router-for-me/CLIProxyAPI/v7/sdk/cliproxy/auth" +) + +func TestHandleHangupForwardsPinnedOAuthCall(t *testing.T) { + gin.SetMode(gin.TestMode) + manager := auth.NewManager(nil, nil, nil) + executor := &captureExecutor{ + statusCode: http.StatusOK, + responseBody: io.NopCloser(strings.NewReader(`{"status":"ok"}`)), + } + manager.RegisterExecutor(executor) + 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.sessions.put("call-123", liveSession{ + authID: "codex-oauth", + model: defaultLiveModel, + ownerPrincipal: "owner-key", + ownerProvider: "static", + }) + + router := gin.New() + router.POST("/v1/realtime/calls/:call_id/hangup", func(c *gin.Context) { + c.Set("userApiKey", "owner-key") + c.Set("accessProvider", "static") + c.Next() + }, handler.HandleHangup) + request := httptest.NewRequest(http.MethodPost, "/v1/realtime/calls/call-123/hangup", nil) + recorder := httptest.NewRecorder() + router.ServeHTTP(recorder, request) + if recorder.Code != http.StatusOK { + t.Fatalf("status = %d, want %d; body=%s", recorder.Code, http.StatusOK, recorder.Body.String()) + } + if executor.request == nil || executor.request.URL.String() != "https://api.openai.com/v1/realtime/calls/call-123/hangup" { + t.Fatalf("upstream request = %#v", executor.request) + } + if _, ok := handler.sessions.peek("call-123"); ok { + t.Fatal("successful hangup retained session") + } +} + +func TestHandleHangupRejectsDifferentAPIPrincipal(t *testing.T) { + gin.SetMode(gin.TestMode) + handler := NewHandler(auth.NewManager(nil, nil, nil), nil) + handler.sessions.put("call-123", liveSession{ + authID: "codex-oauth", + model: defaultLiveModel, + ownerPrincipal: "owner-key", + ownerProvider: "static", + }) + router := gin.New() + router.POST("/v1/realtime/calls/:call_id/hangup", func(c *gin.Context) { + c.Set("userApiKey", "other-key") + c.Set("accessProvider", "static") + c.Next() + }, handler.HandleHangup) + request := httptest.NewRequest(http.MethodPost, "/v1/realtime/calls/call-123/hangup", nil) + 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 TestUnsupportedRealtimeCapabilitiesUseStandardError(t *testing.T) { + gin.SetMode(gin.TestMode) + handler := NewHandler(nil, nil) + router := gin.New() + router.POST("/v1/realtime/transcription_sessions", handler.HandleTranscriptionSession) + router.POST("/v1/realtime/calls/:call_id/accept", handler.HandleSIPControl) + + for _, path := range []string{"/v1/realtime/transcription_sessions", "/v1/realtime/calls/call-123/accept"} { + request := httptest.NewRequest(http.MethodPost, path, nil) + recorder := httptest.NewRecorder() + router.ServeHTTP(recorder, request) + if recorder.Code != http.StatusNotImplemented { + t.Errorf("%s status = %d, want %d", path, recorder.Code, http.StatusNotImplemented) + } + if !strings.Contains(recorder.Body.String(), `"type":"not_supported_error"`) || !strings.Contains(recorder.Body.String(), `"code":"realtime_capability_not_supported"`) { + t.Errorf("%s body = %s", path, recorder.Body.String()) + } + } +} diff --git a/internal/client/codex/live/client_secret.go b/internal/client/codex/live/client_secret.go new file mode 100644 index 000000000..5f44300bd --- /dev/null +++ b/internal/client/codex/live/client_secret.go @@ -0,0 +1,419 @@ +package live + +import ( + "crypto/rand" + "encoding/base64" + "encoding/json" + "errors" + "fmt" + "io" + "net/http" + "strings" + "sync" + "time" + + "github.com/gin-gonic/gin" +) + +const ( + ClientSecretSessionContextKey = "codexLiveClientSecretSession" + ClientSecretPrincipalContextKey = "codexLiveClientSecretPrincipal" + clientSecretPrefix = "ek_" + clientSecretDefaultLifetime = 10 * time.Minute + clientSecretMinimumLifetime = 10 * time.Second + clientSecretMaximumLifetime = 2 * time.Hour + clientSecretMaxBodySize = 64 << 10 + clientSecretMaxEntries = 1024 + clientSecretMaxEntriesPerIssuer = 64 +) + +var ( + errInvalidClientSecret = errors.New("Realtime client secret is invalid or expired") + errClientSecretCapacity = errors.New("Realtime client secret capacity exhausted") + errUnsupportedSessionType = errors.New("Realtime session type is not supported") +) + +// ClientSecretAuthorization contains the local session configuration associated with an ephemeral key. +type ClientSecretAuthorization struct { + Principal string + IssuerPrincipal string + IssuerProvider string + Session json.RawMessage +} + +type clientSecretEntry struct { + authorization ClientSecretAuthorization + expiresAt time.Time +} + +type clientSecretStore struct { + mu sync.Mutex + entries map[string]clientSecretEntry + now func() time.Time +} + +type clientSecretCreateRequest struct { + Session json.RawMessage `json:"session"` + ExpiresAfter *struct { + Anchor string `json:"anchor"` + Seconds int64 `json:"seconds"` + } `json:"expires_after,omitempty"` +} + +type clientSecretCreateResponse struct { + Value string `json:"value"` + ExpiresAt int64 `json:"expires_at"` + Session json.RawMessage `json:"session"` +} + +func newClientSecretStore() *clientSecretStore { + return &clientSecretStore{ + entries: make(map[string]clientSecretEntry), + now: time.Now, + } +} + +func (s *clientSecretStore) create(session json.RawMessage, lifetime time.Duration, issuerPrincipal, issuerProvider string) (string, ClientSecretAuthorization, time.Time, error) { + if s == nil { + return "", ClientSecretAuthorization{}, time.Time{}, errors.New("Realtime client secret store unavailable") + } + token, errToken := randomRealtimeID(clientSecretPrefix, 32) + if errToken != nil { + return "", ClientSecretAuthorization{}, time.Time{}, errToken + } + sessionID, errSessionID := randomRealtimeID("sess_", 18) + if errSessionID != nil { + return "", ClientSecretAuthorization{}, time.Time{}, errSessionID + } + authorization := ClientSecretAuthorization{ + Principal: sessionID, + IssuerPrincipal: strings.TrimSpace(issuerPrincipal), + IssuerProvider: strings.TrimSpace(issuerProvider), + Session: append(json.RawMessage(nil), session...), + } + now := s.currentTime() + expiresAt := now.Add(lifetime) + s.mu.Lock() + s.removeExpiredLocked(now) + if len(s.entries) >= clientSecretMaxEntries { + s.mu.Unlock() + return "", ClientSecretAuthorization{}, time.Time{}, errClientSecretCapacity + } + if authorization.IssuerPrincipal != "" { + issuerEntries := 0 + for _, entry := range s.entries { + if entry.authorization.IssuerPrincipal == authorization.IssuerPrincipal && entry.authorization.IssuerProvider == authorization.IssuerProvider { + issuerEntries++ + } + } + if issuerEntries >= clientSecretMaxEntriesPerIssuer { + s.mu.Unlock() + return "", ClientSecretAuthorization{}, time.Time{}, errClientSecretCapacity + } + } + s.entries[token] = clientSecretEntry{authorization: authorization, expiresAt: expiresAt} + s.mu.Unlock() + return token, authorization, expiresAt, nil +} + +func (s *clientSecretStore) authenticate(token string) (ClientSecretAuthorization, error) { + if s == nil || !strings.HasPrefix(token, clientSecretPrefix) { + return ClientSecretAuthorization{}, errInvalidClientSecret + } + now := s.currentTime() + s.mu.Lock() + entry, ok := s.entries[token] + if !ok || !entry.expiresAt.After(now) { + delete(s.entries, token) + s.mu.Unlock() + return ClientSecretAuthorization{}, errInvalidClientSecret + } + s.mu.Unlock() + entry.authorization.Session = append(json.RawMessage(nil), entry.authorization.Session...) + return entry.authorization, nil +} + +func (s *clientSecretStore) close() { + if s == nil { + return + } + s.mu.Lock() + clear(s.entries) + s.mu.Unlock() +} + +func (s *clientSecretStore) currentTime() time.Time { + if s != nil && s.now != nil { + return s.now() + } + return time.Now() +} + +func (s *clientSecretStore) removeExpiredLocked(now time.Time) { + for token, entry := range s.entries { + if !entry.expiresAt.After(now) { + delete(s.entries, token) + } + } +} + +func readClientSecretBody(body io.Reader) ([]byte, error) { + if body == nil { + return nil, nil + } + payload, errRead := io.ReadAll(io.LimitReader(body, clientSecretMaxBodySize+1)) + if errRead != nil { + return nil, fmt.Errorf("failed to read Realtime client secret request: %w", errRead) + } + if len(payload) > clientSecretMaxBodySize { + return nil, errBodyTooLarge + } + return payload, nil +} + +func randomRealtimeID(prefix string, size int) (string, error) { + payload := make([]byte, size) + if _, errRead := rand.Read(payload); errRead != nil { + return "", fmt.Errorf("generate Realtime identifier: %w", errRead) + } + return prefix + base64.RawURLEncoding.EncodeToString(payload), nil +} + +// AuthenticateClientSecret validates a local ephemeral key when the request carries one. +func (h *Handler) AuthenticateClientSecret(request *http.Request) (ClientSecretAuthorization, bool, error) { + token := bearerToken(request) + if !strings.HasPrefix(token, clientSecretPrefix) { + return ClientSecretAuthorization{}, false, nil + } + if h == nil || h.clientSecrets == nil { + return ClientSecretAuthorization{}, true, errInvalidClientSecret + } + authorization, errAuthenticate := h.clientSecrets.authenticate(token) + return authorization, true, errAuthenticate +} + +func bearerToken(request *http.Request) string { + if request == nil { + return "" + } + authorization := strings.TrimSpace(request.Header.Get("Authorization")) + const bearerPrefix = "Bearer " + if len(authorization) < len(bearerPrefix) || !strings.EqualFold(authorization[:len(bearerPrefix)], bearerPrefix) { + return "" + } + return strings.TrimSpace(authorization[len(bearerPrefix):]) +} + +// CreateClientSecret creates a short-lived credential scoped to this proxy. +func (h *Handler) CreateClientSecret(c *gin.Context) { + if h == nil || h.clientSecrets == nil { + writeRealtimeError(c, http.StatusServiceUnavailable, "Realtime client secret service unavailable", "server_error", "realtime_client_secret_unavailable") + return + } + body, errRead := readClientSecretBody(c.Request.Body) + if errRead != nil { + status := http.StatusBadRequest + if errors.Is(errRead, errBodyTooLarge) { + status = http.StatusRequestEntityTooLarge + } + writeRealtimeError(c, status, errRead.Error(), "invalid_request_error", "invalid_request") + return + } + var request clientSecretCreateRequest + if len(strings.TrimSpace(string(body))) > 0 { + if errUnmarshal := json.Unmarshal(body, &request); errUnmarshal != nil { + writeRealtimeError(c, http.StatusBadRequest, "Invalid Realtime client secret request", "invalid_request_error", "invalid_request") + return + } + } + h.createClientSecret(c, request.Session, request.ExpiresAfter, false) +} + +// CreateLegacySession implements the deprecated Realtime session credential endpoint. +func (h *Handler) CreateLegacySession(c *gin.Context) { + if h == nil || h.clientSecrets == nil { + writeRealtimeError(c, http.StatusServiceUnavailable, "Realtime client secret service unavailable", "server_error", "realtime_client_secret_unavailable") + return + } + body, errRead := readClientSecretBody(c.Request.Body) + if errRead != nil { + status := http.StatusBadRequest + if errors.Is(errRead, errBodyTooLarge) { + status = http.StatusRequestEntityTooLarge + } + writeRealtimeError(c, status, errRead.Error(), "invalid_request_error", "invalid_request") + return + } + h.createClientSecret(c, json.RawMessage(body), nil, true) +} + +func (h *Handler) createClientSecret(c *gin.Context, session json.RawMessage, expiresAfter *struct { + Anchor string `json:"anchor"` + Seconds int64 `json:"seconds"` +}, legacy bool) { + lifetime, errLifetime := clientSecretLifetime(expiresAfter) + if errLifetime != nil { + writeRealtimeError(c, http.StatusBadRequest, errLifetime.Error(), "invalid_request_error", "invalid_expires_after") + return + } + clientSession, upstreamSession, errSession := normalizeClientSecretSession(session) + if errSession != nil { + if errors.Is(errSession, errUnsupportedSessionType) { + writeRealtimeError(c, http.StatusNotImplemented, errSession.Error(), "not_supported_error", "realtime_capability_not_supported") + return + } + writeRealtimeError(c, http.StatusBadRequest, errSession.Error(), "invalid_request_error", "invalid_session") + return + } + issuerPrincipal, _ := c.Get("userApiKey") + issuerProvider, _ := c.Get("accessProvider") + issuerPrincipalValue, _ := issuerPrincipal.(string) + issuerProviderValue, _ := issuerProvider.(string) + token, authorization, expiresAt, errCreate := h.clientSecrets.create(upstreamSession, lifetime, issuerPrincipalValue, issuerProviderValue) + if errCreate != nil { + if errors.Is(errCreate, errClientSecretCapacity) { + c.Header("Retry-After", "1") + writeRealtimeError(c, http.StatusTooManyRequests, errCreate.Error(), "rate_limit_error", "realtime_client_secret_capacity_exhausted") + return + } + writeRealtimeError(c, http.StatusInternalServerError, "Failed to create Realtime client secret", "server_error", "realtime_client_secret_failed") + return + } + responseSession, errResponse := realtimeSessionResponse(clientSession, authorization.Principal, expiresAt) + if errResponse != nil { + writeRealtimeError(c, http.StatusInternalServerError, "Failed to encode Realtime session", "server_error", "realtime_session_failed") + return + } + c.Header("Cache-Control", "no-store") + if legacy { + var response map[string]any + if errUnmarshal := json.Unmarshal(responseSession, &response); errUnmarshal != nil { + writeRealtimeError(c, http.StatusInternalServerError, "Failed to encode Realtime session", "server_error", "realtime_session_failed") + return + } + response["client_secret"] = gin.H{"value": token, "expires_at": expiresAt.Unix()} + c.JSON(http.StatusOK, response) + return + } + c.JSON(http.StatusOK, clientSecretCreateResponse{ + Value: token, + ExpiresAt: expiresAt.Unix(), + Session: responseSession, + }) +} + +func clientSecretLifetime(expiresAfter *struct { + Anchor string `json:"anchor"` + Seconds int64 `json:"seconds"` +}) (time.Duration, error) { + if expiresAfter == nil { + return clientSecretDefaultLifetime, nil + } + if expiresAfter.Anchor != "" && expiresAfter.Anchor != "created_at" { + return 0, errors.New("expires_after.anchor must be created_at") + } + minimumSeconds := int64(clientSecretMinimumLifetime / time.Second) + maximumSeconds := int64(clientSecretMaximumLifetime / time.Second) + if expiresAfter.Seconds < minimumSeconds || expiresAfter.Seconds > maximumSeconds { + return 0, fmt.Errorf("expires_after.seconds must be between %d and %d", minimumSeconds, maximumSeconds) + } + return time.Duration(expiresAfter.Seconds) * time.Second, nil +} + +func normalizeClientSecretSession(session json.RawMessage) (json.RawMessage, json.RawMessage, error) { + trimmedSession := strings.TrimSpace(string(session)) + if trimmedSession == "" || trimmedSession == "null" { + session = json.RawMessage(`{"type":"realtime","model":"gpt-realtime"}`) + } + var clientSession map[string]any + if errUnmarshal := json.Unmarshal(session, &clientSession); errUnmarshal != nil || clientSession == nil { + return nil, nil, errors.New("session must be a valid JSON object") + } + sessionType, _ := clientSession["type"].(string) + if strings.TrimSpace(sessionType) == "" { + sessionType = "realtime" + clientSession["type"] = sessionType + } + if sessionType != "realtime" { + return nil, nil, fmt.Errorf("%w by the Codex OAuth upstream: %q", errUnsupportedSessionType, sessionType) + } + model, _ := clientSession["model"].(string) + if strings.TrimSpace(model) == "" { + model = "gpt-realtime" + clientSession["model"] = model + } + clientEncoded, errMarshal := json.Marshal(clientSession) + if errMarshal != nil { + return nil, nil, fmt.Errorf("encode Realtime session: %w", errMarshal) + } + clientSession["model"] = codexRealtimeModel(model) + upstreamEncoded, errMarshal := json.Marshal(clientSession) + if errMarshal != nil { + return nil, nil, fmt.Errorf("encode Codex Realtime session: %w", errMarshal) + } + return clientEncoded, upstreamEncoded, nil +} + +func realtimeSessionResponse(session json.RawMessage, sessionID string, expiresAt time.Time) (json.RawMessage, error) { + var response map[string]any + if errUnmarshal := json.Unmarshal(session, &response); errUnmarshal != nil { + return nil, errUnmarshal + } + response["id"] = sessionID + response["object"] = "realtime.session" + response["expires_at"] = expiresAt.Unix() + return json.Marshal(response) +} + +func codexRealtimeModel(model string) string { + trimmed := strings.TrimSpace(model) + lower := strings.ToLower(trimmed) + if lower == "" || lower == "gpt-realtime" || strings.HasPrefix(lower, "gpt-realtime-") || strings.Contains(lower, "realtime-preview") { + return defaultLiveModel + } + return trimmed +} + +func liveSelectionHeaders(c *gin.Context) http.Header { + if c == nil || c.Request == nil { + return make(http.Header) + } + headers := c.Request.Header.Clone() + if _, ok := c.Get(ClientSecretPrincipalContextKey); ok { + headers.Del("Authorization") + headers.Del("Proxy-Authorization") + } + return headers +} + +func requestOwner(c *gin.Context) (string, string) { + if c == nil { + return "", "" + } + principalValue, _ := c.Get("userApiKey") + providerValue, _ := c.Get("accessProvider") + principal, _ := principalValue.(string) + provider, _ := providerValue.(string) + return strings.TrimSpace(principal), strings.TrimSpace(provider) +} + +func clientSecretSession(c *gin.Context) json.RawMessage { + if c == nil { + return nil + } + value, ok := c.Get(ClientSecretSessionContextKey) + if !ok { + return nil + } + session, _ := value.(json.RawMessage) + return append(json.RawMessage(nil), session...) +} + +func writeRealtimeError(c *gin.Context, status int, message, errorType, code string) { + c.JSON(status, gin.H{"error": gin.H{ + "message": message, + "type": errorType, + "param": nil, + "code": code, + }}) +} diff --git a/internal/client/codex/live/client_secret_test.go b/internal/client/codex/live/client_secret_test.go new file mode 100644 index 000000000..27474bcf3 --- /dev/null +++ b/internal/client/codex/live/client_secret_test.go @@ -0,0 +1,262 @@ +package live + +import ( + "bytes" + "context" + "encoding/json" + "errors" + "io" + "net/http" + "net/http/httptest" + "strings" + "testing" + "time" + + "github.com/gin-gonic/gin" + "github.com/router-for-me/CLIProxyAPI/v7/sdk/cliproxy/auth" +) + +func TestCreateClientSecretMapsStandardRealtimeModel(t *testing.T) { + gin.SetMode(gin.TestMode) + handler := &Handler{clientSecrets: newClientSecretStore()} + router := gin.New() + router.POST("/v1/realtime/client_secrets", func(c *gin.Context) { + c.Set("userApiKey", "issuer-key") + c.Set("accessProvider", "static") + c.Next() + }, handler.CreateClientSecret) + + request := httptest.NewRequest(http.MethodPost, "/v1/realtime/client_secrets", strings.NewReader(`{ + "session":{"type":"realtime","model":"gpt-realtime","instructions":"help"}, + "expires_after":{"anchor":"created_at","seconds":60} + }`)) + request.Header.Set("Content-Type", "application/json") + recorder := httptest.NewRecorder() + router.ServeHTTP(recorder, request) + if recorder.Code != http.StatusOK { + t.Fatalf("status = %d, want %d; body=%s", recorder.Code, http.StatusOK, recorder.Body.String()) + } + + var response struct { + Value string `json:"value"` + ExpiresAt int64 `json:"expires_at"` + Session struct { + ID string `json:"id"` + Object string `json:"object"` + Type string `json:"type"` + Model string `json:"model"` + Instructions string `json:"instructions"` + } `json:"session"` + } + if errUnmarshal := json.Unmarshal(recorder.Body.Bytes(), &response); errUnmarshal != nil { + t.Fatalf("unmarshal response: %v", errUnmarshal) + } + if !strings.HasPrefix(response.Value, clientSecretPrefix) { + t.Fatalf("client secret = %q", response.Value) + } + if response.ExpiresAt <= time.Now().Unix() { + t.Fatalf("expires_at = %d", response.ExpiresAt) + } + if response.Session.ID == "" || response.Session.Object != "realtime.session" || response.Session.Type != "realtime" { + t.Fatalf("session = %+v", response.Session) + } + if response.Session.Model != "gpt-realtime" || response.Session.Instructions != "help" { + t.Fatalf("client session = %+v", response.Session) + } + + authRequest := httptest.NewRequest(http.MethodPost, "/v1/realtime/calls", nil) + authRequest.Header.Set("Authorization", "Bearer "+response.Value) + authorization, matched, errAuthenticate := handler.AuthenticateClientSecret(authRequest) + if errAuthenticate != nil || !matched { + t.Fatalf("AuthenticateClientSecret() matched=%t error=%v", matched, errAuthenticate) + } + if authorization.Principal != response.Session.ID { + t.Fatalf("principal = %q, want %q", authorization.Principal, response.Session.ID) + } + if authorization.IssuerPrincipal != "issuer-key" || authorization.IssuerProvider != "static" { + t.Fatalf("issuer = %q/%q", authorization.IssuerProvider, authorization.IssuerPrincipal) + } + if got := modelFromJSON(authorization.Session); got != defaultLiveModel { + t.Fatalf("upstream session model = %q, want %q", got, defaultLiveModel) + } +} + +func TestStandardRealtimeCallMapsModelAndLocation(t *testing.T) { + gin.SetMode(gin.TestMode) + manager := auth.NewManager(nil, nil, nil) + executor := &captureExecutor{responseBody: io.NopCloser(strings.NewReader("v=0\r\n"))} + manager.RegisterExecutor(executor) + if _, errRegister := manager.Register(context.Background(), &auth.Auth{ + ID: "codex-oauth", + Provider: "codex", + Status: auth.StatusActive, + Metadata: map[string]any{"access_token": "oauth-token"}, + }); errRegister != nil { + t.Fatalf("register auth: %v", errRegister) + } + handler := NewHandler(manager, nil) + router := gin.New() + router.POST("/v1/realtime/calls", handler.Handle) + + const boundary = "standard-realtime-boundary" + body := multipartBody(boundary, "v=0\r\n", `{"type":"realtime","model":"gpt-realtime"}`) + request := httptest.NewRequest(http.MethodPost, "/v1/realtime/calls", strings.NewReader(body)) + request.Header.Set("Content-Type", "multipart/form-data; boundary="+boundary) + recorder := httptest.NewRecorder() + router.ServeHTTP(recorder, request) + if recorder.Code != http.StatusCreated { + t.Fatalf("status = %d, want %d; body=%s", recorder.Code, http.StatusCreated, recorder.Body.String()) + } + if recorder.Header().Get("Location") != "/v1/realtime/calls/call-123" { + t.Fatalf("Location = %q", recorder.Header().Get("Location")) + } + if got := modelFromJSON(executor.body); got != defaultLiveModel { + t.Fatalf("upstream model = %q, want %q; body=%s", got, defaultLiveModel, executor.body) + } +} + +func TestClientSecretStoreRejectsExpiredToken(t *testing.T) { + store := newClientSecretStore() + now := time.Unix(1700000000, 0) + store.now = func() time.Time { return now } + token, _, _, errCreate := store.create(json.RawMessage(`{"type":"realtime","model":"gpt-live-1-codex"}`), time.Minute, "issuer", "test") + if errCreate != nil { + t.Fatalf("create() error = %v", errCreate) + } + if _, errAuthenticate := store.authenticate(token); errAuthenticate != nil { + t.Fatalf("authenticate() error = %v", errAuthenticate) + } + now = now.Add(time.Minute) + if _, errAuthenticate := store.authenticate(token); errAuthenticate == nil { + t.Fatal("authenticate() accepted expired token") + } +} + +func TestNormalizeClientSecretSessionHandlesWhitespaceNullAndRejectsArrays(t *testing.T) { + clientSession, upstreamSession, errNormalize := normalizeClientSecretSession(json.RawMessage(" null \n")) + if errNormalize != nil { + t.Fatalf("normalize whitespace null: %v", errNormalize) + } + if modelFromJSON(clientSession) != "gpt-realtime" || modelFromJSON(upstreamSession) != defaultLiveModel { + t.Fatalf("client=%s upstream=%s", clientSession, upstreamSession) + } + if _, _, errNormalize = normalizeClientSecretSession(json.RawMessage(`[]`)); errNormalize == nil { + t.Fatal("normalize accepted an array session") + } +} + +func TestReadClientSecretBodyRejectsOversizedSession(t *testing.T) { + _, errRead := readClientSecretBody(bytes.NewReader(make([]byte, clientSecretMaxBodySize+1))) + if !errors.Is(errRead, errBodyTooLarge) { + t.Fatalf("readClientSecretBody() error = %v", errRead) + } +} + +func TestCreateClientSecretRejectsUnsupportedSessionType(t *testing.T) { + gin.SetMode(gin.TestMode) + handler := &Handler{clientSecrets: newClientSecretStore()} + router := gin.New() + router.POST("/v1/realtime/client_secrets", handler.CreateClientSecret) + + request := httptest.NewRequest(http.MethodPost, "/v1/realtime/client_secrets", strings.NewReader(`{"session":{"type":"transcription","model":"gpt-4o-transcribe"}}`)) + recorder := httptest.NewRecorder() + router.ServeHTTP(recorder, request) + if recorder.Code != http.StatusNotImplemented { + t.Fatalf("status = %d, want %d; body=%s", recorder.Code, http.StatusNotImplemented, recorder.Body.String()) + } + if !strings.Contains(recorder.Body.String(), "realtime_capability_not_supported") { + t.Fatalf("body = %s", recorder.Body.String()) + } +} + +func TestLiveSelectionHeadersRemoveLocalClientSecret(t *testing.T) { + gin.SetMode(gin.TestMode) + ginContext, _ := gin.CreateTestContext(httptest.NewRecorder()) + ginContext.Request = httptest.NewRequest(http.MethodPost, "/v1/realtime/calls", nil) + ginContext.Request.Header.Set("Authorization", "Bearer ek_secret") + ginContext.Request.Header.Set("OpenAI-Safety-Identifier", "safe-user") + ginContext.Set(ClientSecretPrincipalContextKey, "sess_123") + headers := liveSelectionHeaders(ginContext) + if headers.Get("Authorization") != "" { + t.Fatalf("Authorization leaked: %q", headers.Get("Authorization")) + } + if headers.Get("OpenAI-Safety-Identifier") != "safe-user" { + t.Fatalf("safety identifier = %q", headers.Get("OpenAI-Safety-Identifier")) + } +} + +func TestSidebandRejectsClientSecretScopeMismatch(t *testing.T) { + gin.SetMode(gin.TestMode) + handler := NewHandler(auth.NewManager(nil, nil, nil), nil) + handler.sessions.put("call-123", liveSession{ + authID: "codex-oauth", + model: defaultLiveModel, + clientSecretPrincipal: "sess_expected", + }) + router := gin.New() + router.GET("/v1/realtime/calls/:call_id", func(c *gin.Context) { + c.Set(ClientSecretPrincipalContextKey, "sess_other") + c.Next() + }, handler.HandleSideband) + request := httptest.NewRequest(http.MethodGet, "/v1/realtime/calls/call-123", 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()) + } + claimed, claim := handler.sessions.claim("call-123") + if claim != sessionClaimAcquired { + t.Fatalf("session claim = %v", claim) + } + handler.sessions.release(claimed) +} + +func TestSidebandRejectsStandardPrincipalScopeMismatch(t *testing.T) { + gin.SetMode(gin.TestMode) + handler := NewHandler(auth.NewManager(nil, nil, nil), nil) + handler.sessions.put("call-123", liveSession{ + authID: "codex-oauth", + model: defaultLiveModel, + ownerPrincipal: "owner-key", + ownerProvider: "static", + }) + router := gin.New() + router.GET("/v1/realtime/calls/:call_id", func(c *gin.Context) { + c.Set("userApiKey", "other-key") + c.Set("accessProvider", "static") + c.Next() + }, handler.HandleSideband) + request := httptest.NewRequest(http.MethodGet, "/v1/realtime/calls/call-123", 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 TestApplyClientSecretCallSession(t *testing.T) { + session := json.RawMessage(`{"type":"realtime","model":"gpt-live-1-codex","instructions":"help"}`) + body, contentType, model, errApply := applyClientSecretCallSession([]byte("v=0\r\n"), "application/sdp", defaultLiveModel, session) + if errApply != nil { + t.Fatalf("applyClientSecretCallSession() error = %v", errApply) + } + if contentType != "application/json" || model != defaultLiveModel { + t.Fatalf("contentType=%q model=%q", contentType, model) + } + var payload struct { + SDP string `json:"sdp"` + Session struct { + Instructions string `json:"instructions"` + } `json:"session"` + } + if errUnmarshal := json.Unmarshal(body, &payload); errUnmarshal != nil { + t.Fatalf("unmarshal body: %v", errUnmarshal) + } + if payload.SDP != "v=0\r\n" || payload.Session.Instructions != "help" { + t.Fatalf("payload = %+v", payload) + } +} diff --git a/internal/client/codex/live/live.go b/internal/client/codex/live/live.go index ac36682ea..4a862e925 100644 --- a/internal/client/codex/live/live.go +++ b/internal/client/codex/live/live.go @@ -38,6 +38,9 @@ var liveProtocolHeaders = []string{ "Session-Id", "Thread-Id", "Originator", + "OpenAI-Safety-Identifier", + "OpenAI-Organization", + "OpenAI-Project", "X-Oai-Attestation", } @@ -46,6 +49,7 @@ type Handler struct { authManager *auth.Manager cfg *config.Config sessions *sessionStore + clientSecrets *clientSecretStore sidebandAPIBaseURL string mediaRelayMu sync.RWMutex mediaRelay mediaRelayFactory @@ -61,6 +65,7 @@ func NewHandler(authManager *auth.Manager, cfg *config.Config) *Handler { authManager: authManager, cfg: cfg, sessions: newSessionStore(), + clientSecrets: newClientSecretStore(), sidebandAPIBaseURL: defaultSidebandAPIBaseURL, } if errUpdate := handler.UpdateConfig(cfg); errUpdate != nil { @@ -162,15 +167,21 @@ func (h *Handler) currentMediaRelay() (mediaRelayFactory, error) { // Close releases all active Codex live sessions. func (h *Handler) Close() { - if h != nil && h.sessions != nil { + if h == nil { + return + } + if h.sessions != nil { h.sessions.closeAll("server_stopped") } + if h.clientSecrets != nil { + h.clientSecrets.close() + } } // Handle forwards a WebRTC SDP bootstrap request to the Codex realtime calls endpoint. func (h *Handler) Handle(c *gin.Context) { if h == nil || h.authManager == nil { - c.JSON(http.StatusServiceUnavailable, gin.H{"error": "Codex auth manager unavailable"}) + writeLiveError(c, http.StatusServiceUnavailable, "Codex auth manager unavailable") return } @@ -180,17 +191,23 @@ func (h *Handler) Handle(c *gin.Context) { if errors.Is(errRead, errBodyTooLarge) { status = http.StatusRequestEntityTooLarge } - c.JSON(status, gin.H{"error": errRead.Error()}) + writeLiveError(c, status, errRead.Error()) return } upstreamBody, upstreamContentType, model, errPayload := prepareCallRequest(body, c.GetHeader("Content-Type")) + if errPayload == nil { + upstreamBody, upstreamContentType, model, errPayload = applyClientSecretCallSession(upstreamBody, upstreamContentType, model, clientSecretSession(c)) + } + if errPayload == nil { + upstreamBody, model, errPayload = rewriteCallRequestModel(upstreamBody, upstreamContentType, model) + } if errPayload != nil { - c.JSON(http.StatusBadRequest, gin.H{"error": errPayload.Error()}) + writeLiveError(c, http.StatusBadRequest, errPayload.Error()) return } runtimeConfig, mediaRelay, mediaRelayErr := h.currentRuntime() if mediaRelayErr != nil { - c.JSON(http.StatusServiceUnavailable, gin.H{"error": mediaRelayErr.Error()}) + writeLiveError(c, http.StatusServiceUnavailable, mediaRelayErr.Error()) return } var mediaSession mediaRelaySession @@ -198,7 +215,7 @@ func (h *Handler) Handle(c *gin.Context) { ctx := context.WithValue(c.Request.Context(), "gin", c) selectionOpts := coreexecutor.Options{ - Headers: c.Request.Header.Clone(), + Headers: liveSelectionHeaders(c), OriginalRequest: body, } selection, selected, errSelect := h.selectOAuth(ctx, model, selectionOpts) @@ -210,7 +227,7 @@ func (h *Handler) Handle(c *gin.Context) { if selection != nil { selection.End("missing_auth") } - c.JSON(http.StatusServiceUnavailable, gin.H{"error": "Codex auth unavailable"}) + writeLiveError(c, http.StatusServiceUnavailable, "Codex auth unavailable") return } @@ -218,7 +235,7 @@ func (h *Handler) Handle(c *gin.Context) { attemptCtx, releaseAttempt, errAttempt := selection.AttemptContext(ctx) if errAttempt != nil { selection.End("attempt_bind_failed") - c.JSON(http.StatusServiceUnavailable, gin.H{"error": errAttempt.Error()}) + writeLiveError(c, http.StatusServiceUnavailable, errAttempt.Error()) return } ctx = attemptCtx @@ -237,7 +254,7 @@ func (h *Handler) Handle(c *gin.Context) { if mediaRelay != nil { clientOffer, errSDP := callRequestSDP(upstreamBody, upstreamContentType) if errSDP != nil { - c.JSON(http.StatusBadRequest, gin.H{"error": errSDP.Error()}) + writeLiveError(c, http.StatusBadRequest, errSDP.Error()) return } var upstreamOffer string @@ -247,7 +264,7 @@ func (h *Handler) Handle(c *gin.Context) { authIndex: selectedIndex, }) if errSDP != nil { - c.JSON(clienterror.HTTPStatusFromErrorOr(errSDP, http.StatusBadGateway), gin.H{"error": errSDP.Error()}) + writeLiveError(c, clienterror.HTTPStatusFromErrorOr(errSDP, http.StatusBadGateway), errSDP.Error()) return } defer func() { @@ -259,7 +276,7 @@ func (h *Handler) Handle(c *gin.Context) { }() upstreamBody, upstreamContentType, errSDP = replaceCallRequestSDP(upstreamBody, upstreamContentType, upstreamOffer) if errSDP != nil { - c.JSON(http.StatusBadRequest, gin.H{"error": errSDP.Error()}) + writeLiveError(c, http.StatusBadRequest, errSDP.Error()) return } } @@ -292,7 +309,7 @@ func (h *Handler) Handle(c *gin.Context) { if selection != nil { selection.End("attempt_canceled") } - c.JSON(clienterror.HTTPStatusFromErrorOr(errContext, http.StatusRequestTimeout), gin.H{"error": errContext.Error()}) + writeLiveError(c, clienterror.HTTPStatusFromErrorOr(errContext, http.StatusRequestTimeout), errContext.Error()) return } resp, errRequest := performRequest(selected) @@ -301,7 +318,7 @@ func (h *Handler) Handle(c *gin.Context) { selection.End("request_failed") } helps.RecordAPIResponseError(ctx, runtimeConfig, errRequest) - c.JSON(clienterror.HTTPStatusFromErrorOr(errRequest, http.StatusBadGateway), gin.H{"error": errRequest.Error()}) + writeLiveError(c, clienterror.HTTPStatusFromErrorOr(errRequest, http.StatusBadGateway), errRequest.Error()) return } if selection != nil && resp.StatusCode == http.StatusUnauthorized { @@ -319,7 +336,7 @@ func (h *Handler) Handle(c *gin.Context) { } if !didRefresh || refreshed == nil { selection.End("refresh_unavailable") - c.JSON(http.StatusUnauthorized, gin.H{"error": "Codex credential unauthorized"}) + writeLiveError(c, http.StatusUnauthorized, "Codex credential unauthorized") return } selected = refreshed @@ -328,7 +345,7 @@ func (h *Handler) Handle(c *gin.Context) { if errRequest != nil { selection.End("retry_failed") helps.RecordAPIResponseError(ctx, runtimeConfig, errRequest) - c.JSON(clienterror.HTTPStatusFromErrorOr(errRequest, http.StatusBadGateway), gin.H{"error": errRequest.Error()}) + writeLiveError(c, clienterror.HTTPStatusFromErrorOr(errRequest, http.StatusBadGateway), errRequest.Error()) return } if resp.StatusCode == http.StatusUnauthorized { @@ -351,7 +368,7 @@ func (h *Handler) Handle(c *gin.Context) { if selection != nil { if errBind := selection.Bind(closeResponseBody); errBind != nil { selection.End("response_bind_failed") - c.JSON(http.StatusServiceUnavailable, gin.H{"error": errBind.Error()}) + writeLiveError(c, http.StatusServiceUnavailable, errBind.Error()) return } } @@ -367,7 +384,7 @@ func (h *Handler) Handle(c *gin.Context) { message = "Codex live response body too large" status = http.StatusBadGateway } - c.JSON(status, gin.H{"error": message}) + writeLiveError(c, status, message) return } helps.AppendAPIResponseChunk(ctx, runtimeConfig, responseBody) @@ -377,22 +394,25 @@ func (h *Handler) Handle(c *gin.Context) { if success { callID = callIDFromLocation(resp.Header.Get("Location")) if callID == "" && mediaSession != nil { - c.JSON(http.StatusBadGateway, gin.H{"error": "Codex live response is missing a valid call ID"}) + writeLiveError(c, http.StatusBadGateway, "Codex live response is missing a valid call ID") return } if mediaSession != nil { mediaSession.SetCallID(callID) } + if callID != "" && strings.HasPrefix(c.Request.URL.Path, "/v1/realtime") { + responseHeaders.Set("Location", "/v1/realtime/calls/"+callID) + } } if success && mediaSession != nil { upstreamAnswer, errSDP := callResponseSDP(responseBody, resp.Header.Get("Content-Type")) if errSDP != nil { - c.JSON(http.StatusBadGateway, gin.H{"error": errSDP.Error()}) + writeLiveError(c, http.StatusBadGateway, errSDP.Error()) return } downstreamAnswer, errAnswer := mediaSession.AcceptUpstreamAnswer(ctx, upstreamAnswer) if errAnswer != nil { - c.JSON(clienterror.HTTPStatusFromErrorOr(errAnswer, http.StatusBadGateway), gin.H{"error": errAnswer.Error()}) + writeLiveError(c, clienterror.HTTPStatusFromErrorOr(errAnswer, http.StatusBadGateway), errAnswer.Error()) return } responseBodyToWrite = []byte(downstreamAnswer) @@ -403,13 +423,17 @@ func (h *Handler) Handle(c *gin.Context) { if success && h.sessions != nil { if callID != "" { session := liveSession{authID: selected.ID, model: model, media: mediaSession} + session.ownerPrincipal, session.ownerProvider = requestOwner(c) + if principal, ok := c.Get(ClientSecretPrincipalContextKey); ok { + session.clientSecretPrincipal, _ = principal.(string) + } if selection != nil { if mediaSession != nil { if errBind := selection.Bind(func() error { return mediaSession.CloseWithReason("home_selection_closed") }); errBind != nil { selection.End("media_bind_failed") - c.JSON(http.StatusServiceUnavailable, gin.H{"error": errBind.Error()}) + writeLiveError(c, http.StatusServiceUnavailable, errBind.Error()) return } } @@ -419,7 +443,7 @@ func (h *Handler) Handle(c *gin.Context) { return nil }); errBind != nil { selection.End("session_drain_bind_failed") - c.JSON(http.StatusServiceUnavailable, gin.H{"error": errBind.Error()}) + writeLiveError(c, http.StatusServiceUnavailable, errBind.Error()) return } selection.Retain() @@ -521,6 +545,78 @@ func prepareCallRequest(body []byte, contentType string) ([]byte, string, string return body, contentType, model, nil } +func applyClientSecretCallSession(body []byte, contentType, model string, session json.RawMessage) ([]byte, string, string, error) { + if len(session) == 0 { + return body, contentType, model, nil + } + mediaType, _, errMediaType := mime.ParseMediaType(contentType) + if errMediaType == nil && (strings.EqualFold(mediaType, "application/sdp") || strings.EqualFold(mediaType, "text/plain")) { + encoded, errEncode := encodeCallRequest(string(body), session) + if errEncode != nil { + return nil, "", "", errEncode + } + return encoded, "application/json", modelFromJSON(session), nil + } + if errMediaType != nil || !strings.EqualFold(mediaType, "application/json") { + return nil, "", "", errors.New("Realtime client secrets require an SDP or JSON call request") + } + var payload map[string]json.RawMessage + if errUnmarshal := json.Unmarshal(body, &payload); errUnmarshal != nil { + return nil, "", "", fmt.Errorf("failed to decode Realtime call request: %w", errUnmarshal) + } + payload["session"] = append(json.RawMessage(nil), session...) + encoded, errMarshal := json.Marshal(payload) + if errMarshal != nil { + return nil, "", "", fmt.Errorf("failed to encode Realtime call request: %w", errMarshal) + } + return encoded, "application/json", modelFromJSON(session), nil +} + +func rewriteCallRequestModel(body []byte, contentType, model string) ([]byte, string, error) { + upstreamModel := codexRealtimeModel(model) + mediaType, _, errMediaType := mime.ParseMediaType(contentType) + if errMediaType != nil || !strings.EqualFold(mediaType, "application/json") || len(bytes.TrimSpace(body)) == 0 { + return body, upstreamModel, nil + } + var payload map[string]json.RawMessage + if errUnmarshal := json.Unmarshal(body, &payload); errUnmarshal != nil { + return nil, "", fmt.Errorf("failed to decode Realtime call request: %w", errUnmarshal) + } + changed := false + if sessionJSON, ok := payload["session"]; ok && len(sessionJSON) > 0 { + var session map[string]json.RawMessage + if errUnmarshal := json.Unmarshal(sessionJSON, &session); errUnmarshal != nil { + return nil, "", fmt.Errorf("failed to decode Realtime session: %w", errUnmarshal) + } + encodedModel, errMarshal := json.Marshal(upstreamModel) + if errMarshal != nil { + return nil, "", fmt.Errorf("failed to encode Realtime model: %w", errMarshal) + } + session["model"] = encodedModel + encodedSession, errMarshal := json.Marshal(session) + if errMarshal != nil { + return nil, "", fmt.Errorf("failed to encode Realtime session: %w", errMarshal) + } + payload["session"] = encodedSession + changed = true + } else if _, ok := payload["model"]; ok { + encodedModel, errMarshal := json.Marshal(upstreamModel) + if errMarshal != nil { + return nil, "", fmt.Errorf("failed to encode Realtime model: %w", errMarshal) + } + payload["model"] = encodedModel + changed = true + } + if !changed { + return body, upstreamModel, nil + } + encoded, errMarshal := json.Marshal(payload) + if errMarshal != nil { + return nil, "", fmt.Errorf("failed to encode Realtime call request: %w", errMarshal) + } + return encoded, upstreamModel, nil +} + func multipartCallRequest(body []byte, boundary string) ([]byte, string, string, error) { if boundary == "" { return nil, "", "", errors.New("Codex live multipart boundary is missing") @@ -704,7 +800,7 @@ func headersForLogging(source http.Header) http.Header { func callResponseHeaders(source http.Header) http.Header { headers := make(http.Header) - for _, name := range []string{"Content-Type", "Location"} { + for _, name := range []string{"Content-Type", "Location", "Retry-After", "X-Request-Id", "OpenAI-Request-Id"} { for _, value := range source.Values(name) { headers.Add(name, value) } @@ -720,10 +816,25 @@ func writeResponseHeaders(destination, source http.Header) { } } +func writeLiveError(c *gin.Context, status int, message string) { + if c != nil && c.Request != nil && c.Request.URL != nil && strings.HasPrefix(c.Request.URL.Path, "/v1/realtime") { + errorType := "api_error" + if status >= http.StatusBadRequest && status < http.StatusInternalServerError { + errorType = "invalid_request_error" + } + if status == http.StatusUnauthorized { + errorType = "authentication_error" + } + writeRealtimeError(c, status, message, errorType, "realtime_request_failed") + return + } + c.JSON(status, gin.H{"error": message}) +} + func writeSelectionError(c *gin.Context, err error) { status := clienterror.HTTPStatusFromErrorOr(err, http.StatusServiceUnavailable) for _, value := range auth.SafeResponseHeaders(err).Values("Retry-After") { c.Writer.Header().Add("Retry-After", value) } - c.JSON(status, gin.H{"error": err.Error()}) + writeLiveError(c, status, err.Error()) } diff --git a/internal/client/codex/live/sideband.go b/internal/client/codex/live/sideband.go index cdbbbd5ba..28b00535d 100644 --- a/internal/client/codex/live/sideband.go +++ b/internal/client/codex/live/sideband.go @@ -42,13 +42,16 @@ var ( ) type liveSession struct { - callID string - authID string - model string - homeSelection *auth.HomeDispatchSelection - media mediaRelaySession - resources *liveSessionResources - token uint64 + callID string + authID string + model string + ownerPrincipal string + ownerProvider string + clientSecretPrincipal string + homeSelection *auth.HomeDispatchSelection + media mediaRelaySession + resources *liveSessionResources + token uint64 } type liveSessionResources struct { @@ -303,28 +306,41 @@ const ( // HandleSideband relays live session sideband WebSocket frames bidirectionally. func (h *Handler) HandleSideband(c *gin.Context) { if h == nil || h.authManager == nil || h.sessions == nil { - c.JSON(http.StatusServiceUnavailable, gin.H{"error": "Codex live sideband unavailable"}) + writeLiveError(c, http.StatusServiceUnavailable, "Codex live sideband unavailable") return } runtimeConfig := h.currentConfig() if !websocket.IsWebSocketUpgrade(c.Request) { - c.JSON(http.StatusUpgradeRequired, gin.H{"error": "WebSocket upgrade required"}) + c.Header("Upgrade", "websocket") + writeLiveError(c, http.StatusUpgradeRequired, "WebSocket upgrade required") return } style, callID, ok := sidebandTarget(c) if !ok { - c.JSON(http.StatusBadRequest, gin.H{"error": "Invalid Codex live call ID"}) + writeLiveError(c, http.StatusBadRequest, "Invalid Codex live call ID") return } session, claim := h.sessions.claim(callID) switch claim { case sessionClaimBusy: - c.JSON(http.StatusConflict, gin.H{"error": "Codex live session already joining"}) + writeLiveError(c, http.StatusConflict, "Codex live session already joining") return case sessionClaimAcquired: default: - c.JSON(http.StatusNotFound, gin.H{"error": "Codex live session not found"}) + writeLiveError(c, http.StatusNotFound, "Codex live session not found") + return + } + if principal, hasClientSecret := c.Get(ClientSecretPrincipalContextKey); hasClientSecret { + principalValue, _ := principal.(string) + if session.clientSecretPrincipal == "" || principalValue != session.clientSecretPrincipal { + h.sessions.release(session) + writeRealtimeError(c, http.StatusForbidden, "Realtime client secret is not valid for this call", "invalid_request_error", "realtime_client_secret_scope_mismatch") + return + } + } else if ownerPrincipal, ownerProvider := requestOwner(c); session.ownerPrincipal != "" && (ownerPrincipal != session.ownerPrincipal || ownerProvider != session.ownerProvider) { + h.sessions.release(session) + writeRealtimeError(c, http.StatusForbidden, "Realtime call belongs to another API principal", "invalid_request_error", "realtime_call_scope_mismatch") return } consumeSession := false @@ -337,20 +353,21 @@ func (h *Handler) HandleSideband(c *gin.Context) { }() ctx := context.WithValue(c.Request.Context(), "gin", c) + ctx = coreexecutor.WithDownstreamWebsocket(ctx) var selection *auth.HomeDispatchSelection var selected *auth.Auth var errSelect error if session.homeSelection != nil { if !session.homeSelection.Active() { consumeSession = true - c.JSON(http.StatusServiceUnavailable, gin.H{"error": "Codex live Home selection unavailable"}) + writeLiveError(c, http.StatusServiceUnavailable, "Codex live Home selection unavailable") return } selection = session.homeSelection selected = selection.CloneAuth() } else { selectionOpts := coreexecutor.Options{ - Headers: c.Request.Header.Clone(), + Headers: liveSelectionHeaders(c), Metadata: map[string]any{ coreexecutor.PinnedAuthMetadataKey: session.authID, coreexecutor.ExecutionSessionMetadataKey: callID, @@ -363,7 +380,7 @@ func (h *Handler) HandleSideband(c *gin.Context) { return } if selected == nil { - c.JSON(http.StatusServiceUnavailable, gin.H{"error": "Codex auth unavailable"}) + writeLiveError(c, http.StatusServiceUnavailable, "Codex auth unavailable") return } @@ -371,7 +388,7 @@ func (h *Handler) HandleSideband(c *gin.Context) { attemptCtx, releaseAttempt, errAttempt := selection.AttemptContext(ctx) if errAttempt != nil { consumeSession = true - c.JSON(http.StatusServiceUnavailable, gin.H{"error": errAttempt.Error()}) + writeLiveError(c, http.StatusServiceUnavailable, errAttempt.Error()) return } ctx = attemptCtx @@ -422,7 +439,7 @@ func (h *Handler) HandleSideband(c *gin.Context) { return } if !didRefresh || refreshed == nil { - c.JSON(http.StatusUnauthorized, gin.H{"error": "Codex credential unauthorized"}) + writeLiveError(c, http.StatusUnauthorized, "Codex credential unauthorized") return } selected = refreshed @@ -449,7 +466,7 @@ func (h *Handler) HandleSideband(c *gin.Context) { if selection != nil { if errBind := selection.Bind(closeUpstream); errBind != nil { consumeSession = true - c.JSON(http.StatusServiceUnavailable, gin.H{"error": errBind.Error()}) + writeLiveError(c, http.StatusServiceUnavailable, errBind.Error()) return } } else { @@ -556,6 +573,7 @@ func handleSidebandDialError(c *gin.Context, ctx context.Context, cfg *config.Co if response.StatusCode > 0 { status = response.StatusCode } + copyRealtimeHandshakeHeaders(c.Writer.Header(), response.Header) helps.RecordAPIWebsocketHandshake(ctx, cfg, response.StatusCode, callResponseHeaders(response.Header)) if response.Body != nil { if errClose := response.Body.Close(); errClose != nil { @@ -564,7 +582,7 @@ func handleSidebandDialError(c *gin.Context, ctx context.Context, cfg *config.Co } } helps.RecordAPIWebsocketError(ctx, cfg, "dial", errDial) - c.JSON(status, gin.H{"error": "Codex live sideband upstream unavailable"}) + writeLiveError(c, status, "Codex live sideband upstream unavailable") } func websocketCloseFunc(name string, conn *websocket.Conn) func() error { diff --git a/internal/client/codex/live/websocket.go b/internal/client/codex/live/websocket.go new file mode 100644 index 000000000..e0a147dd6 --- /dev/null +++ b/internal/client/codex/live/websocket.go @@ -0,0 +1,251 @@ +package live + +import ( + "context" + "encoding/json" + "net/http" + "net/url" + "strings" + + "github.com/gin-gonic/gin" + "github.com/gorilla/websocket" + "github.com/router-for-me/CLIProxyAPI/v7/internal/clienterror" + "github.com/router-for-me/CLIProxyAPI/v7/internal/logging" + "github.com/router-for-me/CLIProxyAPI/v7/internal/runtime/executor/helps" + "github.com/router-for-me/CLIProxyAPI/v7/sdk/cliproxy/auth" + coreexecutor "github.com/router-for-me/CLIProxyAPI/v7/sdk/cliproxy/executor" + log "github.com/sirupsen/logrus" +) + +const defaultStandardRealtimeModel = "gpt-realtime" + +// HandleRealtimeWebsocket dispatches a standard Realtime WebSocket or an existing call sideband. +func (h *Handler) HandleRealtimeWebsocket(c *gin.Context) { + if strings.TrimSpace(c.Query("call_id")) != "" { + h.HandleSideband(c) + return + } + h.HandleDirectWebsocket(c) +} + +// HandleDirectWebsocket relays a standard Realtime WebSocket through Codex OAuth. +func (h *Handler) HandleDirectWebsocket(c *gin.Context) { + if h == nil || h.authManager == nil { + writeRealtimeError(c, http.StatusServiceUnavailable, "Codex auth manager unavailable", "server_error", "codex_auth_unavailable") + return + } + if !websocket.IsWebSocketUpgrade(c.Request) { + c.Header("Upgrade", "websocket") + writeRealtimeError(c, http.StatusUpgradeRequired, "WebSocket upgrade required", "invalid_request_error", "websocket_upgrade_required") + return + } + + requestedModel := strings.TrimSpace(c.Query("model")) + if requestedModel == "" { + requestedModel = defaultStandardRealtimeModel + } + selectionModel := codexRealtimeModel(requestedModel) + tokenSession := clientSecretSession(c) + if len(tokenSession) > 0 { + tokenModel := codexRealtimeModel(modelFromJSON(tokenSession)) + if selectionModel != tokenModel { + writeRealtimeError(c, http.StatusForbidden, "Realtime client secret is not valid for the requested model", "invalid_request_error", "realtime_client_secret_scope_mismatch") + return + } + } + ctx := context.WithValue(c.Request.Context(), "gin", c) + ctx = coreexecutor.WithDownstreamWebsocket(ctx) + selectionOpts := coreexecutor.Options{Headers: liveSelectionHeaders(c)} + selection, selected, errSelect := h.selectOAuth(ctx, selectionModel, selectionOpts) + if errSelect != nil { + writeSelectionError(c, errSelect) + return + } + if selected == nil { + if selection != nil { + selection.End("missing_auth") + } + writeRealtimeError(c, http.StatusServiceUnavailable, "Codex auth unavailable", "server_error", "codex_auth_unavailable") + return + } + if selection != nil { + attemptCtx, releaseAttempt, errAttempt := selection.AttemptContext(ctx) + if errAttempt != nil { + selection.End("attempt_bind_failed") + writeRealtimeError(c, http.StatusServiceUnavailable, errAttempt.Error(), "server_error", "realtime_upstream_unavailable") + return + } + ctx = attemptCtx + defer releaseAttempt() + selection.Retain() + defer selection.End("session_closed") + } + logging.SetGinCPATraceID(c, selected.EnsureIndex()) + + upstreamURL := h.directRealtimeURL(requestedModel) + dialUpstream := func(current *auth.Auth) (*websocket.Conn, *http.Response, error) { + request, errRequest := http.NewRequestWithContext(ctx, http.MethodGet, websocketHTTPURL(upstreamURL), nil) + if errRequest != nil { + return nil, nil, errRequest + } + request.Header = directRealtimeHeaders(c.Request.Header) + setAccountHeader(request.Header, current) + if errPrepare := h.authManager.PrepareHttpRequest(ctx, current, request); errPrepare != nil { + return nil, nil, errPrepare + } + authType, authValue := current.AccountInfo() + helpersConfig := h.currentConfig() + helps.RecordAPIWebsocketRequest(ctx, helpersConfig, helps.UpstreamRequestLog{ + URL: upstreamURL, + Method: "WEBSOCKET", + Headers: headersForLogging(request.Header), + Provider: "codex", + AuthID: current.ID, + AuthLabel: current.Label, + AuthType: authType, + AuthValue: authValue, + }) + dialer := newProxyAwareSidebandDialer(helpersConfig, current) + dialer.Subprotocols = websocket.Subprotocols(c.Request) + return dialer.DialContext(ctx, upstreamURL, request.Header) + } + + upstream, handshakeResponse, errDial := dialUpstream(selected) + if errDial != nil && selection != nil && handshakeResponse != nil && handshakeResponse.StatusCode == http.StatusUnauthorized { + h.authManager.ReportHomeUnauthorized(ctx, selected, "codex", selectionModel) + closeHandshakeBody(handshakeResponse, "direct websocket unauthorized") + refreshed, didRefresh, errRefresh := h.authManager.RefreshHomeSelectionAfterUnauthorized(ctx, selection, selected) + if errRefresh != nil { + writeSelectionError(c, errRefresh) + return + } + if didRefresh && refreshed != nil { + selected = refreshed + logging.SetGinCPATraceID(c, selected.EnsureIndex()) + upstream, handshakeResponse, errDial = dialUpstream(selected) + } + } + if errDial != nil { + status := clienterror.HTTPStatusFromErrorOr(errDial, http.StatusBadGateway) + if handshakeResponse != nil && handshakeResponse.StatusCode > 0 { + status = handshakeResponse.StatusCode + copyRealtimeHandshakeHeaders(c.Writer.Header(), handshakeResponse.Header) + } + closeHandshakeBody(handshakeResponse, "direct websocket rejected") + helpConfig := h.currentConfig() + helpDetails := "Codex Realtime WebSocket upstream unavailable" + helpType := "api_error" + if status == http.StatusNotFound || status == http.StatusNotImplemented { + helpDetails = "Direct Realtime WebSocket is not supported by the Codex OAuth upstream" + helpType = "not_supported_error" + status = http.StatusNotImplemented + } + helpCode := "realtime_websocket_upstream_unavailable" + if helpType == "not_supported_error" { + helpCode = "realtime_capability_not_supported" + } else if status == http.StatusUnauthorized { + helpType = "authentication_error" + helpCode = "realtime_upstream_unauthorized" + } + helps.RecordAPIWebsocketError(ctx, helpConfig, "dial", errDial) + writeRealtimeError(c, status, helpDetails, helpType, helpCode) + return + } + closeHandshakeBody(handshakeResponse, "direct websocket handshake") + closeUpstream := websocketCloseFunc("upstream", upstream) + defer func() { _ = closeUpstream() }() + if len(tokenSession) > 0 { + updateSession, errSession := realtimeSessionUpdate(tokenSession) + if errSession != nil { + _ = closeUpstream() + writeRealtimeError(c, http.StatusInternalServerError, "Failed to apply Realtime client secret session", "server_error", "realtime_session_failed") + return + } + update, errMarshal := json.Marshal(struct { + Type string `json:"type"` + Session json.RawMessage `json:"session"` + }{Type: "session.update", Session: updateSession}) + if errMarshal != nil { + _ = closeUpstream() + writeRealtimeError(c, http.StatusInternalServerError, "Failed to apply Realtime client secret session", "server_error", "realtime_session_failed") + return + } + if errWrite := upstream.WriteMessage(websocket.TextMessage, update); errWrite != nil { + _ = closeUpstream() + writeRealtimeError(c, http.StatusBadGateway, "Failed to apply Realtime client secret session", "api_error", "realtime_upstream_unavailable") + return + } + } + + if selection != nil { + if errBind := selection.Bind(closeUpstream); errBind != nil { + writeRealtimeError(c, http.StatusServiceUnavailable, errBind.Error(), "server_error", "realtime_upstream_unavailable") + return + } + } + + upgradeHeaders := make(http.Header) + if subprotocol := upstream.Subprotocol(); subprotocol != "" { + upgradeHeaders.Set("Sec-WebSocket-Protocol", subprotocol) + } + downstream, errUpgrade := sidebandUpgrader.Upgrade(c.Writer, c.Request, upgradeHeaders) + if errUpgrade != nil { + _ = closeUpstream() + return + } + closeDownstream := websocketCloseFunc("downstream", downstream) + defer func() { _ = closeDownstream() }() + if selection != nil { + if errBind := selection.Bind(closeDownstream); errBind != nil { + return + } + } + + if errRelay := relayWebsockets(downstream, upstream); errRelay != nil && !isNormalWebsocketClose(errRelay) { + helps.RecordAPIWebsocketError(ctx, h.currentConfig(), "relay", errRelay) + log.WithError(errRelay).Debug("codex realtime direct websocket relay closed") + } +} + +func realtimeSessionUpdate(session json.RawMessage) (json.RawMessage, error) { + var update map[string]json.RawMessage + if errUnmarshal := json.Unmarshal(session, &update); errUnmarshal != nil { + return nil, errUnmarshal + } + for _, field := range []string{"model", "id", "object", "expires_at", "client_secret"} { + delete(update, field) + } + return json.Marshal(update) +} + +func (h *Handler) directRealtimeURL(model string) string { + values := make(url.Values) + values.Set("model", strings.TrimSpace(model)) + return strings.TrimRight(h.sidebandAPIBaseURL, "/") + "/realtime?" + values.Encode() +} + +func directRealtimeHeaders(source http.Header) http.Header { + headers := protocolHeaders(source) + headers.Del("OpenAI-Alpha") + if headers.Get("Originator") == "" { + headers.Set("Originator", "Codex Desktop") + } + return headers +} + +func copyRealtimeHandshakeHeaders(destination, source http.Header) { + for _, name := range []string{"Retry-After", "X-Request-Id", "OpenAI-Request-Id"} { + for _, value := range source.Values(name) { + destination.Add(name, value) + } + } +} + +func closeHandshakeBody(response *http.Response, label string) { + if response == nil || response.Body == nil { + return + } + if errClose := response.Body.Close(); errClose != nil { + log.Errorf("codex realtime: close %s response body error: %v", label, errClose) + } +} diff --git a/internal/client/codex/live/websocket_test.go b/internal/client/codex/live/websocket_test.go new file mode 100644 index 000000000..42eb78a77 --- /dev/null +++ b/internal/client/codex/live/websocket_test.go @@ -0,0 +1,203 @@ +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") + } +}