From bda79b21bb9484bc95875c709ae777ce46186a06 Mon Sep 17 00:00:00 2001 From: Luis Pater Date: Sat, 25 Jul 2026 05:11:28 +0800 Subject: [PATCH] feat(live): relay realtime WebRTC media --- config.example.yaml | 22 + go.mod | 18 + go.sum | 41 +- internal/api/server.go | 28 +- internal/client/codex/live/live.go | 223 +++++++- internal/client/codex/live/live_test.go | 351 ++++++++++++- internal/client/codex/live/media.go | 618 +++++++++++++++++++++++ internal/client/codex/live/media_test.go | 295 +++++++++++ internal/client/codex/live/sideband.go | 107 +++- internal/config/codex_live.go | 63 +++ internal/config/codex_live_test.go | 92 ++++ internal/config/config.go | 23 + 12 files changed, 1849 insertions(+), 32 deletions(-) create mode 100644 internal/client/codex/live/media.go create mode 100644 internal/client/codex/live/media_test.go create mode 100644 internal/config/codex_live.go create mode 100644 internal/config/codex_live_test.go diff --git a/config.example.yaml b/config.example.yaml index ee981bd40..58eaa03e5 100644 --- a/config.example.yaml +++ b/config.example.yaml @@ -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 diff --git a/go.mod b/go.mod index 7264ba4a1..b1544a510 100644 --- a/go.mod +++ b/go.mod @@ -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 ( diff --git a/go.sum b/go.sum index 0f9a92e29..3d3458d60 100644 --- a/go.sum +++ b/go.sum @@ -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= diff --git a/internal/api/server.go b/internal/api/server.go index f23fe6a1c..f263e82fb 100644 --- a/internal/api/server.go +++ b/internal/api/server.go @@ -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) diff --git a/internal/client/codex/live/live.go b/internal/client/codex/live/live.go index 74557b6ac..ac2949258 100644 --- a/internal/client/codex/live/live.go +++ b/internal/client/codex/live/live.go @@ -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 { diff --git a/internal/client/codex/live/live_test.go b/internal/client/codex/live/live_test.go index e1f4585ed..931cb8eff 100644 --- a/internal/client/codex/live/live_test.go +++ b/internal/client/codex/live/live_test.go @@ -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) diff --git a/internal/client/codex/live/media.go b/internal/client/codex/live/media.go new file mode 100644 index 000000000..ba7b77ad1 --- /dev/null +++ b/internal/client/codex/live/media.go @@ -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) + } +} diff --git a/internal/client/codex/live/media_test.go b/internal/client/codex/live/media_test.go new file mode 100644 index 000000000..a972500f0 --- /dev/null +++ b/internal/client/codex/live/media_test.go @@ -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) + } +} diff --git a/internal/client/codex/live/sideband.go b/internal/client/codex/live/sideband.go index ab57db6bb..f16ecb924 100644 --- a/internal/client/codex/live/sideband.go +++ b/internal/client/codex/live/sideband.go @@ -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) { diff --git a/internal/config/codex_live.go b/internal/config/codex_live.go new file mode 100644 index 000000000..611ed5ed3 --- /dev/null +++ b/internal/config/codex_live.go @@ -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 +} diff --git a/internal/config/codex_live_test.go b/internal/config/codex_live_test.go new file mode 100644 index 000000000..d9f93e13b --- /dev/null +++ b/internal/config/codex_live_test.go @@ -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") + } + }) + } +} diff --git a/internal/config/config.go b/internal/config/config.go index 64548297f..c0e27b312 100644 --- a/internal/config/config.go +++ b/internal/config/config.go @@ -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).