feat(live): relay realtime WebRTC media

This commit is contained in:
Luis Pater
2026-07-25 05:11:28 +08:00
parent 46172dd452
commit bda79b21bb
12 changed files with 1849 additions and 32 deletions

View File

@@ -217,6 +217,28 @@ codex:
# normalizes encrypted agent_message content for Codex, and converts agent_message input
# into standard user messages for non-Codex upstream protocols.
optimize-multi-agent-v2: false
# Terminate and relay Codex Live WebRTC audio and DataChannel traffic in this process.
# This requires inbound UDP reachability. Keep disabled to preserve direct media behavior.
live-media-relay:
enabled: false
# Maximum concurrent media sessions. Zero uses the default of 32.
max-sessions: 32
# Allow downstream SDP candidates that target private, loopback, link-local, or unspecified IPs.
# Enable only when Codex Desktop reaches CPA over a trusted local network.
allow-private-remote-ips: false
# Public IPv4 or IPv6 address advertised when CPA is behind 1:1 NAT.
public-ip: ""
# Optional UDP allocation range. Both values must be set together and provide at least two ports per session.
udp-port-min: 0
udp-port-max: 0
# Optional STUN/TURN servers. TURN credentials are never returned by the JSON config API.
# ice-servers:
# - urls:
# - "stun:stun.example.com:3478"
# - urls:
# - "turn:turn.example.com:3478?transport=udp"
# username: "user"
# credential: "secret"
# When true, enable authentication for the WebSocket API (/v1/ws).
ws-auth: true

18
go.mod
View File

@@ -17,6 +17,9 @@ require (
github.com/joho/godotenv v1.5.1
github.com/klauspost/compress v1.17.4
github.com/minio/minio-go/v7 v7.0.66
github.com/pion/interceptor v0.1.45
github.com/pion/rtp v1.10.4
github.com/pion/webrtc/v4 v4.2.17
github.com/redis/go-redis/v9 v9.19.0
github.com/refraction-networking/utls v1.8.2
github.com/sirupsen/logrus v1.9.3
@@ -36,8 +39,23 @@ require (
require (
github.com/cespare/xxhash/v2 v2.3.0 // indirect
github.com/dlclark/regexp2/v2 v2.5.1 // indirect
github.com/pion/datachannel v1.6.2 // indirect
github.com/pion/dtls/v3 v3.1.5 // indirect
github.com/pion/ice/v4 v4.3.0 // indirect
github.com/pion/logging v0.2.4 // indirect
github.com/pion/mdns/v2 v2.1.0 // indirect
github.com/pion/randutil v0.1.0 // indirect
github.com/pion/rtcp v1.2.17 // indirect
github.com/pion/sctp v1.11.0 // indirect
github.com/pion/sdp/v3 v3.0.19 // indirect
github.com/pion/srtp/v3 v3.0.12 // indirect
github.com/pion/stun/v3 v3.1.6 // indirect
github.com/pion/transport/v4 v4.0.2 // indirect
github.com/pion/turn/v5 v5.0.12 // indirect
github.com/rogpeppe/go-internal v1.15.0 // indirect
github.com/wlynxg/anet v0.0.5 // indirect
go.uber.org/atomic v1.11.0 // indirect
golang.org/x/time v0.14.0 // indirect
)
require (

41
go.sum
View File

@@ -121,8 +121,9 @@ github.com/kr/pretty v0.1.0/go.mod h1:dAy3ld7l9f0ibDNOQOHHMYYIIbhfbHSm3C4ZsoJORN
github.com/kr/pretty v0.3.0 h1:WgNl7dwNpEZ6jJ9k1snq4pZsg7DOEN8hP9Xw0Tsjwk0=
github.com/kr/pretty v0.3.0/go.mod h1:640gp4NfQd8pI5XOwp5fnNeVWj67G7CFk/SaSQn7NBk=
github.com/kr/pty v1.1.1/go.mod h1:pFQYn66WHrOpPYNljwOMqo10TkYh1fy3cYio2l3bCsQ=
github.com/kr/text v0.1.0 h1:45sCR5RtlFHMR4UwH9sdQ5TC8v0qDQCHnXt+kaKSTVE=
github.com/kr/text v0.1.0/go.mod h1:4Jbv+DJW3UT/LiOwJeYQe1efqtUx/iVham/4vfdArNI=
github.com/kr/text v0.2.0 h1:5Nx0Ya0ZqY2ygV366QzturHI13Jq95ApcVaJBhpS+AY=
github.com/kr/text v0.2.0/go.mod h1:eLer722TekiGuMkidMxC/pM04lWEeraHUUmBw8l2grE=
github.com/leodido/go-urn v1.4.0 h1:WT9HwE9SGECu3lg4d/dIA+jxlljEa1/ffXKmRjqdmIQ=
github.com/leodido/go-urn v1.4.0/go.mod h1:bvxc+MVxLKB4z00jd1z+Dvzr47oO32F/QSNjSBOlFxI=
github.com/lucasb-eyer/go-colorful v1.3.0 h1:2/yBRLdWBZKrf7gB40FoiKfAWYQ0lqNcbuQwVHXptag=
@@ -154,6 +155,40 @@ github.com/pelletier/go-toml/v2 v2.2.2 h1:aYUidT7k73Pcl9nb2gScu7NSrKCSHIDE89b3+6
github.com/pelletier/go-toml/v2 v2.2.2/go.mod h1:1t835xjRzz80PqgE6HHgN2JOsmgYu/h4qDAS4n929Rs=
github.com/pierrec/xxHash v0.1.5 h1:n/jBpwTHiER4xYvK3/CdPVnLDPchj8eTJFFLUb4QHBo=
github.com/pierrec/xxHash v0.1.5/go.mod h1:w2waW5Zoa/Wc4Yqe0wgrIYAGKqRMf7czn2HNKXmuL+I=
github.com/pion/datachannel v1.6.2 h1:7EXQ8TH3vTouBUdRWYbcX2edSx9Yj6k5zl5P+qyxEPc=
github.com/pion/datachannel v1.6.2/go.mod h1:pzbdAZvyGtXbcHM1hBbsFaOTf40lZizU/dNlvVOak6E=
github.com/pion/dtls/v3 v3.1.5 h1:9xJtVsHwMYeSjPp5Hh1FTis4DchnQWtnOa5o+6ygqfc=
github.com/pion/dtls/v3 v3.1.5/go.mod h1:gz1K4jg6c+fq86oQMH4pilpCEOEPwmEr2jY+VcF/mkU=
github.com/pion/ice/v4 v4.3.0 h1:X8l4s9zV2HeTKX33nulWAFXAEo5KhIVzOsY62/3t/LM=
github.com/pion/ice/v4 v4.3.0/go.mod h1:obAyD+J+Hzs7QA7Y8YXHp5uIn6gb7z87pKedXZkrcFU=
github.com/pion/interceptor v0.1.45 h1:6PUo/5829bIfRFIPPJQzuDn8EjxRTSB/CSD7QVCOaqo=
github.com/pion/interceptor v0.1.45/go.mod h1:gNDYM/uFKcLe/B3gS2/7+aw6z+RDiMy2qKTnF1LO31w=
github.com/pion/logging v0.2.4 h1:tTew+7cmQ+Mc1pTBLKH2puKsOvhm32dROumOZ655zB8=
github.com/pion/logging v0.2.4/go.mod h1:DffhXTKYdNZU+KtJ5pyQDjvOAh/GsNSyv1lbkFbe3so=
github.com/pion/mdns/v2 v2.1.0 h1:3IJ9+Xio6tWYjhN6WwuY142P/1jA0D5ERaIqawg/fOY=
github.com/pion/mdns/v2 v2.1.0/go.mod h1:pcez23GdynwcfRU1977qKU0mDxSeucttSHbCSfFOd9A=
github.com/pion/randutil v0.1.0 h1:CFG1UdESneORglEsnimhUjf33Rwjubwj6xfiOXBa3mA=
github.com/pion/randutil v0.1.0/go.mod h1:XcJrSMMbbMRhASFVOlj/5hQial/Y8oH/HVo7TBZq+j8=
github.com/pion/rtcp v1.2.17 h1:PxiT6L79yPZKtXIsXdG1eakBl6dtBj4x+4oVEL0DlSw=
github.com/pion/rtcp v1.2.17/go.mod h1:7kBpuBJaWwax4hzc/pgexY8vkOpvh8atgYDbaKZq0iU=
github.com/pion/rtp v1.10.4 h1:4sCUwUd35Nllcpyp8V7lRgb4DV/ulHJaRTjbrkAcpQ4=
github.com/pion/rtp v1.10.4/go.mod h1:Au8fc6cEByy8RLTwKTQTEeQqDB/SJDxwL4mZuxYA5Pk=
github.com/pion/sctp v1.11.0 h1:sAxv9Qp3uIcaF5wu1XntwshtnW93CEuxhpkYzSbnfMs=
github.com/pion/sctp v1.11.0/go.mod h1:7KFmTwLcoYgJs/Z+99nJvsWL0qDpuyloSI0RbAqlrz0=
github.com/pion/sdp/v3 v3.0.19 h1:1VMKs3gIkTQV5M3hNKfTAPrDXSNrYtOlmOD8+mSZUGQ=
github.com/pion/sdp/v3 v3.0.19/go.mod h1:dE5WOSlzXrtiE/iuZqe9n+AcEbOjtAd3k5m5NtlV/qU=
github.com/pion/srtp/v3 v3.0.12 h1:U7V17bckl7sI4mb3sepiojByDuBY0wNCqQE+6IlQBbc=
github.com/pion/srtp/v3 v3.0.12/go.mod h1:EeZOi/sd6glM1EXapg051gdNWO9yWT1YSsgQ4SlJkns=
github.com/pion/stun/v3 v3.1.6 h1:WnhsD0eHCiwCfKNkVx0VJJwr2Y3eV4Ueih3KJ+dfZy8=
github.com/pion/stun/v3 v3.1.6/go.mod h1:zRUghXSQU32Lx5orJsz3uYMkIihweXb3mu5gIns02fs=
github.com/pion/transport/v3 v3.1.1 h1:Tr684+fnnKlhPceU+ICdrw6KKkTms+5qHMgw6bIkYOM=
github.com/pion/transport/v3 v3.1.1/go.mod h1:+c2eewC5WJQHiAA46fkMMzoYZSuGzA/7E2FPrOYHctQ=
github.com/pion/transport/v4 v4.0.2 h1:ifYlPqNwsy6aKQ9y8yzxXlHae5431ZrH2avkD/Rn6Tk=
github.com/pion/transport/v4 v4.0.2/go.mod h1:06hFI+jCFcok2X2MekVufNZ/uzNZXivGBPfviSVcjgM=
github.com/pion/turn/v5 v5.0.12 h1:6+b69ivQQXSlyfkp2AKripqD2k3W32qXK8QzCzpJWPI=
github.com/pion/turn/v5 v5.0.12/go.mod h1:CQACsRDJtjQ+6RSrGHrS2PCIerLwbW3uqXRqOvtjAFg=
github.com/pion/webrtc/v4 v4.2.17 h1:no7rmszKV1jkGz7GvErGp/VlnzGu/koVHO9CRjItiVU=
github.com/pion/webrtc/v4 v4.2.17/go.mod h1:xRtWZDJ0FbyW98WVCCgOvxaBM5gxqqJa7pCc4f+x/LI=
github.com/pjbgf/sha1cd v0.6.0 h1:3WJ8Wz8gvDz29quX1OcEmkAlUg9diU4GxJHqs0/XiwU=
github.com/pjbgf/sha1cd v0.6.0/go.mod h1:lhpGlyHLpQZoxMv8HcgXvZEhcGs0PG/vsZnEJ7H0iCM=
github.com/pmezard/go-difflib v1.0.0 h1:4DBwDE0NGyQoBHbLQYPwSUPoCMWR5BEzIk/f1lZbAQM=
@@ -203,6 +238,8 @@ github.com/twitchyliquid64/golang-asm v0.15.1 h1:SU5vSMR7hnwNxj24w34ZyCi/FmDZTkS
github.com/twitchyliquid64/golang-asm v0.15.1/go.mod h1:a1lVb/DtPvCB8fslRZhAngC2+aY1QWCk3Cedj/Gdt08=
github.com/ugorji/go/codec v1.2.12 h1:9LC83zGrHhuUA9l16C9AHXAqEV/2wBQ4nkvumAE65EE=
github.com/ugorji/go/codec v1.2.12/go.mod h1:UNopzCgEMSXjBc6AOMqYvWC1ktqTAfzJZUZgYf6w6lg=
github.com/wlynxg/anet v0.0.5 h1:J3VJGi1gvo0JwZ/P1/Yc/8p63SoW98B5dHkYDmpgvvU=
github.com/wlynxg/anet v0.0.5/go.mod h1:eay5PRQr7fIVAMbTbchTnO9gG65Hg/uYGdc7mguHxoA=
github.com/xo/terminfo v0.0.0-20220910002029-abceb7e1c41e h1:JVG44RsyaB9T2KIHavMF/ppJZNG9ZpyihvCd0w101no=
github.com/xo/terminfo v0.0.0-20220910002029-abceb7e1c41e/go.mod h1:RbqR21r5mrJuqunuUZ/Dhy/avygyECGrLceyNeo4LiM=
github.com/zeebo/xxh3 v1.1.0 h1:s7DLGDK45Dyfg7++yxI0khrfwq9661w9EN78eP/UZVs=
@@ -231,6 +268,8 @@ golang.org/x/term v0.45.0 h1:NwWyBmoJCbfTHpxrWoZ9C6/VxOf7ic219I8xZZFdrf0=
golang.org/x/term v0.45.0/go.mod h1:9aqxs0blBcrm/n0L9QW0aRVD+ktan8ssZromtqJC43w=
golang.org/x/text v0.40.0 h1:Ub2Z6/xjgF1WrYQz2nuITOEegKFtiIy+rieRJ5lHZKs=
golang.org/x/text v0.40.0/go.mod h1:hpnzDAfGV753zIKo+wk3u1bVKCGPbrnF7+7LBF/UHVY=
golang.org/x/time v0.14.0 h1:MRx4UaLrDotUKUdCIqzPC48t1Y9hANFKIRpNx+Te8PI=
golang.org/x/time v0.14.0/go.mod h1:eL/Oa2bBBK0TkX57Fyni+NgnyQQN4LitPmob2Hjnqw4=
golang.org/x/xerrors v0.0.0-20191204190536-9bdfabe68543 h1:E7g+9GITq07hpfrRu66IVDexMakfv52eLZ2CXBWiKr4=
golang.org/x/xerrors v0.0.0-20191204190536-9bdfabe68543/go.mod h1:I/5z698sn9Ka8TeJc9MKroUUfqBBauWjQqLJ2OPfmY0=
google.golang.org/protobuf v1.34.1 h1:9ddQBjfCyZPOHPUiPxpYESBLc+T8P3E+Vo4IbKZgFWg=

View File

@@ -215,7 +215,8 @@ type Server struct {
muxHTTPListener *muxListener
// handlers contains the API handlers for processing requests.
handlers *handlers.BaseAPIHandler
handlers *handlers.BaseAPIHandler
codexLiveHandler *codexlive.Handler
// cfg holds the current server configuration.
cfg *config.Config
@@ -524,7 +525,7 @@ func (s *Server) setupRoutes() {
geminiHandlers := gemini.NewGeminiAPIHandler(s.handlers)
claudeCodeHandlers := claude.NewClaudeCodeAPIHandler(s.handlers)
openaiResponsesHandlers := openai.NewOpenAIResponsesAPIHandler(s.handlers)
codexLiveHandler := codexlive.NewHandler(s.handlers.AuthManager, s.cfg)
s.codexLiveHandler = codexlive.NewHandler(s.handlers.AuthManager, s.cfg)
// OpenAI compatible API routes
v1 := s.engine.Group("/v1")
@@ -546,11 +547,11 @@ func (s *Server) setupRoutes() {
v1.POST("/responses", openaiResponsesHandlers.Responses)
v1.POST("/responses/compact", openaiResponsesHandlers.Compact)
v1.POST("/alpha/search", s.codexAlphaSearch)
v1.POST("/live", codexLiveHandler.Handle)
v1.GET("/live/:call_id", codexLiveHandler.HandleSideband)
v1.POST("/realtime/calls", codexLiveHandler.Handle)
v1.GET("/realtime/calls/:call_id", codexLiveHandler.HandleSideband)
v1.GET("/realtime", codexLiveHandler.HandleSideband)
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)
}
openaiV1 := s.engine.Group("/openai/v1")
@@ -1875,8 +1876,12 @@ func (s *Server) Stop(ctx context.Context) error {
}
// Shutdown the HTTP server.
if err := s.server.Shutdown(ctx); err != nil {
return fmt.Errorf("failed to shutdown HTTP server: %v", err)
errShutdown := s.server.Shutdown(ctx)
if s.codexLiveHandler != nil {
s.codexLiveHandler.Close()
}
if errShutdown != nil {
return fmt.Errorf("failed to shutdown HTTP server: %v", errShutdown)
}
log.Debug("API server stopped")
@@ -2048,6 +2053,11 @@ func (s *Server) UpdateClientsContext(ctx context.Context, cfg *config.Config) b
s.exampleAPIKeySafeModeActive.Store(exampleAPIKeySafeModeRequired)
}
s.cfg = cfg
if s.codexLiveHandler != nil {
if errUpdate := s.codexLiveHandler.UpdateConfig(cfg); errUpdate != nil {
log.WithError(errUpdate).Error("failed to update Codex Live media relay configuration")
}
}
s.wsAuthEnabled.Store(cfg.WebsocketAuth)
if oldCfg != nil && s.wsAuthChanged != nil && oldCfg.WebsocketAuth != cfg.WebsocketAuth {
s.wsAuthChanged(oldCfg.WebsocketAuth, cfg.WebsocketAuth)

View File

@@ -44,16 +44,56 @@ type Handler struct {
cfg *config.Config
sessions *sessionStore
sidebandAPIBaseURL string
mediaRelayMu sync.RWMutex
mediaRelay mediaRelayFactory
mediaRelayErr error
}
// NewHandler creates a Codex live session handler.
func NewHandler(authManager *auth.Manager, cfg *config.Config) *Handler {
return &Handler{
handler := &Handler{
authManager: authManager,
cfg: cfg,
sessions: newSessionStore(),
sidebandAPIBaseURL: defaultSidebandAPIBaseURL,
}
_ = handler.UpdateConfig(cfg)
return handler
}
// UpdateConfig atomically applies Codex Live media relay settings to new sessions.
func (h *Handler) UpdateConfig(cfg *config.Config) error {
if h == nil {
return nil
}
var relay mediaRelayFactory
var relayErr error
if cfg != nil && cfg.Codex.LiveMediaRelay.Enabled {
relay, relayErr = newPionMediaRelay(cfg.Codex.LiveMediaRelay)
}
h.mediaRelayMu.Lock()
h.mediaRelay = relay
h.mediaRelayErr = relayErr
h.mediaRelayMu.Unlock()
return relayErr
}
func (h *Handler) currentMediaRelay() (mediaRelayFactory, error) {
if h == nil {
return nil, nil
}
h.mediaRelayMu.RLock()
relay := h.mediaRelay
relayErr := h.mediaRelayErr
h.mediaRelayMu.RUnlock()
return relay, relayErr
}
// Close releases all active Codex live sessions.
func (h *Handler) Close() {
if h != nil && h.sessions != nil {
h.sessions.closeAll("server_stopped")
}
}
// Handle forwards a WebRTC SDP bootstrap request to the Codex realtime calls endpoint.
@@ -77,6 +117,13 @@ func (h *Handler) Handle(c *gin.Context) {
c.JSON(http.StatusBadRequest, gin.H{"error": errPayload.Error()})
return
}
mediaRelay, mediaRelayErr := h.currentMediaRelay()
if mediaRelayErr != nil {
c.JSON(http.StatusServiceUnavailable, gin.H{"error": mediaRelayErr.Error()})
return
}
var mediaSession mediaRelaySession
mediaRetained := false
ctx := context.WithValue(c.Request.Context(), "gin", c)
selectionOpts := coreexecutor.Options{
@@ -107,6 +154,39 @@ func (h *Handler) Handle(c *gin.Context) {
defer releaseAttempt()
}
logging.SetGinCPATraceID(c, selected.EnsureIndex())
if selection != nil {
defer func() {
if selection.Active() && !selection.Retained() {
selection.End("request_closed")
}
}()
}
if mediaRelay != nil {
clientOffer, errSDP := callRequestSDP(upstreamBody, upstreamContentType)
if errSDP != nil {
c.JSON(http.StatusBadRequest, gin.H{"error": errSDP.Error()})
return
}
var upstreamOffer string
mediaSession, upstreamOffer, errSDP = mediaRelay.NewSession(ctx, clientOffer)
if errSDP != nil {
c.JSON(http.StatusBadGateway, gin.H{"error": errSDP.Error()})
return
}
defer func() {
if !mediaRetained {
if errClose := mediaSession.Close(); errClose != nil {
log.WithError(errClose).Debug("codex live media: close unretained session")
}
}
}()
upstreamBody, upstreamContentType, errSDP = replaceCallRequestSDP(upstreamBody, upstreamContentType, upstreamOffer)
if errSDP != nil {
c.JSON(http.StatusBadRequest, gin.H{"error": errSDP.Error()})
return
}
}
headers := protocolHeaders(c.Request.Header)
headers.Set("Content-Type", upstreamContentType)
@@ -168,11 +248,6 @@ func (h *Handler) Handle(c *gin.Context) {
c.JSON(http.StatusServiceUnavailable, gin.H{"error": errBind.Error()})
return
}
defer func() {
if !selection.Retained() {
selection.End("response_closed")
}
}()
}
responseHeaders := callResponseHeaders(resp.Header)
@@ -188,10 +263,40 @@ func (h *Handler) Handle(c *gin.Context) {
return
}
helps.AppendAPIResponseChunk(ctx, h.cfg, responseBody)
if resp.StatusCode >= http.StatusOK && resp.StatusCode < http.StatusMultipleChoices && h.sessions != nil {
if callID := callIDFromLocation(resp.Header.Get("Location")); callID != "" {
session := liveSession{authID: selected.ID, model: model}
responseBodyToWrite := responseBody
success := resp.StatusCode >= http.StatusOK && resp.StatusCode < http.StatusMultipleChoices
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()})
return
}
downstreamAnswer, errAnswer := mediaSession.AcceptUpstreamAnswer(ctx, upstreamAnswer)
if errAnswer != nil {
c.JSON(http.StatusBadGateway, gin.H{"error": errAnswer.Error()})
return
}
responseBodyToWrite = []byte(downstreamAnswer)
responseHeaders.Set("Content-Type", "application/sdp")
}
var storedSession liveSession
sessionStored := false
if success && h.sessions != nil {
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"})
return
}
if callID != "" {
session := liveSession{authID: selected.ID, model: model, media: mediaSession}
if selection != nil {
if mediaSession != nil {
if errBind := selection.Bind(mediaSession.Close); errBind != nil {
selection.End("media_bind_failed")
c.JSON(http.StatusServiceUnavailable, gin.H{"error": errBind.Error()})
return
}
}
if errBind := selection.Bind(func() error {
// End outside the resource closer to avoid waiting on the closer itself.
go selection.End("session_drained")
@@ -204,12 +309,22 @@ func (h *Handler) Handle(c *gin.Context) {
selection.Retain()
session.homeSelection = selection
}
h.sessions.put(callID, session)
storedSession = h.sessions.put(callID, session)
sessionStored = storedSession.callID != ""
if mediaSession != nil {
mediaSession.SetCloseHandler(func(reason string) {
h.sessions.complete(storedSession, reason)
})
mediaRetained = true
}
}
}
writeResponseHeaders(c.Writer.Header(), responseHeaders)
c.Status(resp.StatusCode)
if _, errWrite := c.Writer.Write(responseBody); errWrite != nil {
if _, errWrite := c.Writer.Write(responseBodyToWrite); errWrite != nil {
if sessionStored {
h.sessions.complete(storedSession, "response_write_failed")
}
helps.RecordAPIResponseError(ctx, h.cfg, errWrite)
log.WithError(errWrite).Warn("codex live: write response body failed")
}
@@ -265,7 +380,6 @@ func prepareCallRequest(body []byte, contentType string) ([]byte, string, string
if errMediaType == nil && strings.EqualFold(mediaType, "multipart/form-data") {
return multipartCallRequest(body, strings.TrimSpace(params["boundary"]))
}
model := modelFromJSON(body)
if model == "" {
model = defaultLiveModel
@@ -321,18 +435,97 @@ func multipartCallRequest(body []byte, boundary string) ([]byte, string, string,
model = defaultLiveModel
}
encoded, errEncode := encodeCallRequest(*sdp, session)
if errEncode != nil {
return nil, "", "", errEncode
}
return encoded, "application/json", model, nil
}
func encodeCallRequest(sdp string, session json.RawMessage) ([]byte, error) {
payload := struct {
SDP string `json:"sdp"`
Session json.RawMessage `json:"session,omitempty"`
}{
SDP: *sdp,
SDP: sdp,
Session: session,
}
encoded, errMarshal := json.Marshal(payload)
if errMarshal != nil {
return nil, "", "", fmt.Errorf("failed to encode Codex live request: %w", errMarshal)
return nil, fmt.Errorf("failed to encode Codex live request: %w", errMarshal)
}
return encoded, "application/json", model, nil
return encoded, nil
}
func callRequestSDP(body []byte, contentType string) (string, error) {
mediaType, _, errMediaType := mime.ParseMediaType(contentType)
if errMediaType == nil && (strings.EqualFold(mediaType, "application/sdp") || strings.EqualFold(mediaType, "text/plain")) {
if strings.TrimSpace(string(body)) == "" {
return "", errors.New("Codex live call request requires an SDP offer")
}
return string(body), nil
}
if errMediaType != nil || !strings.EqualFold(mediaType, "application/json") {
return "", errors.New("Codex live media relay requires an SDP or JSON call request")
}
var payload struct {
SDP string `json:"sdp"`
}
if errUnmarshal := json.Unmarshal(body, &payload); errUnmarshal != nil {
return "", fmt.Errorf("failed to decode Codex live call request: %w", errUnmarshal)
}
if strings.TrimSpace(payload.SDP) == "" {
return "", errors.New("Codex live call request requires an SDP offer")
}
return payload.SDP, nil
}
func replaceCallRequestSDP(body []byte, contentType, sdp string) ([]byte, string, error) {
mediaType, _, errMediaType := mime.ParseMediaType(contentType)
if errMediaType == nil && (strings.EqualFold(mediaType, "application/sdp") || strings.EqualFold(mediaType, "text/plain")) {
encoded, errEncode := encodeCallRequest(sdp, nil)
if errEncode != nil {
return nil, "", errEncode
}
return encoded, "application/json", nil
}
if errMediaType != nil || !strings.EqualFold(mediaType, "application/json") {
return nil, "", errors.New("Codex live media relay requires 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 Codex live call request: %w", errUnmarshal)
}
encodedSDP, errMarshal := json.Marshal(sdp)
if errMarshal != nil {
return nil, "", fmt.Errorf("failed to encode Codex live SDP offer: %w", errMarshal)
}
payload["sdp"] = encodedSDP
encoded, errMarshal := json.Marshal(payload)
if errMarshal != nil {
return nil, "", fmt.Errorf("failed to encode Codex live call request: %w", errMarshal)
}
return encoded, "application/json", nil
}
func callResponseSDP(body []byte, contentType string) (string, error) {
mediaType, _, errMediaType := mime.ParseMediaType(contentType)
if errMediaType == nil && strings.EqualFold(mediaType, "application/json") {
var payload struct {
SDP string `json:"sdp"`
}
if errUnmarshal := json.Unmarshal(body, &payload); errUnmarshal != nil {
return "", fmt.Errorf("failed to decode Codex live response: %w", errUnmarshal)
}
if strings.TrimSpace(payload.SDP) == "" {
return "", errors.New("Codex live response requires an SDP answer")
}
return payload.SDP, nil
}
if strings.TrimSpace(string(body)) == "" {
return "", errors.New("Codex live response requires an SDP answer")
}
return string(body), nil
}
func modelFromJSON(body []byte) string {

View File

@@ -3,6 +3,7 @@ package live
import (
"context"
"encoding/json"
"errors"
"io"
"net/http"
"net/http/httptest"
@@ -38,6 +39,7 @@ type captureExecutor struct {
body []byte
selectedAuth *auth.Auth
responseBody io.ReadCloser
statusCode int
}
func (*captureExecutor) Identifier() string { return "codex" }
@@ -72,8 +74,12 @@ func (e *captureExecutor) HttpRequest(_ context.Context, credential *auth.Auth,
return nil, errRead
}
e.body = body
statusCode := e.statusCode
if statusCode == 0 {
statusCode = http.StatusCreated
}
return &http.Response{
StatusCode: http.StatusCreated,
StatusCode: statusCode,
Header: http.Header{
"Connection": []string{"X-Connection-Secret"},
"Content-Type": []string{"application/sdp"},
@@ -114,6 +120,23 @@ func (d *homeDispatcher) RPopAuth(_ context.Context, model string, _ string, _ h
func (*homeDispatcher) AbortAmbiguousDispatch() {}
type failingHTTPWriter struct {
header http.Header
status int
}
func (w *failingHTTPWriter) Header() http.Header {
return w.header
}
func (*failingHTTPWriter) Write([]byte) (int, error) {
return 0, errors.New("downstream write failed")
}
func (w *failingHTTPWriter) WriteHeader(statusCode int) {
w.status = statusCode
}
type trackedResponseBody struct {
io.Reader
closed atomic.Bool
@@ -124,6 +147,40 @@ func (b *trackedResponseBody) Close() error {
return nil
}
type fakeMediaRelay struct {
clientOffer string
upstreamOffer string
session *fakeMediaSession
err error
}
func (r *fakeMediaRelay) NewSession(_ context.Context, clientOffer string) (mediaRelaySession, string, error) {
r.clientOffer = clientOffer
return r.session, r.upstreamOffer, r.err
}
type fakeMediaSession struct {
upstreamAnswer string
downstreamSDP string
closeHandler func(string)
closed atomic.Bool
err error
}
func (s *fakeMediaSession) AcceptUpstreamAnswer(_ context.Context, answer string) (string, error) {
s.upstreamAnswer = answer
return s.downstreamSDP, s.err
}
func (s *fakeMediaSession) SetCloseHandler(handler func(string)) {
s.closeHandler = handler
}
func (s *fakeMediaSession) Close() error {
s.closed.Store(true)
return nil
}
func registerCredential(t *testing.T, manager *auth.Manager, credential *auth.Auth) {
t.Helper()
if _, errRegister := manager.Register(context.Background(), credential); errRegister != nil {
@@ -250,6 +307,207 @@ func TestHandlerRewritesLiveCallAndSchedulesOAuth(t *testing.T) {
}
}
func TestHandlerRelaysWebRTCMediaSDP(t *testing.T) {
gin.SetMode(gin.TestMode)
manager := auth.NewManager(nil, nil, nil)
executor := &captureExecutor{
responseBody: &trackedResponseBody{Reader: strings.NewReader("v=0\r\no=upstream-answer\r\n")},
}
manager.RegisterExecutor(executor)
registerCredential(t, manager, &auth.Auth{
ID: "codex-oauth",
Provider: "codex",
Status: auth.StatusActive,
Metadata: map[string]any{"access_token": "oauth-token"},
})
mediaSession := &fakeMediaSession{downstreamSDP: "v=0\r\no=downstream-answer\r\n"}
mediaRelay := &fakeMediaRelay{
upstreamOffer: "v=0\r\no=gateway-offer\r\n",
session: mediaSession,
}
handler := NewHandler(manager, nil)
handler.mediaRelay = mediaRelay
router := gin.New()
router.POST("/v1/live", handler.Handle)
const boundary = "media-relay-boundary"
body := multipartBody(boundary, "v=0\r\no=desktop-offer\r\n", `{"model":"gpt-live-1-codex"}`)
req := httptest.NewRequest(http.MethodPost, "/v1/live", strings.NewReader(body))
req.Header.Set("Content-Type", "multipart/form-data; boundary="+boundary)
recorder := httptest.NewRecorder()
router.ServeHTTP(recorder, req)
if recorder.Code != http.StatusCreated {
t.Fatalf("status = %d, want %d; body=%s", recorder.Code, http.StatusCreated, recorder.Body.String())
}
if mediaRelay.clientOffer != "v=0\r\no=desktop-offer\r\n" {
t.Fatalf("media client offer = %q", mediaRelay.clientOffer)
}
var upstreamPayload struct {
SDP string `json:"sdp"`
}
if errUnmarshal := json.Unmarshal(executor.body, &upstreamPayload); errUnmarshal != nil {
t.Fatalf("unmarshal upstream body: %v", errUnmarshal)
}
if upstreamPayload.SDP != mediaRelay.upstreamOffer {
t.Fatalf("upstream SDP = %q, want gateway offer", upstreamPayload.SDP)
}
if mediaSession.upstreamAnswer != "v=0\r\no=upstream-answer\r\n" {
t.Fatalf("accepted upstream answer = %q", mediaSession.upstreamAnswer)
}
if got := recorder.Body.String(); got != mediaSession.downstreamSDP {
t.Fatalf("downstream SDP = %q, want %q", got, mediaSession.downstreamSDP)
}
if got := recorder.Header().Get("Content-Type"); got != "application/sdp" {
t.Fatalf("Content-Type = %q, want application/sdp", got)
}
if mediaSession.closed.Load() {
t.Fatal("retained media session was closed before session completion")
}
if mediaSession.closeHandler == nil {
t.Fatal("media session close handler was not installed")
}
mediaSession.closeHandler("test_closed")
if !mediaSession.closed.Load() {
t.Fatal("completed media session was not closed")
}
if _, ok := handler.sessions.peek("call-123"); ok {
t.Fatal("completed media session remained stored")
}
}
func TestHandlerClosesUnretainedMediaSession(t *testing.T) {
for name, testCase := range map[string]struct {
upstreamStatus int
answerError error
wantStatus int
}{
"upstream rejection": {
upstreamStatus: http.StatusUnauthorized,
wantStatus: http.StatusUnauthorized,
},
"invalid upstream answer": {
upstreamStatus: http.StatusCreated,
answerError: errors.New("invalid answer"),
wantStatus: http.StatusBadGateway,
},
} {
t.Run(name, func(t *testing.T) {
gin.SetMode(gin.TestMode)
manager := auth.NewManager(nil, nil, nil)
executor := &captureExecutor{
responseBody: &trackedResponseBody{Reader: strings.NewReader("v=0\r\no=upstream-answer\r\n")},
statusCode: testCase.upstreamStatus,
}
manager.RegisterExecutor(executor)
registerCredential(t, manager, &auth.Auth{
ID: "codex-oauth",
Provider: "codex",
Status: auth.StatusActive,
Metadata: map[string]any{"access_token": "oauth-token"},
})
mediaSession := &fakeMediaSession{
downstreamSDP: "v=0\r\no=downstream-answer\r\n",
err: testCase.answerError,
}
handler := NewHandler(manager, nil)
handler.mediaRelay = &fakeMediaRelay{
upstreamOffer: "v=0\r\no=gateway-offer\r\n",
session: mediaSession,
}
router := gin.New()
router.POST("/v1/live", handler.Handle)
const boundary = "media-error-boundary"
body := multipartBody(boundary, "v=0\r\no=desktop-offer\r\n", `{"model":"gpt-live-1-codex"}`)
req := httptest.NewRequest(http.MethodPost, "/v1/live", strings.NewReader(body))
req.Header.Set("Content-Type", "multipart/form-data; boundary="+boundary)
recorder := httptest.NewRecorder()
router.ServeHTTP(recorder, req)
if recorder.Code != testCase.wantStatus {
t.Fatalf("status = %d, want %d; body=%s", recorder.Code, testCase.wantStatus, recorder.Body.String())
}
if !mediaSession.closed.Load() {
t.Fatal("failed request retained its media session")
}
if _, ok := handler.sessions.peek("call-123"); ok {
t.Fatal("failed request stored its media session")
}
})
}
}
func TestHandlerReleasesHomeSelectionWhenMediaSetupFails(t *testing.T) {
gin.SetMode(gin.TestMode)
manager := auth.NewManager(nil, nil, nil)
manager.SetConfig(&config.Config{Home: config.HomeConfig{Enabled: true}})
registry := executionregistry.New()
manager.PublishHomeDispatch(&homeDispatcher{}, registry, 1)
manager.RegisterExecutor(&captureExecutor{})
handler := NewHandler(manager, nil)
handler.mediaRelay = &fakeMediaRelay{err: errors.New("media setup failed")}
router := gin.New()
router.POST("/v1/live", handler.Handle)
const boundary = "home-media-error-boundary"
body := multipartBody(boundary, "v=0\r\no=desktop-offer\r\n", `{"model":"gpt-live-1-codex"}`)
req := httptest.NewRequest(http.MethodPost, "/v1/live", strings.NewReader(body))
req.Header.Set("Content-Type", "multipart/form-data; boundary="+boundary)
recorder := httptest.NewRecorder()
router.ServeHTTP(recorder, req)
if recorder.Code != http.StatusBadGateway {
t.Fatalf("status = %d, want %d; body=%s", recorder.Code, http.StatusBadGateway, recorder.Body.String())
}
if got := len(registry.FreezeInFlight(time.Now()).Executions); got != 0 {
t.Fatalf("active Home executions = %d, want 0", got)
}
if errDrain := registry.Drain(context.Background()); errDrain != nil {
t.Fatalf("Drain() error = %v", errDrain)
}
}
func TestHandlerClosesMediaWhenResponseWriteFails(t *testing.T) {
gin.SetMode(gin.TestMode)
manager := auth.NewManager(nil, nil, nil)
manager.RegisterExecutor(&captureExecutor{
responseBody: &trackedResponseBody{Reader: strings.NewReader("v=0\r\no=upstream-answer\r\n")},
})
registerCredential(t, manager, &auth.Auth{
ID: "codex-oauth",
Provider: "codex",
Status: auth.StatusActive,
Metadata: map[string]any{"access_token": "oauth-token"},
})
mediaSession := &fakeMediaSession{downstreamSDP: "v=0\r\no=downstream-answer\r\n"}
handler := NewHandler(manager, nil)
handler.mediaRelay = &fakeMediaRelay{
upstreamOffer: "v=0\r\no=gateway-offer\r\n",
session: mediaSession,
}
router := gin.New()
router.POST("/v1/live", handler.Handle)
const boundary = "response-write-error-boundary"
body := multipartBody(boundary, "v=0\r\no=desktop-offer\r\n", `{"model":"gpt-live-1-codex"}`)
req := httptest.NewRequest(http.MethodPost, "/v1/live", strings.NewReader(body))
req.Header.Set("Content-Type", "multipart/form-data; boundary="+boundary)
writer := &failingHTTPWriter{header: make(http.Header)}
router.ServeHTTP(writer, req)
if writer.status != http.StatusCreated {
t.Fatalf("status = %d, want %d", writer.status, http.StatusCreated)
}
if !mediaSession.closed.Load() {
t.Fatal("response write failure retained its media session")
}
if _, ok := handler.sessions.peek("call-123"); ok {
t.Fatal("response write failure retained a stored session")
}
}
func TestHandlerUsesLiveModelForHomeDispatch(t *testing.T) {
gin.SetMode(gin.TestMode)
@@ -454,6 +712,74 @@ func TestPrepareCallRequestRewritesMultipart(t *testing.T) {
}
}
func TestPrepareCallRequestPreservesRawSDPWhenRelayDisabled(t *testing.T) {
body := []byte("v=0\r\no=raw-offer\r\n")
prepared, contentType, model, errPrepare := prepareCallRequest(body, "application/sdp")
if errPrepare != nil {
t.Fatalf("prepareCallRequest() error = %v", errPrepare)
}
if string(prepared) != string(body) {
t.Fatalf("prepared SDP = %q, want original body", prepared)
}
if contentType != "application/sdp" {
t.Fatalf("content type = %q, want application/sdp", contentType)
}
if model != defaultLiveModel {
t.Fatalf("model = %q, want %q", model, defaultLiveModel)
}
}
func TestMediaRelayWrapsRawSDPForCodexBackend(t *testing.T) {
body := []byte("v=0\r\no=raw-offer\r\n")
clientOffer, errSDP := callRequestSDP(body, "application/sdp")
if errSDP != nil {
t.Fatalf("callRequestSDP() error = %v", errSDP)
}
if clientOffer != string(body) {
t.Fatalf("client offer = %q, want original body", clientOffer)
}
prepared, contentType, errReplace := replaceCallRequestSDP(body, "application/sdp", "v=0\r\no=gateway-offer\r\n")
if errReplace != nil {
t.Fatalf("replaceCallRequestSDP() error = %v", errReplace)
}
if contentType != "application/json" {
t.Fatalf("content type = %q, want application/json", contentType)
}
var payload struct {
SDP string `json:"sdp"`
}
if errUnmarshal := json.Unmarshal(prepared, &payload); errUnmarshal != nil {
t.Fatalf("unmarshal prepared request: %v", errUnmarshal)
}
if payload.SDP != "v=0\r\no=gateway-offer\r\n" {
t.Fatalf("upstream SDP = %q", payload.SDP)
}
}
func TestHandlerUpdatesMediaRelayConfig(t *testing.T) {
handler := NewHandler(nil, nil)
if relay, errRelay := handler.currentMediaRelay(); relay != nil || errRelay != nil {
t.Fatalf("initial media relay = %#v, error = %v", relay, errRelay)
}
enabled := &config.Config{Codex: config.CodexConfig{LiveMediaRelay: config.CodexLiveMediaRelayConfig{
Enabled: true,
MaxSessions: 1,
AllowPrivateRemoteIPs: true,
}}}
if errUpdate := handler.UpdateConfig(enabled); errUpdate != nil {
t.Fatalf("enable media relay: %v", errUpdate)
}
if relay, errRelay := handler.currentMediaRelay(); relay == nil || errRelay != nil {
t.Fatalf("enabled media relay = %#v, error = %v", relay, errRelay)
}
if errUpdate := handler.UpdateConfig(&config.Config{}); errUpdate != nil {
t.Fatalf("disable media relay: %v", errUpdate)
}
if relay, errRelay := handler.currentMediaRelay(); relay != nil || errRelay != nil {
t.Fatalf("disabled media relay = %#v, error = %v", relay, errRelay)
}
}
func TestPrepareCallRequestRejectsInvalidMultipart(t *testing.T) {
const boundary = "invalid-live-boundary"
body := "--" + boundary + "\r\n" +
@@ -509,6 +835,29 @@ func TestSessionStoreClaimsAndExpiresSessions(t *testing.T) {
t.Fatal("released live session did not expire")
}
func TestSessionStoreCloseAllReleasesMediaAndResources(t *testing.T) {
store := newSessionStore()
mediaSession := &fakeMediaSession{}
stored := store.put("call-close-all", liveSession{media: mediaSession})
var resourceClosed atomic.Bool
stored.resources.add(func() error {
resourceClosed.Store(true)
return nil
})
store.closeAll("test_shutdown")
if !mediaSession.closed.Load() {
t.Fatal("closeAll() did not close the media session")
}
if !resourceClosed.Load() {
t.Fatal("closeAll() did not close session resources")
}
if _, ok := store.peek("call-close-all"); ok {
t.Fatal("closeAll() retained a session")
}
}
func TestSidebandURLShapes(t *testing.T) {
if got := buildSidebandURL(defaultSidebandAPIBaseURL, sidebandFrameless, "rtc_1"); got != "wss://api.openai.com/v1/live/rtc_1" {
t.Fatalf("Frameless sideband URL = %q", got)

View File

@@ -0,0 +1,618 @@
package live
import (
"context"
"errors"
"fmt"
"io"
"net"
"strings"
"sync"
"github.com/pion/interceptor"
"github.com/pion/rtp"
"github.com/pion/webrtc/v4"
"github.com/router-for-me/CLIProxyAPI/v7/internal/config"
log "github.com/sirupsen/logrus"
)
const (
realtimeDataChannelLabel = "oai-events"
mediaDataQueueSize = 64
mediaDataMessageMaxSize = 256 << 10
mediaDataBufferedMaxSize = 1 << 20
)
var opusCodec = webrtc.RTPCodecCapability{
MimeType: webrtc.MimeTypeOpus,
ClockRate: 48000,
Channels: 2,
SDPFmtpLine: "minptime=10;useinbandfec=1",
}
type mediaRelaySession interface {
AcceptUpstreamAnswer(context.Context, string) (string, error)
SetCloseHandler(func(string))
Close() error
}
type mediaRelayFactory interface {
NewSession(context.Context, string) (mediaRelaySession, string, error)
}
type pionMediaRelay struct {
downstreamAPI *webrtc.API
upstreamAPI *webrtc.API
configuration webrtc.Configuration
slots chan struct{}
}
type pionMediaSession struct {
downstream *webrtc.PeerConnection
upstream *webrtc.PeerConnection
bridge *dataChannelBridge
done chan struct{}
closeOnce sync.Once
closeErr error
failureOnce sync.Once
handlerMu sync.Mutex
onClose func(string)
failureReason string
handlerCalled bool
releaseSlot func()
}
type dataChannelMessage struct {
data []byte
isString bool
}
type dataChannelPipe struct {
name string
done <-chan struct{}
queue chan dataChannelMessage
ready chan struct{}
readyOnce sync.Once
writable chan struct{}
destination *webrtc.DataChannel
mu sync.RWMutex
onError func(error)
}
type dataChannelBridge struct {
done <-chan struct{}
downToUp *dataChannelPipe
upToDown *dataChannelPipe
closeOnce sync.Once
downstreamMu sync.Mutex
downstream *webrtc.DataChannel
upstreamMu sync.Mutex
upstream *webrtc.DataChannel
}
func newPionMediaRelay(relayConfig config.CodexLiveMediaRelayConfig) (*pionMediaRelay, error) {
if errValidate := relayConfig.Validate(); errValidate != nil {
return nil, errValidate
}
downstreamAPI, errAPI := newPionAPI(relayConfig, !relayConfig.AllowPrivateRemoteIPs)
if errAPI != nil {
return nil, errAPI
}
upstreamAPI, errAPI := newPionAPI(relayConfig, false)
if errAPI != nil {
return nil, errAPI
}
iceServers := make([]webrtc.ICEServer, 0, len(relayConfig.ICEServers))
for _, server := range relayConfig.ICEServers {
urls := make([]string, 0, len(server.URLs))
for _, rawURL := range server.URLs {
urls = append(urls, strings.TrimSpace(rawURL))
}
iceServers = append(iceServers, webrtc.ICEServer{
URLs: urls,
Username: server.Username,
Credential: server.Credential,
CredentialType: webrtc.ICECredentialTypePassword,
})
}
return &pionMediaRelay{
downstreamAPI: downstreamAPI,
upstreamAPI: upstreamAPI,
configuration: webrtc.Configuration{ICEServers: iceServers},
slots: make(chan struct{}, relayConfig.EffectiveMaxSessions()),
}, nil
}
func newPionAPI(relayConfig config.CodexLiveMediaRelayConfig, filterPrivateRemoteIPs bool) (*webrtc.API, error) {
mediaEngine := &webrtc.MediaEngine{}
if errRegister := mediaEngine.RegisterCodec(webrtc.RTPCodecParameters{
RTPCodecCapability: opusCodec,
PayloadType: 111,
}, webrtc.RTPCodecTypeAudio); errRegister != nil {
return nil, fmt.Errorf("register Opus codec: %w", errRegister)
}
interceptorRegistry := &interceptor.Registry{}
if errRegister := webrtc.RegisterDefaultInterceptors(mediaEngine, interceptorRegistry); errRegister != nil {
return nil, fmt.Errorf("register WebRTC interceptors: %w", errRegister)
}
settingEngine := webrtc.SettingEngine{}
if relayConfig.UDPPortMin != 0 {
if errPorts := settingEngine.SetEphemeralUDPPortRange(relayConfig.UDPPortMin, relayConfig.UDPPortMax); errPorts != nil {
return nil, fmt.Errorf("configure WebRTC UDP port range: %w", errPorts)
}
}
if publicIP := strings.TrimSpace(relayConfig.PublicIP); publicIP != "" {
settingEngine.SetNAT1To1IPs([]string{publicIP}, webrtc.ICECandidateTypeHost)
}
if filterPrivateRemoteIPs {
settingEngine.SetRemoteIPFilter(isPublicRemoteIP)
}
return webrtc.NewAPI(
webrtc.WithMediaEngine(mediaEngine),
webrtc.WithInterceptorRegistry(interceptorRegistry),
webrtc.WithSettingEngine(settingEngine),
), nil
}
func isPublicRemoteIP(ip net.IP) bool {
return ip != nil && !ip.IsUnspecified() && !ip.IsLoopback() && !ip.IsPrivate() &&
!ip.IsLinkLocalUnicast() && !ip.IsLinkLocalMulticast() && !ip.IsMulticast()
}
func (r *pionMediaRelay) NewSession(ctx context.Context, clientOffer string) (mediaRelaySession, string, error) {
if r == nil || r.downstreamAPI == nil || r.upstreamAPI == nil {
return nil, "", errors.New("Codex live media relay unavailable")
}
select {
case r.slots <- struct{}{}:
case <-ctx.Done():
return nil, "", ctx.Err()
default:
return nil, "", errors.New("Codex live media relay capacity exhausted")
}
releaseSlot := func() { <-r.slots }
downstream, errDownstream := r.downstreamAPI.NewPeerConnection(r.configuration)
if errDownstream != nil {
releaseSlot()
return nil, "", fmt.Errorf("create downstream PeerConnection: %w", errDownstream)
}
upstream, errUpstream := r.upstreamAPI.NewPeerConnection(r.configuration)
if errUpstream != nil {
releaseSlot()
if errClose := downstream.Close(); errClose != nil {
log.WithError(errClose).Debug("codex live media: close downstream PeerConnection after setup error")
}
return nil, "", fmt.Errorf("create upstream PeerConnection: %w", errUpstream)
}
session := &pionMediaSession{
downstream: downstream,
upstream: upstream,
done: make(chan struct{}),
releaseSlot: releaseSlot,
}
session.bridge = newDataChannelBridge(session.done, func(err error) {
session.fail("data_channel_failed", err)
})
session.installStateHandlers()
if errRemote := downstream.SetRemoteDescription(webrtc.SessionDescription{
Type: webrtc.SDPTypeOffer,
SDP: clientOffer,
}); errRemote != nil {
_ = session.Close()
return nil, "", fmt.Errorf("set downstream WebRTC offer: %w", errRemote)
}
toDesktop, errTrack := webrtc.NewTrackLocalStaticRTP(opusCodec, "audio", "codex-live")
if errTrack != nil {
_ = session.Close()
return nil, "", fmt.Errorf("create downstream audio track: %w", errTrack)
}
downstreamSender, errTrack := downstream.AddTrack(toDesktop)
if errTrack != nil {
_ = session.Close()
return nil, "", fmt.Errorf("add downstream audio track: %w", errTrack)
}
go drainRTCP("downstream", downstreamSender, session.done)
toOpenAI, errTrack := webrtc.NewTrackLocalStaticRTP(opusCodec, "audio", "codex-live")
if errTrack != nil {
_ = session.Close()
return nil, "", fmt.Errorf("create upstream audio track: %w", errTrack)
}
upstreamSender, errTrack := upstream.AddTrack(toOpenAI)
if errTrack != nil {
_ = session.Close()
return nil, "", fmt.Errorf("add upstream audio track: %w", errTrack)
}
go drainRTCP("upstream", upstreamSender, session.done)
downstream.OnTrack(func(track *webrtc.TrackRemote, _ *webrtc.RTPReceiver) {
if !strings.EqualFold(track.Codec().MimeType, webrtc.MimeTypeOpus) {
return
}
go relayRTP("downstream-to-upstream", track, toOpenAI, session.done)
})
upstream.OnTrack(func(track *webrtc.TrackRemote, _ *webrtc.RTPReceiver) {
if !strings.EqualFold(track.Codec().MimeType, webrtc.MimeTypeOpus) {
return
}
go relayRTP("upstream-to-downstream", track, toDesktop, session.done)
})
downstream.OnDataChannel(func(channel *webrtc.DataChannel) {
if channel.Label() != realtimeDataChannelLabel {
if errClose := channel.Close(); errClose != nil {
log.WithError(errClose).Debug("codex live media: close unsupported downstream DataChannel")
}
return
}
session.bridge.attachDownstream(channel)
})
upstreamChannel, errChannel := upstream.CreateDataChannel(realtimeDataChannelLabel, nil)
if errChannel != nil {
_ = session.Close()
return nil, "", fmt.Errorf("create upstream DataChannel: %w", errChannel)
}
session.bridge.attachUpstream(upstreamChannel)
gatherComplete := webrtc.GatheringCompletePromise(upstream)
offer, errOffer := upstream.CreateOffer(nil)
if errOffer != nil {
_ = session.Close()
return nil, "", fmt.Errorf("create upstream WebRTC offer: %w", errOffer)
}
if errLocal := upstream.SetLocalDescription(offer); errLocal != nil {
_ = session.Close()
return nil, "", fmt.Errorf("set upstream WebRTC offer: %w", errLocal)
}
select {
case <-gatherComplete:
case <-ctx.Done():
_ = session.Close()
return nil, "", fmt.Errorf("gather upstream WebRTC candidates: %w", ctx.Err())
}
localDescription := upstream.LocalDescription()
if localDescription == nil || strings.TrimSpace(localDescription.SDP) == "" {
_ = session.Close()
return nil, "", errors.New("upstream WebRTC offer is empty")
}
return session, localDescription.SDP, nil
}
func (s *pionMediaSession) AcceptUpstreamAnswer(ctx context.Context, upstreamAnswer string) (string, error) {
if s == nil || s.upstream == nil || s.downstream == nil {
return "", errors.New("Codex live media session unavailable")
}
if errRemote := s.upstream.SetRemoteDescription(webrtc.SessionDescription{
Type: webrtc.SDPTypeAnswer,
SDP: upstreamAnswer,
}); errRemote != nil {
return "", fmt.Errorf("set upstream WebRTC answer: %w", errRemote)
}
gatherComplete := webrtc.GatheringCompletePromise(s.downstream)
answer, errAnswer := s.downstream.CreateAnswer(nil)
if errAnswer != nil {
return "", fmt.Errorf("create downstream WebRTC answer: %w", errAnswer)
}
if errLocal := s.downstream.SetLocalDescription(answer); errLocal != nil {
return "", fmt.Errorf("set downstream WebRTC answer: %w", errLocal)
}
select {
case <-gatherComplete:
case <-ctx.Done():
return "", fmt.Errorf("gather downstream WebRTC candidates: %w", ctx.Err())
}
localDescription := s.downstream.LocalDescription()
if localDescription == nil || strings.TrimSpace(localDescription.SDP) == "" {
return "", errors.New("downstream WebRTC answer is empty")
}
return localDescription.SDP, nil
}
func (s *pionMediaSession) SetCloseHandler(handler func(string)) {
if s == nil {
return
}
s.handlerMu.Lock()
s.onClose = handler
reason := s.failureReason
callHandler := handler != nil && reason != "" && !s.handlerCalled
if callHandler {
s.handlerCalled = true
}
s.handlerMu.Unlock()
if callHandler {
handler(reason)
}
}
func (s *pionMediaSession) Close() error {
if s == nil {
return nil
}
s.closeOnce.Do(func() {
close(s.done)
if s.bridge != nil {
s.bridge.close()
}
var closeErrors []error
if s.downstream != nil {
if errClose := s.downstream.Close(); errClose != nil {
closeErrors = append(closeErrors, fmt.Errorf("close downstream PeerConnection: %w", errClose))
}
}
if s.upstream != nil {
if errClose := s.upstream.Close(); errClose != nil {
closeErrors = append(closeErrors, fmt.Errorf("close upstream PeerConnection: %w", errClose))
}
}
if s.releaseSlot != nil {
s.releaseSlot()
}
s.closeErr = errors.Join(closeErrors...)
})
return s.closeErr
}
func (s *pionMediaSession) installStateHandlers() {
handle := func(leg string) func(webrtc.PeerConnectionState) {
return func(state webrtc.PeerConnectionState) {
switch state {
case webrtc.PeerConnectionStateFailed:
s.fail(leg+"_failed", fmt.Errorf("%s PeerConnection failed", leg))
case webrtc.PeerConnectionStateClosed:
select {
case <-s.done:
return
default:
s.fail(leg+"_closed", fmt.Errorf("%s PeerConnection closed", leg))
}
}
}
}
s.downstream.OnConnectionStateChange(handle("downstream"))
s.upstream.OnConnectionStateChange(handle("upstream"))
}
func (s *pionMediaSession) fail(reason string, err error) {
s.failureOnce.Do(func() {
if err != nil {
log.WithError(err).Debug("codex live media relay closed")
}
if errClose := s.Close(); errClose != nil {
log.WithError(errClose).Debug("codex live media: close failed session")
}
s.handlerMu.Lock()
s.failureReason = reason
handler := s.onClose
callHandler := handler != nil && !s.handlerCalled
if callHandler {
s.handlerCalled = true
}
s.handlerMu.Unlock()
if callHandler {
handler(reason)
}
})
}
func relayRTP(name string, source *webrtc.TrackRemote, destination *webrtc.TrackLocalStaticRTP, done <-chan struct{}) {
for {
packet, _, errRead := source.ReadRTP()
if errRead != nil {
if !isClosedMediaError(errRead, done) {
log.WithError(errRead).Debugf("codex live media: %s RTP read stopped", name)
}
return
}
normalizeRTPPacket(packet)
if errWrite := destination.WriteRTP(packet); errWrite != nil {
if !isClosedMediaError(errWrite, done) {
log.WithError(errWrite).Debugf("codex live media: %s RTP write stopped", name)
}
return
}
}
}
func normalizeRTPPacket(packet *rtp.Packet) {
if packet == nil {
return
}
packet.Extension = false
packet.ExtensionProfile = 0
packet.Extensions = nil
}
func drainRTCP(name string, sender *webrtc.RTPSender, done <-chan struct{}) {
for {
if _, _, errRead := sender.ReadRTCP(); errRead != nil {
if !isClosedMediaError(errRead, done) {
log.WithError(errRead).Debugf("codex live media: %s RTCP reader stopped", name)
}
return
}
}
}
func isClosedMediaError(err error, done <-chan struct{}) bool {
select {
case <-done:
return true
default:
}
return errors.Is(err, io.EOF) || errors.Is(err, net.ErrClosed)
}
func newDataChannelBridge(done <-chan struct{}, onError func(error)) *dataChannelBridge {
bridge := &dataChannelBridge{done: done}
bridge.downToUp = newDataChannelPipe("downstream-to-upstream", done, onError)
bridge.upToDown = newDataChannelPipe("upstream-to-downstream", done, onError)
return bridge
}
func newDataChannelPipe(name string, done <-chan struct{}, onError func(error)) *dataChannelPipe {
pipe := &dataChannelPipe{
name: name,
done: done,
queue: make(chan dataChannelMessage, mediaDataQueueSize),
ready: make(chan struct{}),
writable: make(chan struct{}, 1),
onError: onError,
}
go pipe.run()
return pipe
}
func (b *dataChannelBridge) attachDownstream(channel *webrtc.DataChannel) {
b.downstreamMu.Lock()
if b.downstream != nil {
b.downstreamMu.Unlock()
if errClose := channel.Close(); errClose != nil {
log.WithError(errClose).Debug("codex live media: close duplicate downstream DataChannel")
}
return
}
b.downstream = channel
b.downstreamMu.Unlock()
b.upToDown.setDestination(channel)
b.bindSource(channel, b.downToUp)
}
func (b *dataChannelBridge) attachUpstream(channel *webrtc.DataChannel) {
b.upstreamMu.Lock()
if b.upstream != nil {
b.upstreamMu.Unlock()
if errClose := channel.Close(); errClose != nil {
log.WithError(errClose).Debug("codex live media: close duplicate upstream DataChannel")
}
return
}
b.upstream = channel
b.upstreamMu.Unlock()
b.downToUp.setDestination(channel)
b.bindSource(channel, b.upToDown)
}
func (b *dataChannelBridge) bindSource(channel *webrtc.DataChannel, destination *dataChannelPipe) {
channel.OnMessage(func(message webrtc.DataChannelMessage) {
if len(message.Data) > mediaDataMessageMaxSize {
destination.reportError(fmt.Errorf("%s DataChannel message exceeds %d bytes", destination.name, mediaDataMessageMaxSize))
return
}
payload := append([]byte(nil), message.Data...)
select {
case destination.queue <- dataChannelMessage{data: payload, isString: message.IsString}:
case <-b.done:
}
})
channel.OnError(func(err error) {
destination.reportError(fmt.Errorf("%s DataChannel error: %w", destination.name, err))
})
channel.OnClose(func() {
select {
case <-b.done:
return
default:
destination.reportError(fmt.Errorf("%s DataChannel closed", destination.name))
}
})
}
func (b *dataChannelBridge) close() {
if b == nil {
return
}
b.closeOnce.Do(func() {
b.downstreamMu.Lock()
downstream := b.downstream
b.downstreamMu.Unlock()
if downstream != nil {
if errClose := downstream.Close(); errClose != nil {
log.WithError(errClose).Debug("codex live media: close downstream DataChannel")
}
}
b.upstreamMu.Lock()
upstream := b.upstream
b.upstreamMu.Unlock()
if upstream != nil {
if errClose := upstream.Close(); errClose != nil {
log.WithError(errClose).Debug("codex live media: close upstream DataChannel")
}
}
})
}
func (p *dataChannelPipe) setDestination(channel *webrtc.DataChannel) {
p.mu.Lock()
p.destination = channel
p.mu.Unlock()
markReady := func() {
p.readyOnce.Do(func() { close(p.ready) })
}
channel.SetBufferedAmountLowThreshold(mediaDataBufferedMaxSize / 2)
channel.OnBufferedAmountLow(func() {
select {
case p.writable <- struct{}{}:
default:
}
})
channel.OnOpen(markReady)
if channel.ReadyState() == webrtc.DataChannelStateOpen {
markReady()
}
}
func (p *dataChannelPipe) run() {
select {
case <-p.ready:
case <-p.done:
return
}
for {
select {
case message := <-p.queue:
p.mu.RLock()
destination := p.destination
p.mu.RUnlock()
if destination == nil {
p.reportError(fmt.Errorf("%s DataChannel destination unavailable", p.name))
return
}
if !p.waitWritable(destination, len(message.data)) {
return
}
var errSend error
if message.isString {
errSend = destination.SendText(string(message.data))
} else {
errSend = destination.Send(message.data)
}
if errSend != nil {
p.reportError(fmt.Errorf("send %s DataChannel message: %w", p.name, errSend))
return
}
case <-p.done:
return
}
}
}
func (p *dataChannelPipe) waitWritable(destination *webrtc.DataChannel, messageSize int) bool {
for destination.BufferedAmount()+uint64(messageSize) > mediaDataBufferedMaxSize {
select {
case <-p.writable:
case <-p.done:
return false
}
}
return true
}
func (p *dataChannelPipe) reportError(err error) {
if p.onError != nil {
p.onError(err)
}
}

View File

@@ -0,0 +1,295 @@
package live
import (
"context"
"net"
"testing"
"time"
"github.com/pion/interceptor"
"github.com/pion/rtp"
"github.com/pion/webrtc/v4"
"github.com/router-for-me/CLIProxyAPI/v7/internal/config"
)
func TestPionMediaRelayBridgesAudioAndDataChannel(t *testing.T) {
clientAPI := newTestWebRTCAPI(t)
client, errClient := clientAPI.NewPeerConnection(webrtc.Configuration{})
if errClient != nil {
t.Fatalf("create client PeerConnection: %v", errClient)
}
defer closeTestPeerConnection(t, client)
clientDone := make(chan struct{})
defer close(clientDone)
clientAudio, errTrack := webrtc.NewTrackLocalStaticRTP(opusCodec, "client-audio", "client")
if errTrack != nil {
t.Fatalf("create client audio track: %v", errTrack)
}
clientSender, errTrack := client.AddTrack(clientAudio)
if errTrack != nil {
t.Fatalf("add client audio track: %v", errTrack)
}
go drainRTCP("test-client", clientSender, clientDone)
clientData, errData := client.CreateDataChannel(realtimeDataChannelLabel, nil)
if errData != nil {
t.Fatalf("create client DataChannel: %v", errData)
}
clientMessages := make(chan webrtc.DataChannelMessage, 4)
clientData.OnMessage(func(message webrtc.DataChannelMessage) {
message.Data = append([]byte(nil), message.Data...)
clientMessages <- message
})
clientAudioMessages := make(chan []byte, 1)
client.OnTrack(func(track *webrtc.TrackRemote, _ *webrtc.RTPReceiver) {
packet, _, errRead := track.ReadRTP()
if errRead == nil {
clientAudioMessages <- append([]byte(nil), packet.Payload...)
}
})
clientOffer := completeOffer(t, client)
relay, errRelay := newPionMediaRelay(config.CodexLiveMediaRelayConfig{
Enabled: true,
MaxSessions: 1,
AllowPrivateRemoteIPs: true,
})
if errRelay != nil {
t.Fatalf("create media relay: %v", errRelay)
}
session, relayOffer, errSession := relay.NewSession(context.Background(), clientOffer)
if errSession != nil {
t.Fatalf("create media relay session: %v", errSession)
}
defer func() {
if errClose := session.Close(); errClose != nil {
t.Errorf("close media relay session: %v", errClose)
}
}()
if _, _, errCapacity := relay.NewSession(context.Background(), clientOffer); errCapacity == nil {
t.Fatal("media relay accepted a session beyond its configured capacity")
}
upstreamAPI := newTestWebRTCAPI(t)
upstream, errUpstream := upstreamAPI.NewPeerConnection(webrtc.Configuration{})
if errUpstream != nil {
t.Fatalf("create upstream PeerConnection: %v", errUpstream)
}
defer closeTestPeerConnection(t, upstream)
upstreamDone := make(chan struct{})
defer close(upstreamDone)
upstreamDataChannels := make(chan *webrtc.DataChannel, 1)
upstreamMessages := make(chan webrtc.DataChannelMessage, 4)
upstream.OnDataChannel(func(channel *webrtc.DataChannel) {
if channel.Label() != realtimeDataChannelLabel {
return
}
channel.OnMessage(func(message webrtc.DataChannelMessage) {
message.Data = append([]byte(nil), message.Data...)
upstreamMessages <- message
})
upstreamDataChannels <- channel
})
upstreamAudioMessages := make(chan []byte, 1)
upstream.OnTrack(func(track *webrtc.TrackRemote, _ *webrtc.RTPReceiver) {
packet, _, errRead := track.ReadRTP()
if errRead == nil {
upstreamAudioMessages <- append([]byte(nil), packet.Payload...)
}
})
if errRemote := upstream.SetRemoteDescription(webrtc.SessionDescription{Type: webrtc.SDPTypeOffer, SDP: relayOffer}); errRemote != nil {
t.Fatalf("set upstream offer: %v", errRemote)
}
upstreamAudio, errTrack := webrtc.NewTrackLocalStaticRTP(opusCodec, "upstream-audio", "upstream")
if errTrack != nil {
t.Fatalf("create upstream audio track: %v", errTrack)
}
upstreamSender, errTrack := upstream.AddTrack(upstreamAudio)
if errTrack != nil {
t.Fatalf("add upstream audio track: %v", errTrack)
}
go drainRTCP("test-upstream", upstreamSender, upstreamDone)
upstreamAnswer := completeAnswer(t, upstream)
downstreamAnswer, errAnswer := session.AcceptUpstreamAnswer(context.Background(), upstreamAnswer)
if errAnswer != nil {
t.Fatalf("accept upstream answer: %v", errAnswer)
}
if errRemote := client.SetRemoteDescription(webrtc.SessionDescription{Type: webrtc.SDPTypeAnswer, SDP: downstreamAnswer}); errRemote != nil {
t.Fatalf("set client answer: %v", errRemote)
}
upstreamData := receiveDataChannel(t, upstreamDataChannels)
waitDataChannelOpen(t, clientData)
waitDataChannelOpen(t, upstreamData)
if errSend := clientData.SendText("from-client"); errSend != nil {
t.Fatalf("send client DataChannel message: %v", errSend)
}
if got := receiveDataMessage(t, upstreamMessages); !got.IsString || string(got.Data) != "from-client" {
t.Fatalf("upstream DataChannel message = %#v, want text from-client", got)
}
if errSend := upstreamData.SendText("from-upstream"); errSend != nil {
t.Fatalf("send upstream DataChannel message: %v", errSend)
}
if got := receiveDataMessage(t, clientMessages); !got.IsString || string(got.Data) != "from-upstream" {
t.Fatalf("client DataChannel message = %#v, want text from-upstream", got)
}
if errSend := clientData.Send([]byte{0x01, 0x02, 0x03}); errSend != nil {
t.Fatalf("send client binary DataChannel message: %v", errSend)
}
if got := receiveDataMessage(t, upstreamMessages); got.IsString || string(got.Data) != string([]byte{0x01, 0x02, 0x03}) {
t.Fatalf("upstream binary DataChannel message = %#v", got)
}
clientPayload := []byte{0xf8, 0xff, 0xfe}
sendTestRTP(t, clientAudio, clientPayload, upstreamAudioMessages)
upstreamPayload := []byte{0xf8, 0xfe, 0xfd}
sendTestRTP(t, upstreamAudio, upstreamPayload, clientAudioMessages)
}
func TestIsPublicRemoteIP(t *testing.T) {
for rawIP, want := range map[string]bool{
"8.8.8.8": true,
"2001:4860::1": true,
"127.0.0.1": false,
"10.0.0.1": false,
"169.254.1.1": false,
"224.0.0.1": false,
"::1": false,
"fc00::1": false,
"fe80::1": false,
"ff02::1": false,
"0.0.0.0": false,
} {
if got := isPublicRemoteIP(net.ParseIP(rawIP)); got != want {
t.Errorf("isPublicRemoteIP(%q) = %t, want %t", rawIP, got, want)
}
}
if isPublicRemoteIP(nil) {
t.Fatal("isPublicRemoteIP(nil) = true, want false")
}
}
func newTestWebRTCAPI(t *testing.T) *webrtc.API {
t.Helper()
mediaEngine := &webrtc.MediaEngine{}
if errRegister := mediaEngine.RegisterCodec(webrtc.RTPCodecParameters{
RTPCodecCapability: opusCodec,
PayloadType: 111,
}, webrtc.RTPCodecTypeAudio); errRegister != nil {
t.Fatalf("register test Opus codec: %v", errRegister)
}
interceptorRegistry := &interceptor.Registry{}
if errRegister := webrtc.RegisterDefaultInterceptors(mediaEngine, interceptorRegistry); errRegister != nil {
t.Fatalf("register test interceptors: %v", errRegister)
}
return webrtc.NewAPI(
webrtc.WithMediaEngine(mediaEngine),
webrtc.WithInterceptorRegistry(interceptorRegistry),
)
}
func completeOffer(t *testing.T, connection *webrtc.PeerConnection) string {
t.Helper()
gatherComplete := webrtc.GatheringCompletePromise(connection)
offer, errOffer := connection.CreateOffer(nil)
if errOffer != nil {
t.Fatalf("create offer: %v", errOffer)
}
if errLocal := connection.SetLocalDescription(offer); errLocal != nil {
t.Fatalf("set local offer: %v", errLocal)
}
select {
case <-gatherComplete:
case <-time.After(5 * time.Second):
t.Fatal("offer ICE gathering did not complete")
}
return connection.LocalDescription().SDP
}
func completeAnswer(t *testing.T, connection *webrtc.PeerConnection) string {
t.Helper()
gatherComplete := webrtc.GatheringCompletePromise(connection)
answer, errAnswer := connection.CreateAnswer(nil)
if errAnswer != nil {
t.Fatalf("create answer: %v", errAnswer)
}
if errLocal := connection.SetLocalDescription(answer); errLocal != nil {
t.Fatalf("set local answer: %v", errLocal)
}
select {
case <-gatherComplete:
case <-time.After(5 * time.Second):
t.Fatal("answer ICE gathering did not complete")
}
return connection.LocalDescription().SDP
}
func waitDataChannelOpen(t *testing.T, channel *webrtc.DataChannel) {
t.Helper()
deadline := time.Now().Add(5 * time.Second)
for time.Now().Before(deadline) {
if channel.ReadyState() == webrtc.DataChannelStateOpen {
return
}
time.Sleep(10 * time.Millisecond)
}
t.Fatalf("DataChannel %q did not open", channel.Label())
}
func receiveDataChannel(t *testing.T, channels <-chan *webrtc.DataChannel) *webrtc.DataChannel {
t.Helper()
select {
case channel := <-channels:
return channel
case <-time.After(5 * time.Second):
t.Fatal("upstream DataChannel was not created")
return nil
}
}
func receiveDataMessage(t *testing.T, messages <-chan webrtc.DataChannelMessage) webrtc.DataChannelMessage {
t.Helper()
select {
case message := <-messages:
return message
case <-time.After(5 * time.Second):
t.Fatal("DataChannel message was not relayed")
return webrtc.DataChannelMessage{}
}
}
func sendTestRTP(t *testing.T, track *webrtc.TrackLocalStaticRTP, payload []byte, received <-chan []byte) {
t.Helper()
for sequence := uint16(1); sequence <= 25; sequence++ {
packet := &rtp.Packet{
Header: rtp.Header{
Version: 2,
PayloadType: 111,
SequenceNumber: sequence,
Timestamp: uint32(sequence) * 960,
SSRC: 1234,
},
Payload: payload,
}
if errWrite := track.WriteRTP(packet); errWrite != nil {
t.Fatalf("write test RTP: %v", errWrite)
}
select {
case got := <-received:
if string(got) != string(payload) {
t.Fatalf("relayed RTP payload = %v, want %v", got, payload)
}
return
case <-time.After(20 * time.Millisecond):
}
}
t.Fatal("RTP packet was not relayed")
}
func closeTestPeerConnection(t *testing.T, connection *webrtc.PeerConnection) {
t.Helper()
if errClose := connection.Close(); errClose != nil {
t.Errorf("close test PeerConnection: %v", errClose)
}
}

View File

@@ -45,9 +45,17 @@ type liveSession struct {
authID string
model string
homeSelection *auth.HomeDispatchSelection
media mediaRelaySession
resources *liveSessionResources
token uint64
}
type liveSessionResources struct {
mu sync.Mutex
closed bool
closers []func() error
}
type storedSession struct {
session liveSession
claimed bool
@@ -76,12 +84,15 @@ func newSessionStore() *sessionStore {
}
}
func (s *sessionStore) put(callID string, session liveSession) {
func (s *sessionStore) put(callID string, session liveSession) liveSession {
if s == nil || !callIDPattern.MatchString(callID) {
endHomeSelection(session, "invalid_call_id")
return
endLiveSession(session, "invalid_call_id")
return liveSession{}
}
if session.resources == nil {
session.resources = &liveSessionResources{}
}
s.mu.Lock()
s.next++
session.callID = callID
@@ -98,10 +109,19 @@ func (s *sessionStore) put(callID string, session liveSession) {
if previous.timer != nil {
previous.timer.Stop()
}
if previous.session.resources != nil && previous.session.resources != session.resources {
previous.session.resources.close()
}
if previous.session.media != nil && previous.session.media != session.media {
if errClose := previous.session.media.Close(); errClose != nil {
log.WithError(errClose).Debug("codex live media: close replaced session")
}
}
if previous.session.homeSelection != session.homeSelection {
endHomeSelection(previous.session, "session_replaced")
}
}
return session
}
func (s *sessionStore) claim(callID string) (liveSession, sessionClaim) {
@@ -144,7 +164,7 @@ func (s *sessionStore) release(session liveSession) {
func (s *sessionStore) complete(session liveSession, reason string) {
if s == nil || session.callID == "" {
endHomeSelection(session, reason)
endLiveSession(session, reason)
return
}
s.mu.Lock()
@@ -158,7 +178,26 @@ func (s *sessionStore) complete(session liveSession, reason string) {
entry.timer.Stop()
}
s.mu.Unlock()
endHomeSelection(entry.session, reason)
endLiveSession(entry.session, reason)
}
func (s *sessionStore) closeAll(reason string) {
if s == nil {
return
}
s.mu.Lock()
entries := make([]*storedSession, 0, len(s.sessions))
for callID, entry := range s.sessions {
delete(s.sessions, callID)
if entry.timer != nil {
entry.timer.Stop()
}
entries = append(entries, entry)
}
s.mu.Unlock()
for _, entry := range entries {
endLiveSession(entry.session, reason)
}
}
func (s *sessionStore) expiryDuration() time.Duration {
@@ -177,7 +216,7 @@ func (s *sessionStore) expire(callID string, token uint64) {
}
delete(s.sessions, callID)
s.mu.Unlock()
endHomeSelection(entry.session, "session_expired")
endLiveSession(entry.session, "session_expired")
}
func (s *sessionStore) peek(callID string) (liveSession, bool) {
@@ -193,12 +232,65 @@ func (s *sessionStore) peek(callID string) (liveSession, bool) {
return entry.session, true
}
func endLiveSession(session liveSession, reason string) {
if session.resources != nil {
session.resources.close()
}
if session.media != nil {
if errClose := session.media.Close(); errClose != nil {
log.WithError(errClose).Debug("codex live media: close stored session")
}
}
endHomeSelection(session, reason)
}
func endHomeSelection(session liveSession, reason string) {
if session.homeSelection != nil {
session.homeSelection.End(reason)
}
}
func (r *liveSessionResources) add(closers ...func() error) {
if r == nil {
return
}
r.mu.Lock()
if !r.closed {
r.closers = append(r.closers, closers...)
r.mu.Unlock()
return
}
r.mu.Unlock()
closeSessionResources(closers)
}
func (r *liveSessionResources) close() {
if r == nil {
return
}
r.mu.Lock()
if r.closed {
r.mu.Unlock()
return
}
r.closed = true
closers := r.closers
r.closers = nil
r.mu.Unlock()
closeSessionResources(closers)
}
func closeSessionResources(closers []func() error) {
for _, closer := range closers {
if closer == nil {
continue
}
if errClose := closer(); errClose != nil && !isNormalWebsocketClose(errClose) {
log.WithError(errClose).Debug("codex live: close session resource")
}
}
}
type sidebandStyle int
const (
@@ -356,6 +448,9 @@ func (h *Handler) HandleSideband(c *gin.Context) {
} else {
defer func() { _ = closeDownstream() }()
}
if session.resources != nil {
session.resources.add(closeUpstream, closeDownstream)
}
consumeSession = true
if errRelay := relayWebsockets(downstream, upstream); errRelay != nil && !isNormalWebsocketClose(errRelay) {

View File

@@ -0,0 +1,63 @@
package config
import (
"errors"
"fmt"
"net"
"net/url"
"strings"
)
// DefaultCodexLiveMediaMaxSessions is the default in-process media session limit.
const DefaultCodexLiveMediaMaxSessions = 32
// EffectiveMaxSessions returns the configured media session limit.
func (c CodexLiveMediaRelayConfig) EffectiveMaxSessions() int {
if c.MaxSessions > 0 {
return c.MaxSessions
}
return DefaultCodexLiveMediaMaxSessions
}
// Validate verifies the Codex Live media relay configuration.
func (c CodexLiveMediaRelayConfig) Validate() error {
if !c.Enabled {
return nil
}
if c.MaxSessions < 0 {
return errors.New("codex.live-media-relay.max-sessions must not be negative")
}
if publicIP := strings.TrimSpace(c.PublicIP); publicIP != "" && net.ParseIP(publicIP) == nil {
return fmt.Errorf("codex.live-media-relay.public-ip is invalid: %q", publicIP)
}
if (c.UDPPortMin == 0) != (c.UDPPortMax == 0) {
return errors.New("codex.live-media-relay UDP port minimum and maximum must both be set")
}
if c.UDPPortMin > c.UDPPortMax {
return errors.New("codex.live-media-relay.udp-port-min must not exceed udp-port-max")
}
if c.UDPPortMin != 0 {
availablePorts := int(c.UDPPortMax) - int(c.UDPPortMin) + 1
requiredPorts := c.EffectiveMaxSessions() * 2
if availablePorts < requiredPorts {
return fmt.Errorf("codex.live-media-relay UDP range requires at least %d ports for %d sessions", requiredPorts, c.EffectiveMaxSessions())
}
}
for serverIndex, server := range c.ICEServers {
if len(server.URLs) == 0 {
return fmt.Errorf("codex.live-media-relay.ice-servers[%d].urls is required", serverIndex)
}
for _, rawURL := range server.URLs {
parsed, errParse := url.Parse(strings.TrimSpace(rawURL))
if errParse != nil || parsed.Scheme == "" {
return fmt.Errorf("codex.live-media-relay.ice-servers[%d] contains an invalid URL", serverIndex)
}
switch strings.ToLower(parsed.Scheme) {
case "stun", "stuns", "turn", "turns":
default:
return fmt.Errorf("codex.live-media-relay.ice-servers[%d] uses unsupported scheme %q", serverIndex, parsed.Scheme)
}
}
}
return nil
}

View File

@@ -0,0 +1,92 @@
package config
import (
"encoding/json"
"strings"
"testing"
"gopkg.in/yaml.v3"
)
func TestCodexLiveMediaRelayConfigParsesAndValidates(t *testing.T) {
var cfg Config
raw := []byte(`codex:
live-media-relay:
enabled: true
max-sessions: 64
allow-private-remote-ips: true
public-ip: "203.0.113.10"
udp-port-min: 40000
udp-port-max: 40150
ice-servers:
- urls: ["stun:stun.example.com:3478"]
- urls: ["turn:turn.example.com:3478?transport=udp"]
username: "relay-user"
credential: "relay-secret"
`)
if errUnmarshal := yaml.Unmarshal(raw, &cfg); errUnmarshal != nil {
t.Fatalf("unmarshal Codex Live media relay config: %v", errUnmarshal)
}
relay := cfg.Codex.LiveMediaRelay
if !relay.Enabled || relay.MaxSessions != 64 || !relay.AllowPrivateRemoteIPs || relay.PublicIP != "203.0.113.10" {
t.Fatalf("parsed media relay = %#v", relay)
}
if relay.UDPPortMin != 40000 || relay.UDPPortMax != 40150 {
t.Fatalf("parsed UDP range = %d-%d", relay.UDPPortMin, relay.UDPPortMax)
}
if len(relay.ICEServers) != 2 || relay.ICEServers[1].Credential != "relay-secret" {
t.Fatalf("parsed ICE servers = %#v", relay.ICEServers)
}
if errValidate := relay.Validate(); errValidate != nil {
t.Fatalf("Validate() error = %v", errValidate)
}
encoded, errMarshal := json.Marshal(relay)
if errMarshal != nil {
t.Fatalf("marshal media relay config: %v", errMarshal)
}
if strings.Contains(string(encoded), "relay-secret") || strings.Contains(string(encoded), "credential") {
t.Fatalf("JSON media relay config leaked TURN credential: %s", encoded)
}
}
func TestCodexLiveMediaRelayConfigRejectsInvalidValues(t *testing.T) {
for name, relay := range map[string]CodexLiveMediaRelayConfig{
"negative session limit": {
Enabled: true,
MaxSessions: -1,
},
"invalid public IP": {
Enabled: true,
PublicIP: "not-an-ip",
},
"partial UDP range": {
Enabled: true,
UDPPortMin: 40000,
},
"reversed UDP range": {
Enabled: true,
UDPPortMin: 40100,
UDPPortMax: 40000,
},
"undersized UDP range": {
Enabled: true,
MaxSessions: 2,
UDPPortMin: 40000,
UDPPortMax: 40002,
},
"missing ICE URLs": {
Enabled: true,
ICEServers: []CodexLiveICEServer{{Username: "user"}},
},
"unsupported ICE URL": {
Enabled: true,
ICEServers: []CodexLiveICEServer{{URLs: []string{"https://example.com"}}},
},
} {
t.Run(name, func(t *testing.T) {
if errValidate := relay.Validate(); errValidate == nil {
t.Fatal("Validate() accepted invalid media relay config")
}
})
}
}

View File

@@ -293,6 +293,26 @@ type CodexConfig struct {
IdentityConfuse bool `yaml:"identity-confuse" json:"identity-confuse"`
// OptimizeMultiAgentV2 optimizes official Codex multi-agent requests.
OptimizeMultiAgentV2 bool `yaml:"optimize-multi-agent-v2" json:"optimize-multi-agent-v2"`
// LiveMediaRelay terminates and relays Codex Live WebRTC media in this process.
LiveMediaRelay CodexLiveMediaRelayConfig `yaml:"live-media-relay" json:"live-media-relay"`
}
// CodexLiveMediaRelayConfig configures the in-process Codex Live WebRTC gateway.
type CodexLiveMediaRelayConfig struct {
Enabled bool `yaml:"enabled" json:"enabled"`
MaxSessions int `yaml:"max-sessions" json:"max-sessions"`
AllowPrivateRemoteIPs bool `yaml:"allow-private-remote-ips" json:"allow-private-remote-ips"`
PublicIP string `yaml:"public-ip" json:"public-ip"`
UDPPortMin uint16 `yaml:"udp-port-min" json:"udp-port-min"`
UDPPortMax uint16 `yaml:"udp-port-max" json:"udp-port-max"`
ICEServers []CodexLiveICEServer `yaml:"ice-servers" json:"ice-servers"`
}
// CodexLiveICEServer configures a STUN or TURN server for the media relay.
type CodexLiveICEServer struct {
URLs []string `yaml:"urls" json:"urls"`
Username string `yaml:"username" json:"username"`
Credential string `yaml:"credential" json:"-"`
}
// TLSConfig holds HTTPS server settings.
@@ -785,6 +805,9 @@ func LoadConfigOptional(configFile string, optional bool) (*Config, error) {
if errValidate := cfg.CredentialInFlight.Validate(); errValidate != nil {
return nil, errValidate
}
if errValidate := cfg.Codex.LiveMediaRelay.Validate(); errValidate != nil {
return nil, errValidate
}
// Hash remote management key if plaintext is detected (nested)
// We consider a value to be already hashed if it looks like a bcrypt hash ($2a$, $2b$, or $2y$ prefix).