Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
37 changes: 32 additions & 5 deletions srtgo.go
Original file line number Diff line number Diff line change
Expand Up @@ -311,13 +311,29 @@ type ListenCallbackFunc func(socket *SrtSocket, version int, addr *net.UDPAddr,

//export srtListenCBWrapper
func srtListenCBWrapper(arg unsafe.Pointer, socket C.SRTSOCKET, hsVersion C.int, peeraddr *C.struct_sockaddr, streamid *C.char) C.int {
userCB := gopointer.Restore(arg).(ListenCallbackFunc)
callbackMutex.Lock()
restored := gopointer.Restore(arg)
callbackMutex.Unlock()

userCB, ok := restored.(ListenCallbackFunc)
if !ok || userCB == nil {
return SRT_ERROR
}

s := new(SrtSocket)
s.socket = socket
udpAddr, _ := udpAddrFromSockaddr((*syscall.RawSockaddrAny)(unsafe.Pointer(peeraddr)))

if userCB(s, int(hsVersion), udpAddr, C.GoString(streamid)) {
var udpAddr *net.UDPAddr
if peeraddr != nil {
udpAddr, _ = udpAddrFromSockaddr((*syscall.RawSockaddrAny)(unsafe.Pointer(peeraddr)))
}

var sid string
if streamid != nil {
sid = C.GoString(streamid)
}

if userCB(s, int(hsVersion), udpAddr, sid) {
return 0
}
return SRT_ERROR
Expand All @@ -344,11 +360,22 @@ type ConnectCallbackFunc func(socket *SrtSocket, err error, addr *net.UDPAddr, t

//export srtConnectCBWrapper
func srtConnectCBWrapper(arg unsafe.Pointer, socket C.SRTSOCKET, errcode C.int, peeraddr *C.struct_sockaddr, token C.int) {
userCB := gopointer.Restore(arg).(ConnectCallbackFunc)
callbackMutex.Lock()
restored := gopointer.Restore(arg)
callbackMutex.Unlock()

userCB, ok := restored.(ConnectCallbackFunc)
if !ok || userCB == nil {
return
}

s := new(SrtSocket)
s.socket = socket
udpAddr, _ := udpAddrFromSockaddr((*syscall.RawSockaddrAny)(unsafe.Pointer(peeraddr)))

var udpAddr *net.UDPAddr
if peeraddr != nil {
udpAddr, _ = udpAddrFromSockaddr((*syscall.RawSockaddrAny)(unsafe.Pointer(peeraddr)))
}

userCB(s, SRTErrno(errcode), udpAddr, int(token))
}
Expand Down
34 changes: 34 additions & 0 deletions srtgo_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -6,6 +6,8 @@ import (
"sync"
"testing"
"time"

gopointer "github.com/mattn/go-pointer"
)

func init() {
Expand Down Expand Up @@ -341,3 +343,35 @@ func TestClose(t *testing.T) {
t.Error("Failed to delete connect callback")
}
}

func TestListenCallbackWrapperStalePointerDoesNotPanic(t *testing.T) {
ptr := gopointer.Save(ListenCallbackFunc(func(socket *SrtSocket, version int, addr *net.UDPAddr, streamid string) bool {
return true
}))
gopointer.Unref(ptr)

defer func() {
if r := recover(); r != nil {
t.Fatalf("listen callback wrapper panicked with stale pointer: %v", r)
}
}()

ret := srtListenCBWrapper(ptr, 0, 0, nil, nil)
if ret != SRT_ERROR {
t.Fatalf("expected SRT_ERROR for stale callback pointer, got %d", int(ret))
}
}

func TestConnectCallbackWrapperStalePointerDoesNotPanic(t *testing.T) {
ptr := gopointer.Save(ConnectCallbackFunc(func(socket *SrtSocket, err error, addr *net.UDPAddr, token int) {
}))
gopointer.Unref(ptr)

defer func() {
if r := recover(); r != nil {
t.Fatalf("connect callback wrapper panicked with stale pointer: %v", r)
}
}()

srtConnectCBWrapper(ptr, 0, 0, nil, 0)
}