mirror of
https://github.com/router-for-me/CLIProxyAPI.git
synced 2026-10-10 17:50:44 +08:00
feat(codex): add realtime hangup forwarding and local client-secret support
- Add Codex live handlers for unsupported translation/transcription/SIP endpoints, returning standardized `realtime_capability_not_supported` errors. - Implement `HandleHangup` to validate call ownership, select/refresh pinned OAuth credentials, forward hangup requests to upstream, and complete local session on success. - Add ephemeral client-secret infrastructure for local ephemeral auth: create/authenticate endpoints, token storage with expiry/capacity limits, session normalization/model mapping, and unified realtime error handling. Closes: #4726
This commit is contained in:
89
examples/realtime-openai-go/README.md
Normal file
89
examples/realtime-openai-go/README.md
Normal file
@@ -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.
|
||||
15
examples/realtime-openai-go/go.mod
Normal file
15
examples/realtime-openai-go/go.mod
Normal file
@@ -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
|
||||
)
|
||||
14
examples/realtime-openai-go/go.sum
Normal file
14
examples/realtime-openai-go/go.sum
Normal file
@@ -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=
|
||||
341
examples/realtime-openai-go/main.go
Normal file
341
examples/realtime-openai-go/main.go
Normal file
@@ -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
|
||||
}
|
||||
226
examples/realtime-openai-go/main_test.go
Normal file
226
examples/realtime-openai-go/main_test.go
Normal file
@@ -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)
|
||||
}
|
||||
}
|
||||
147
examples/realtime-openai-go/wav.go
Normal file
147
examples/realtime-openai-go/wav.go
Normal file
@@ -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
|
||||
}
|
||||
@@ -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()
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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))
|
||||
{
|
||||
|
||||
@@ -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{}
|
||||
|
||||
215
internal/client/codex/live/capabilities.go
Normal file
215
internal/client/codex/live/capabilities.go
Normal file
@@ -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")
|
||||
}
|
||||
97
internal/client/codex/live/capabilities_test.go
Normal file
97
internal/client/codex/live/capabilities_test.go
Normal file
@@ -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())
|
||||
}
|
||||
}
|
||||
}
|
||||
419
internal/client/codex/live/client_secret.go
Normal file
419
internal/client/codex/live/client_secret.go
Normal file
@@ -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,
|
||||
}})
|
||||
}
|
||||
262
internal/client/codex/live/client_secret_test.go
Normal file
262
internal/client/codex/live/client_secret_test.go
Normal file
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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())
|
||||
}
|
||||
|
||||
@@ -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 {
|
||||
|
||||
251
internal/client/codex/live/websocket.go
Normal file
251
internal/client/codex/live/websocket.go
Normal file
@@ -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)
|
||||
}
|
||||
}
|
||||
203
internal/client/codex/live/websocket_test.go
Normal file
203
internal/client/codex/live/websocket_test.go
Normal file
@@ -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")
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user