diff --git a/pkg/sip/inbound.go b/pkg/sip/inbound.go index 95fded36..7dcb94d3 100644 --- a/pkg/sip/inbound.go +++ b/pkg/sip/inbound.go @@ -45,7 +45,6 @@ import ( lksip "github.com/livekit/protocol/sip" "github.com/livekit/protocol/utils/traceid" "github.com/livekit/psrpc" - lksdk "github.com/livekit/server-sdk-go/v2" "github.com/livekit/sipgo/sip" "github.com/livekit/sip/pkg/config" @@ -718,7 +717,6 @@ type inboundCall struct { lkRoom RoomInterface // LiveKit room; only active after correct pin is entered callDur func() time.Duration joinDur func() time.Duration - forwardDTMF atomic.Bool done atomic.Bool started core.Fuse stats Stats @@ -978,7 +976,7 @@ func (c *inboundCall) handleInvite(ctx context.Context, tid traceid.ID, req *sip // Start this timer right after the Accept. ackTimeout = time.After(inviteOkAckLateTimeout) } - if old := c.audioOut.Swap(c.media.GetAudioWriter()); old != nil { + if old := c.audioOut.Swap(c.media.GetOutboundAudioWriter()); old != nil { c.log().Warnw("unexpected audio out writer", nil) old.Close() } @@ -1015,6 +1013,7 @@ func (c *inboundCall) handleInvite(ctx context.Context, tid traceid.ID, req *sip return err // already sent a response } } + p := &disp.Room.Participant p.Attributes = HeadersToAttrs(p.Attributes, disp.HeadersToAttributes, disp.IncludeHeaders, c.cc, nil) if disp.MaxCallDuration <= 0 || disp.MaxCallDuration > maxCallDuration { @@ -1113,25 +1112,33 @@ func (c *inboundCall) waitForCallEnd(ctx context.Context, ackReceived <-chan str } } -// TODO(alexfish): Update Room so that we don't need this adapater. -type dtmfEventWriter struct { - handler func(msg *livekit.SipDTMF) +type pinDTMFWriter struct { + dtmfEvents chan<- dtmf.Event } -func (w *dtmfEventWriter) String() string { - return "dtmfEventWriter" +func (w *pinDTMFWriter) String() string { + return "pinDTMFWriter" } -func (w *dtmfEventWriter) SampleRate() int { +func (w *pinDTMFWriter) SampleRate() int { return dtmf.SampleRate } -func (w *dtmfEventWriter) Close() error { +func (w *pinDTMFWriter) Close() error { return nil } -func (w *dtmfEventWriter) WriteSample(sample *livekit.SipDTMF) error { - w.handler(sample) +func (w *pinDTMFWriter) WriteSample(msg *livekit.SipDTMF) error { + if msg == nil { + return nil + } + + event := dtmfEventFromSipDTMF(msg) + // We should have enough buffer here. + select { + case w.dtmfEvents <- event: + default: + } return nil } @@ -1173,13 +1180,20 @@ func (c *inboundCall) runMediaConn(tid traceid.ID, offerData []byte, mconf *sipM c.mon.SDPSize(len(answerData), false) c.log().Debugw("SDP answer", "sdp", string(answerData)) - mp.WriteDTMFTo(&dtmfEventWriter{handler: c.handleDTMF}) + if old := mp.WriteInboundDTMFTo(&pinDTMFWriter{c.dtmf}); old != nil { + c.log().Warnw("media port has unexpected inbound DTMF writer", nil) + } // Must be set earlier to send the pin prompts. - if w := c.lkRoom.SwapOutput(c.audioOut); w != nil { - _ = w.Close() + if old := c.lkRoom.WriteOutboundAudioTo(c.audioOut); old != nil { + c.log().Warnw("room has unexpected outbound audio writer", nil) + old.Close() + } + + if old := c.lkRoom.WriteOutboundDTMFTo(c.media.GetOutboundDTMFWriter()); old != nil { + c.log().Warnw("room has unexpected outbound audio DTMF writer", nil) + old.Close() } - c.lkRoom.SetDTMFOutput(c.media.GetDTMFWriter()) audio := mp.NegotiatedAudio() if audio == nil { @@ -1576,7 +1590,6 @@ func (c *inboundCall) createLiveKitParticipant(ctx context.Context, rconf RoomCo partConf.Attributes[k] = v } partConf.Attributes[livekit.AttrSIPCallStatus] = status.Attribute() - c.forwardDTMF.Store(true) select { case <-ctx.Done(): return ctx.Err() @@ -1607,16 +1620,23 @@ func (c *inboundCall) createLiveKitParticipant(ctx context.Context, rconf RoomCo func (c *inboundCall) publishTrack(features []livekit.SIPFeature, featureFlags map[string]string) error { defer c.mon.StageDurTimer("track-publish")() - local, err := c.lkRoom.NewParticipantTrack(RoomSampleRate) + inboundAudio, err := c.lkRoom.GetInboundAudioWriter() if err != nil { _ = c.lkRoom.Close() return err } if audioInProcessor := c.s.handler.GetMediaProcessor(features, featureFlags, string(c.cc.ID()), MediaProcessorOpts{InputSampleRate: RoomSampleRate}); audioInProcessor != nil { - local = audioInProcessor(local) + inboundAudio = audioInProcessor(inboundAudio) + } + if old := c.media.WriteInboundAudioTo(inboundAudio); old != nil { + c.log().Warnw("media port has unexpected inbound audio writer", nil) + old.Close() + } + if old := c.media.WriteInboundDTMFTo(c.lkRoom.GetInboundDTMFWriter()); old != nil { + c.log().Warnw("media port has unexpected inbound dtmf writer", nil) + old.Close() } - c.media.WriteAudioTo(local) return nil } @@ -1657,11 +1677,7 @@ func (c *inboundCall) playAudio(ctx context.Context, frames []msdk.PCM16Sample) _ = msdk.PlayAudio[msdk.PCM16Sample](ctx, t, rtp.DefFrameDur, frames) } -func (c *inboundCall) handleDTMF(msg *livekit.SipDTMF) { - if msg == nil { - return - } - +func dtmfEventFromSipDTMF(msg *livekit.SipDTMF) dtmf.Event { code := byte(msg.Code) digit := byte(0) if len(msg.Digit) == 1 { @@ -1669,23 +1685,10 @@ func (c *inboundCall) handleDTMF(msg *livekit.SipDTMF) { } else { digit = dtmf.CodeToChar(code) } - event := dtmf.Event{ + return dtmf.Event{ Code: code, Digit: digit, } - - if c.forwardDTMF.Load() { - _ = c.lkRoom.SendData(&livekit.SipDTMF{ - Code: uint32(code), - Digit: string([]byte{digit}), - }, lksdk.WithDataPublishReliable(true)) - return - } - // We should have enough buffer here. - select { - case c.dtmf <- event: - default: - } } func (c *inboundCall) transferCall(ctx context.Context, transferTo string, headers map[string]string, dialtone bool) (retErr error) { @@ -1701,14 +1704,13 @@ func (c *inboundCall) transferCall(ctx context.Context, transferTo string, heade rctx, rcancel := context.WithCancel(ctx) defer rcancel() - // mute the room audio to the SIP participant - w := c.lkRoom.SwapOutput(nil) + // Mute the room audio to the SIP participant. + // Skip closing the existing writer, which is c.audioOut. + _ = c.lkRoom.WriteOutboundAudioTo(nil) defer func() { if retErr != nil && !c.done.Load() { - c.lkRoom.SwapOutput(w) - } else if w != nil { - w.Close() + c.lkRoom.WriteOutboundAudioTo(c.audioOut) } }() diff --git a/pkg/sip/media_port.go b/pkg/sip/media_port.go index 64c6dee9..4e6f4f91 100644 --- a/pkg/sip/media_port.go +++ b/pkg/sip/media_port.go @@ -390,20 +390,21 @@ func (o *MediaOptions) ApplyDefaults() { } } -type MediaSegment interface { - GetAudioWriter() msdk.PCM16Writer - WriteAudioTo(w msdk.PCM16Writer) msdk.PCM16Writer - GetDTMFWriter() msdk.WriteCloser[*livekit.SipDTMF] - WriteDTMFTo(w msdk.WriteCloser[*livekit.SipDTMF]) msdk.WriteCloser[*livekit.SipDTMF] -} - // MediaPort is the insulated media-plane API: UDP/RTP to the wire, SDP negotiation, // and audio/DTMF endpoints. It does not know about calls, rooms, or SIP dialogs. type MediaPort interface { Close() CloseWait() - MediaSegment + // GetOutboundAudioWriter returns the LK room -> SIP writer. + GetOutboundAudioWriter() msdk.PCM16Writer + // GetOutboundDTMFWriter returns the LK room -> SIP DTMF writer. + GetOutboundDTMFWriter() msdk.WriteCloser[*livekit.SipDTMF] + + // WriteInboundAudioTo tells the MediaSegment where to write inbound SIP audio. + WriteInboundAudioTo(w msdk.PCM16Writer) msdk.PCM16Writer + // WriteInboundDTMFTo tells the MediaSegment where to write inbound SIP DTMF. + WriteInboundDTMFTo(w msdk.WriteCloser[*livekit.SipDTMF]) msdk.WriteCloser[*livekit.SipDTMF] // If there is no offer, this generates an offer. // If there is an offer, this simply returns the SDP of that offer. @@ -731,20 +732,20 @@ func (p *mediaPort) reportPeerCodecs(d sdp.MediaDesc) { // Plumbing -func (p *mediaPort) GetAudioWriter() msdk.PCM16Writer { +func (p *mediaPort) GetOutboundAudioWriter() msdk.PCM16Writer { return p.audioOut } -// WriteAudioTo sets audio writer that will receive decoded PCM from incoming RTP packets. -func (p *mediaPort) WriteAudioTo(w msdk.PCM16Writer) msdk.PCM16Writer { +// WriteInboundAudioTo sets audio writer that will receive decoded PCM from incoming RTP packets. +func (p *mediaPort) WriteInboundAudioTo(w msdk.PCM16Writer) msdk.PCM16Writer { return p.audioIn.Swap(w) } -func (p *mediaPort) GetDTMFWriter() msdk.WriteCloser[*livekit.SipDTMF] { +func (p *mediaPort) GetOutboundDTMFWriter() msdk.WriteCloser[*livekit.SipDTMF] { return &p.dtmfOut } -func (p *mediaPort) WriteDTMFTo(w msdk.WriteCloser[*livekit.SipDTMF]) msdk.WriteCloser[*livekit.SipDTMF] { +func (p *mediaPort) WriteInboundDTMFTo(w msdk.WriteCloser[*livekit.SipDTMF]) msdk.WriteCloser[*livekit.SipDTMF] { return p.dtmfIn.Swap(w) } diff --git a/pkg/sip/media_port_negotiation_test.go b/pkg/sip/media_port_negotiation_test.go index 5a984237..9c5b877c 100644 --- a/pkg/sip/media_port_negotiation_test.go +++ b/pkg/sip/media_port_negotiation_test.go @@ -72,7 +72,7 @@ func roomFrame() msdk.PCM16Sample { // writeFrames pushes room audio into the port. Write errors are ignored: the in-memory // UDP pipe is bounded, and a peer that is mid-renegotiation may not be draining it. func writeFrames(m *mediaPort, frames int) { - w := m.GetAudioWriter() + w := m.GetOutboundAudioWriter() frame := roomFrame() for range frames { _ = w.WriteSample(frame) @@ -169,7 +169,7 @@ func TestMediaPortRenegotiation(t *testing.T) { m1, m2 := newMediaPair(t, nil, nil, "") recv2 := &recvBuffer{} - m2.WriteAudioTo(recv2) + m2.WriteInboundAudioTo(recv2) requireAudioFlows(t, m1, recv2) for range 3 { @@ -203,7 +203,7 @@ func TestMediaPortRenegotiation(t *testing.T) { require.Equal(t, g711.ULawSDPNameAndRate, answerCodec(t, answerData)) recv2 := &recvBuffer{} - m2.WriteAudioTo(recv2) + m2.WriteInboundAudioTo(recv2) requireAudioFlows(t, m1, recv2) // G722 samples at 16k, so the encode leaf changes sample rate under the same @@ -247,9 +247,9 @@ func TestMediaPortHold(t *testing.T) { m1, m2 := newMediaPair(t, nil, nil, "") recv1 := &recvBuffer{} - m1.WriteAudioTo(recv1) + m1.WriteInboundAudioTo(recv1) recv2 := &recvBuffer{} - m2.WriteAudioTo(recv2) + m2.WriteInboundAudioTo(recv2) // Baseline: m1 sends to m2. require.NotNil(t, m1.audioOut.Get()) @@ -268,7 +268,7 @@ func TestMediaPortHold(t *testing.T) { assert.True(t, dst.Addr().IsUnspecified(), "held port kept a destination: %v", dst) } assert.Nil(t, m1.audioOut.Get(), "held port still accepts room audio") - assert.NoError(t, m1.GetAudioWriter().WriteSample(roomFrame())) + assert.NoError(t, m1.GetOutboundAudioWriter().WriteSample(roomFrame())) sent := recv2.count() writeFrames(m1, 10) diff --git a/pkg/sip/media_port_pending_test.go b/pkg/sip/media_port_pending_test.go index 50c006ea..a8575934 100644 --- a/pkg/sip/media_port_pending_test.go +++ b/pkg/sip/media_port_pending_test.go @@ -197,11 +197,11 @@ func TestMediaPort(t *testing.T) { var aliceRecvBuf msdk.PCM16Sample aliceHandler := msdk.NewPCM16BufferWriter(&aliceRecvBuf, testRate) - alicePort.WriteAudioTo(aliceHandler) + alicePort.WriteInboundAudioTo(aliceHandler) var bobRecvBuf msdk.PCM16Sample bobHandler := msdk.NewPCM16BufferWriter(&bobRecvBuf, testRate) - bobPort.WriteAudioTo(bobHandler) + bobPort.WriteInboundAudioTo(bobHandler) aliceToBob := alicePort.GetAudioWriter() bobToAlice := bobPort.GetAudioWriter() diff --git a/pkg/sip/media_port_test.go b/pkg/sip/media_port_test.go index 0317fc80..64866e88 100644 --- a/pkg/sip/media_port_test.go +++ b/pkg/sip/media_port_test.go @@ -299,7 +299,7 @@ func newMediaPairWithAddr(t testing.TB, ip1, ip2 netip.Addr, opt1, opt2 *MediaOp // TODO(port-refactor): the encode chain gained the always-on 48k resampler, refresh // the expected string once the package builds and this can be run. - // w2 := m2.GetAudioWriter() + // w2 := m2.GetOutboundAudioWriter() // require.Equal(t, "Switch(16000) -> LatencyEntry -> G722(encode) -> ByteEncoder(16000) -> StatsWriter(G722/8000) -> LatencyExit -> RTPWriteStream(1.1.1.1:10000)", w2.String()) return m1, m2 @@ -339,7 +339,7 @@ func TestMediaTimeout(t *testing.T) { MediaTimeout: timeout, }, nil, codec) - w2 := m2.GetAudioWriter() + w2 := m2.GetOutboundAudioWriter() err := w2.WriteSample(msdk.PCM16Sample{0, 0}) require.NoError(t, err) @@ -362,7 +362,7 @@ func TestMediaTimeout(t *testing.T) { MediaTimeout: timeout, }, nil, codec) - w2 := m2.GetAudioWriter() + w2 := m2.GetOutboundAudioWriter() for i := 0; i < 10; i++ { err := w2.WriteSample(msdk.PCM16Sample{0, 0}) @@ -382,7 +382,7 @@ func TestMediaTimeout(t *testing.T) { MediaTimeout: timeout, }, nil, codec) - w2 := m2.GetAudioWriter() + w2 := m2.GetOutboundAudioWriter() for i := 0; i < 5; i++ { err := w2.WriteSample(msdk.PCM16Sample{0, 0}) @@ -440,7 +440,7 @@ func TestMediaTimeout(t *testing.T) { MediaTimeout: timeout, }, nil, codec) - w2 := m2.GetAudioWriter() + w2 := m2.GetOutboundAudioWriter() for i := 0; i < 5; i++ { err := w2.WriteSample(msdk.PCM16Sample{0, 0}) @@ -480,7 +480,7 @@ func TestSymmetricRTP(t *testing.T) { newAddr := netip.AddrPortFrom(newIP("9.9.9.9"), 9999) c2.addr = newAddr - err := m2.GetAudioWriter().WriteSample(msdk.PCM16Sample{0, 0}) + err := m2.GetOutboundAudioWriter().WriteSample(msdk.PCM16Sample{0, 0}) require.NoError(t, err) select { @@ -504,7 +504,7 @@ func TestSymmetricRTP(t *testing.T) { newAddr := netip.AddrPortFrom(newIP("9.9.9.9"), 9999) c2.addr = newAddr - err := m2.GetAudioWriter().WriteSample(msdk.PCM16Sample{0, 0}) + err := m2.GetOutboundAudioWriter().WriteSample(msdk.PCM16Sample{0, 0}) require.NoError(t, err) select { @@ -534,7 +534,7 @@ func TestSymmetricRTP(t *testing.T) { newAddr := netip.AddrPortFrom(newIP("3.3.3.3"), 9999) c2.addr = newAddr - err := m2.GetAudioWriter().WriteSample(msdk.PCM16Sample{0, 0}) + err := m2.GetOutboundAudioWriter().WriteSample(msdk.PCM16Sample{0, 0}) require.NoError(t, err) select { diff --git a/pkg/sip/outbound.go b/pkg/sip/outbound.go index cee0c713..2cf4f36a 100644 --- a/pkg/sip/outbound.go +++ b/pkg/sip/outbound.go @@ -31,14 +31,12 @@ import ( "golang.org/x/exp/maps" msdk "github.com/livekit/media-sdk" - "github.com/livekit/media-sdk/dtmf" "github.com/livekit/media-sdk/tones" "github.com/livekit/protocol/livekit" "github.com/livekit/protocol/logger" "github.com/livekit/protocol/utils/guid" "github.com/livekit/protocol/utils/traceid" "github.com/livekit/psrpc" - lksdk "github.com/livekit/server-sdk-go/v2" "github.com/livekit/sipgo" "github.com/livekit/sipgo/sip" @@ -385,7 +383,6 @@ func (c *outboundCall) close(ctx context.Context, end EndCall) bool { } if r := c.lkRoom; r != nil { - _ = r.CloseOutput() _ = r.CloseWithReason(end.Status.DisconnectReason()) } @@ -519,7 +516,7 @@ func (c *outboundCall) dialSIP(ctx context.Context, tid traceid.ID) error { if digits := c.sipConf.dtmf; digits != "" { c.setStatus(CallAutomation) // Write initial DTMF to SIP - dtmfWriter := c.media.GetDTMFWriter() + dtmfWriter := c.media.GetOutboundDTMFWriter() if err := dtmfWriter.WriteSample(&livekit.SipDTMF{ Digit: digits, }); err != nil { @@ -545,16 +542,29 @@ func (c *outboundCall) updateRemoteFromSDP(body []byte) error { } func (c *outboundCall) connectMedia() { - if w := c.lkRoom.SwapOutput(c.audioOut); w != nil { - _ = w.Close() + if old := c.lkRoom.WriteOutboundAudioTo(c.audioOut); old != nil { + old.Close() + c.log.Warnw("room has unexpected outbound audio writer", nil) + } + + if old := c.lkRoom.WriteOutboundDTMFTo(c.media.GetOutboundDTMFWriter()); old != nil { + old.Close() + c.log.Warnw("room has unexpected outbound DTMF writer", nil) } - c.lkRoom.SetDTMFOutput(c.media.GetDTMFWriter()) if processor := c.c.handler.GetMediaProcessor(c.sipConf.enabledFeatures, c.sipConf.featureFlags, string(c.cc.ID()), MediaProcessorOpts{InputSampleRate: RoomSampleRate}); processor != nil { c.lkRoomIn = processor(c.lkRoomIn) } - c.media.WriteAudioTo(c.lkRoomIn) - c.media.WriteDTMFTo(&dtmfEventWriter{handler: c.handleDTMF}) + + if old := c.media.WriteInboundAudioTo(c.lkRoomIn); old != nil { + old.Close() + c.log.Warnw("media port has unexpected inbound audio writer", nil) + } + + if old := c.media.WriteInboundDTMFTo(c.lkRoom.GetInboundDTMFWriter()); old != nil { + old.Close() + c.log.Warnw("media port has unexpected inbound DTMF writer", nil) + } } type sipRespFunc func(code sip.StatusCode, hdrs Headers) @@ -763,7 +773,7 @@ func (c *outboundCall) sipSignal(ctx context.Context, tid traceid.ID) error { } c.mon.InviteAccept() - if old := c.audioOut.Swap(c.media.GetAudioWriter()); old != nil { + if old := c.audioOut.Swap(c.media.GetOutboundAudioWriter()); old != nil { c.log.Warnw("unexpected audio out writer", nil) old.Close() } @@ -799,30 +809,6 @@ func (c *outboundCall) sipSignal(ctx context.Context, tid traceid.ID) error { return nil } -// TODO(alexfish): Update Room so that we don't need this adapter. -func (c *outboundCall) handleDTMF(msg *livekit.SipDTMF) { - if c.lkRoom == nil { - return - } - - if msg == nil { - return - } - - code := byte(msg.Code) - digit := byte(0) - if len(msg.Digit) == 1 { - digit = msg.Digit[0] - } else { - digit = dtmf.CodeToChar(code) - } - - _ = c.lkRoom.SendData(&livekit.SipDTMF{ - Code: uint32(code), - Digit: string([]byte{digit}), - }, lksdk.WithDataPublishReliable(true)) -} - func (c *outboundCall) transferCall(ctx context.Context, transferTo string, headers map[string]string, dialtone bool) (retErr error) { ctx, span := Tracer.Start(ctx, "sip.outbound.transferCall") defer span.End() @@ -838,14 +824,13 @@ func (c *outboundCall) transferCall(ctx context.Context, transferTo string, head rctx, rcancel := context.WithCancel(ctx) defer rcancel() - // mute the room audio to the SIP participant - w := c.lkRoom.SwapOutput(nil) + // Mute the room audio to the SIP participant. + // Skip closing the existing writer, which is c.audioOut. + _ = c.lkRoom.WriteOutboundAudioTo(nil) defer func() { if retErr != nil && !c.stopped.IsBroken() { - c.lkRoom.SwapOutput(w) - } else { - w.Close() + c.lkRoom.WriteOutboundAudioTo(c.audioOut) } }() diff --git a/pkg/sip/outbound_utilities_test.go b/pkg/sip/outbound_utilities_test.go index 9e850fe5..53262b77 100644 --- a/pkg/sip/outbound_utilities_test.go +++ b/pkg/sip/outbound_utilities_test.go @@ -102,6 +102,8 @@ type testRoom struct { room *Room } +var _ RoomInterface = (*testRoom)(nil) + type testRoomConfig struct { ringForever bool } @@ -122,15 +124,17 @@ func newTestRoomWithConfig(log logger.Logger, st *RoomStats, cfg *testRoomConfig } // Create a Room with all the necessary structure but skip connection room := &Room{ - log: log, - stats: st, - out: msdk.NewSwitchWriter(RoomSampleRate), - subscribe: atomic.Bool{}, + log: log, + stats: st, + outboundAudio: msdk.NewWriteCloserSwitch[msdk.PCM16Sample](RoomSampleRate), + outboundDTMF: msdk.NewWriteCloserSwitch[*livekit.SipDTMF](0), + subscribe: atomic.Bool{}, } + room.inboundDTMF = inboundDTMFWriter{room} // Create mixer var err error - room.mix, err = mixer.NewMixer(room.out, rtp.DefFrameDur, 1, mixer.WithStats(&st.Mixer), mixer.WithOutputChannel()) + room.mix, err = mixer.NewMixer(room.outboundAudio, rtp.DefFrameDur, 1, mixer.WithStats(&st.Mixer), mixer.WithOutputChannel()) if err != nil { panic(err) } @@ -202,20 +206,20 @@ func (r *testRoom) Subscribe() { r.room.Subscribe() } -func (r *testRoom) Output() msdk.Writer[msdk.PCM16Sample] { - return r.room.Output() +func (r *testRoom) WriteOutboundAudioTo(w msdk.PCM16Writer) msdk.PCM16Writer { + return r.room.WriteOutboundAudioTo(w) } -func (r *testRoom) SwapOutput(out msdk.PCM16Writer) msdk.PCM16Writer { - return r.room.SwapOutput(out) +func (r *testRoom) WriteOutboundDTMFTo(w msdk.WriteCloser[*livekit.SipDTMF]) msdk.WriteCloser[*livekit.SipDTMF] { + return r.room.WriteOutboundDTMFTo(w) } -func (r *testRoom) CloseOutput() error { - return r.room.CloseOutput() +func (r *testRoom) GetInboundAudioWriter() (msdk.PCM16Writer, error) { + return r.NewParticipantTrack(RoomSampleRate) } -func (r *testRoom) SetDTMFOutput(w msdk.WriteCloser[*livekit.SipDTMF]) { - r.room.SetDTMFOutput(w) +func (r *testRoom) GetInboundDTMFWriter() msdk.WriteCloser[*livekit.SipDTMF] { + return r.room.GetInboundDTMFWriter() } func (r *testRoom) Close() error { @@ -256,10 +260,6 @@ func (w *noOpWriter) Close() error { return nil } -func (r *testRoom) SendData(data lksdk.DataPacket, opts ...lksdk.DataPublishOption) error { - return r.room.SendData(data, opts...) -} - func (r *testRoom) NewTrack() *mixer.Input { return r.room.NewTrack() } diff --git a/pkg/sip/room.go b/pkg/sip/room.go index 9e362341..6b98c7b9 100644 --- a/pkg/sip/room.go +++ b/pkg/sip/room.go @@ -28,6 +28,7 @@ import ( "github.com/pion/webrtc/v4" msdk "github.com/livekit/media-sdk" + "github.com/livekit/media-sdk/dtmf" "github.com/livekit/media-sdk/g711" "github.com/livekit/media-sdk/jitter" "github.com/livekit/media-sdk/mixer" @@ -184,17 +185,27 @@ type RoomInterface interface { Subscribed() <-chan struct{} Room() *lksdk.Room Subscribe() - Output() msdk.Writer[msdk.PCM16Sample] - SwapOutput(out msdk.PCM16Writer) msdk.PCM16Writer - CloseOutput() error - SetDTMFOutput(w msdk.WriteCloser[*livekit.SipDTMF]) Close() error CloseWithReason(reason livekit.DisconnectReason) error Participant() ParticipantInfo NewParticipantTrack(sampleRate int) (msdk.WriteCloser[msdk.PCM16Sample], error) - SendData(data lksdk.DataPacket, opts ...lksdk.DataPublishOption) error NewTrack() *mixer.Input lksdk.RoomRPCInterface + + // WriteOutboundAudioTo tells the room where to send audio to. + // Returns the previously-set writer (if one exists). + WriteOutboundAudioTo(w msdk.PCM16Writer) msdk.PCM16Writer + + // WriteOutboundDTMFTo tells the room where to send DTMF to. + // Returns the previously-set writer (if one exists). + WriteOutboundDTMFTo(w msdk.WriteCloser[*livekit.SipDTMF]) msdk.WriteCloser[*livekit.SipDTMF] + + // GetInboundAudioWriter returns a writer that, when written to, writes + // audio to the room. + GetInboundAudioWriter() (msdk.PCM16Writer, error) + // GetInboundDTMFWriter returns a writer that, when weritten to, writes DTMF + // to the room. + GetInboundDTMFWriter() msdk.WriteCloser[*livekit.SipDTMF] } type GetRoomFunc func(log logger.Logger, st *RoomStats) RoomInterface @@ -207,10 +218,13 @@ type Room struct { log logger.Logger roomLog logger.Logger // deferred logger // room is cleared on close while SDK callback goroutines still read it. - room atomic.Pointer[lksdk.Room] - mix *mixer.Mixer - out *msdk.SwitchWriter - outDtmf atomic.Pointer[msdk.WriteCloser[*livekit.SipDTMF]] + room atomic.Pointer[lksdk.Room] + mix *mixer.Mixer + + outboundAudio *msdk.WriteCloserSwitch[msdk.PCM16Sample] + outboundDTMF *msdk.WriteCloserSwitch[*livekit.SipDTMF] + inboundDTMF inboundDTMFWriter + // p is replaced on every reconnect, since the server issues a new // participant SID, and read concurrently by Participant(). p atomic.Pointer[ParticipantInfo] @@ -245,10 +259,17 @@ func NewRoom(log logger.Logger, st *RoomStats) *Room { if st == nil { st = &RoomStats{} } - r := &Room{log: log, stats: st, out: msdk.NewSwitchWriter(RoomSampleRate)} + r := &Room{ + log: log, + stats: st, + + outboundAudio: msdk.NewWriteCloserSwitch[msdk.PCM16Sample](RoomSampleRate), + outboundDTMF: msdk.NewWriteCloserSwitch[*livekit.SipDTMF](0), + } + r.inboundDTMF = inboundDTMFWriter{r} var err error - r.mix, err = mixer.NewMixer(r.out, rtp.DefFrameDur, 1, mixer.WithStats(&st.Mixer), mixer.WithOutputChannel()) + r.mix, err = mixer.NewMixer(r.outboundAudio, rtp.DefFrameDur, 1, mixer.WithStats(&st.Mixer), mixer.WithOutputChannel()) if err != nil { panic(err) } @@ -675,50 +696,10 @@ func (r *Room) subscribeAll(room *lksdk.Room) { } } -func (r *Room) Output() msdk.Writer[msdk.PCM16Sample] { - return r.out.Get() -} - -// SwapOutput sets room audio output and returns the old one. -// Caller is responsible for closing the old writer. -func (r *Room) SwapOutput(out msdk.PCM16Writer) msdk.PCM16Writer { - if r == nil { - return nil - } - if out == nil { - return r.out.Swap(nil) - } - return r.out.Swap(msdk.ResampleWriter(out, r.mix.SampleRate())) -} - -func (r *Room) CloseOutput() error { - w := r.SwapOutput(nil) - if w == nil { - return nil - } - return w.Close() -} - -func (r *Room) SetDTMFOutput(w msdk.WriteCloser[*livekit.SipDTMF]) { - if r == nil { - return - } - if w == nil { - r.outDtmf.Store(nil) - return - } - r.outDtmf.Store(&w) -} - func (r *Room) sendDTMF(ctx context.Context, msg *livekit.SipDTMF) { - outDTMF := r.outDtmf.Load() - if outDTMF == nil { - r.log.Infow("ignoring dtmf", "digit", msg.Digit) - return - } // TODO: Separate goroutine? r.log.Debugw("forwarding dtmf to sip", "digit", msg.Digit) - (*outDTMF).WriteSample(msg) + r.outboundDTMF.WriteSample(msg) } func (r *Room) Close() error { @@ -729,13 +710,13 @@ func (r *Room) CloseWithReason(reason livekit.DisconnectReason) error { if r == nil { return nil } - var err error + var errs []error r.closed.Once(func() { defer r.stats.Closed.Store(true) r.subscribe.Store(false) - err = r.CloseOutput() - r.SetDTMFOutput(nil) + errs = append(errs, r.outboundAudio.Close()) + errs = append(errs, r.outboundDTMF.Close()) if room := r.room.Swap(nil); room != nil { room.DisconnectWithReason(reason) } @@ -743,7 +724,7 @@ func (r *Room) CloseWithReason(reason livekit.DisconnectReason) error { r.mix.Stop() } }) - return err + return errors.Join(errs...) } func (r *Room) Participant() ParticipantInfo { @@ -756,6 +737,7 @@ func (r *Room) Participant() ParticipantInfo { return ParticipantInfo{} } +// TODO(alexfish): Remove this from the public interface. func (r *Room) NewParticipantTrack(sampleRate int) (msdk.WriteCloser[msdk.PCM16Sample], error) { track, err := webrtc.NewTrackLocalStaticSample(webrtc.RTPCodecCapability{MimeType: webrtc.MimeTypeOpus}, "audio", "pion") if err != nil { @@ -797,6 +779,45 @@ func (r *Room) NewTrack() *mixer.Input { return r.mix.NewInput() } +func (r *Room) WriteOutboundAudioTo(w msdk.PCM16Writer) msdk.PCM16Writer { + return r.outboundAudio.Swap(w) +} + +func (r *Room) WriteOutboundDTMFTo(w msdk.WriteCloser[*livekit.SipDTMF]) msdk.WriteCloser[*livekit.SipDTMF] { + return r.outboundDTMF.Swap(w) +} + +func (r *Room) GetInboundAudioWriter() (msdk.PCM16Writer, error) { + return r.NewParticipantTrack(RoomSampleRate) +} + +func (r *Room) GetInboundDTMFWriter() msdk.WriteCloser[*livekit.SipDTMF] { + return &r.inboundDTMF +} + +type inboundDTMFWriter struct { + r *Room +} + +func (w *inboundDTMFWriter) String() string { + return "inboundDTMFWriter" +} + +func (w *inboundDTMFWriter) SampleRate() int { + return dtmf.SampleRate +} + +func (w *inboundDTMFWriter) Close() error { + return nil +} + +func (w *inboundDTMFWriter) WriteSample(sample *livekit.SipDTMF) error { + if sample == nil { + return nil + } + return w.r.SendData(sample, lksdk.WithDataPublishReliable(true)) +} + // roomOverrideLogger converts errors to warnings and ignore debug type roomOverrideLogger struct { logger.Logger