mirror of
https://github.com/router-for-me/CLIProxyAPI.git
synced 2026-09-03 06:35:00 +08:00
feat(live): relay realtime WebRTC media
This commit is contained in:
@@ -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
18
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 (
|
||||
|
||||
41
go.sum
41
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=
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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)
|
||||
|
||||
618
internal/client/codex/live/media.go
Normal file
618
internal/client/codex/live/media.go
Normal 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)
|
||||
}
|
||||
}
|
||||
295
internal/client/codex/live/media_test.go
Normal file
295
internal/client/codex/live/media_test.go
Normal 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)
|
||||
}
|
||||
}
|
||||
@@ -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) {
|
||||
|
||||
63
internal/config/codex_live.go
Normal file
63
internal/config/codex_live.go
Normal 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
|
||||
}
|
||||
92
internal/config/codex_live_test.go
Normal file
92
internal/config/codex_live_test.go
Normal 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")
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -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).
|
||||
|
||||
Reference in New Issue
Block a user