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
3 changes: 3 additions & 0 deletions unix/linux/types.go
Original file line number Diff line number Diff line change
Expand Up @@ -775,6 +775,8 @@ type PacketMreq C.struct_packet_mreq

type Msghdr C.struct_msghdr

type Mmsghdr C.struct_mmsghdr

type Cmsghdr C.struct_cmsghdr

type Inet4Pktinfo C.struct_in_pktinfo
Expand Down Expand Up @@ -829,6 +831,7 @@ const (
SizeofIPv6Mreq = C.sizeof_struct_ipv6_mreq
SizeofPacketMreq = C.sizeof_struct_packet_mreq
SizeofMsghdr = C.sizeof_struct_msghdr
SizeofMmsghdr = C.sizeof_struct_mmsghdr
SizeofCmsghdr = C.sizeof_struct_cmsghdr
SizeofInet4Pktinfo = C.sizeof_struct_in_pktinfo
SizeofInet6Pktinfo = C.sizeof_struct_in6_pktinfo
Expand Down
81 changes: 81 additions & 0 deletions unix/syscall_linux.go
Original file line number Diff line number Diff line change
Expand Up @@ -1617,6 +1617,87 @@ func sendmsgN(fd int, iov []Iovec, oob []byte, ptr unsafe.Pointer, salen _Sockle
return n, nil
}

// RecvmmsgData holds the per-message buffers and results for [Recvmmsg].
// A RecvmmsgData slice can be pre-allocated and reused across calls without
// any per-call allocation.
type RecvmmsgData struct {
Data [][]byte // payload buffers; each entry is a separate iovec
OOB []byte // control buffer
N int // set on return: payload bytes received into Data
OOBN int // set on return: OOB bytes received into OOB; interpret with [ParseSocketControlMessage]
Flags int // set on return: per-message flags
From RawSockaddrAny // set on return: raw sender address; zero (AF_UNSPEC) for connected sockets
}

// Recvmmsg receives multiple messages from a socket using the recvmmsg system
// call. msgs is a caller-provided slice, one entry per message slot. Data and
// OOB in each entry must be pre-allocated; the call writes N, OOBN, Flags,
// and From back into each entry. n is the number of messages received.
//
// To convert From to a [Sockaddr], call [AnyToSockaddr].
func Recvmmsg(fd int, msgs []RecvmmsgData, flags int) (n int, err error) {
vlen := len(msgs)
if vlen == 0 {
return 0, EINVAL
}

totalIovecs := 0
for i := range msgs {
totalIovecs += len(msgs[i].Data)
}

msghdrs := make([]Mmsghdr, vlen)
var iovecs []Iovec
if totalIovecs > 0 {
iovecs = make([]Iovec, totalIovecs)
}

iovIdx := 0
for i := range vlen {
m := &msgs[i]
// clear stale address from previous call on reuse
m.From = RawSockaddrAny{}
if len(m.Data) > 0 {
startIdx := iovIdx
for _, b := range m.Data {
if len(b) > 0 {
iovecs[iovIdx].Base = &b[0]
iovecs[iovIdx].SetLen(len(b))
} else {
iovecs[iovIdx].Base = (*byte)(unsafe.Pointer(&_zero))
}
iovIdx++
}
msghdrs[i].Hdr.Iov = &iovecs[startIdx]
msghdrs[i].Hdr.SetIovlen(len(m.Data))
}
if len(m.OOB) > 0 {
msghdrs[i].Hdr.Control = &m.OOB[0]
msghdrs[i].Hdr.SetControllen(len(m.OOB))
}
msghdrs[i].Hdr.Name = (*byte)(unsafe.Pointer(&m.From))
msghdrs[i].Hdr.Namelen = uint32(SizeofSockaddrAny)
}

n, err = recvmmsg(fd, &msghdrs[0], vlen, flags, nil)
if err != nil {
return 0, err
}

for i := range n {
msgs[i].N = int(msghdrs[i].Len)
msgs[i].OOBN = int(msghdrs[i].Hdr.Controllen)
msgs[i].Flags = int(msghdrs[i].Hdr.Flags)
}

return
}

// AnyToSockaddr converts a raw socket address to a [Sockaddr] interface.
func AnyToSockaddr(fd int, rsa *RawSockaddrAny) (Sockaddr, error) {
return anyToSockaddr(fd, rsa)
}

// BindToDevice binds the socket associated with fd to device.
func BindToDevice(fd int, device string) (err error) {
return SetsockoptString(fd, SOL_SOCKET, SO_BINDTODEVICE, device)
Expand Down
8 changes: 8 additions & 0 deletions unix/syscall_linux_386.go
Original file line number Diff line number Diff line change
Expand Up @@ -252,6 +252,14 @@ func sendmsg(s int, msg *Msghdr, flags int) (n int, err error) {
return
}

func recvmmsg(s int, mmsg *Mmsghdr, vlen int, flags int, timeout *Timespec) (n int, err error) {
n, e := socketcall(_RECVMMSG, uintptr(s), uintptr(unsafe.Pointer(mmsg)), uintptr(vlen), uintptr(flags), uintptr(unsafe.Pointer(timeout)), 0)
if e != 0 {
err = e
}
return
}

func Listen(s int, n int) (err error) {
_, e := socketcall(_LISTEN, uintptr(s), uintptr(n), 0, 0, 0, 0)
if e != 0 {
Expand Down
1 change: 1 addition & 0 deletions unix/syscall_linux_amd64.go
Original file line number Diff line number Diff line change
Expand Up @@ -71,6 +71,7 @@ func Stat(path string, stat *Stat_t) (err error) {
//sys sendto(s int, buf []byte, flags int, to unsafe.Pointer, addrlen _Socklen) (err error)
//sys recvmsg(s int, msg *Msghdr, flags int) (n int, err error)
//sys sendmsg(s int, msg *Msghdr, flags int) (n int, err error)
//sys recvmmsg(s int, mmsg *Mmsghdr, vlen int, flags int, timeout *Timespec) (n int, err error)
//sys mmap(addr uintptr, length uintptr, prot int, flags int, fd int, offset int64) (xaddr uintptr, err error)

//sys futimesat(dirfd int, path string, times *[2]Timeval) (err error)
Expand Down
1 change: 1 addition & 0 deletions unix/syscall_linux_arm.go
Original file line number Diff line number Diff line change
Expand Up @@ -41,6 +41,7 @@ func Seek(fd int, offset int64, whence int) (newoffset int64, err error) {
//sysnb socketpair(domain int, typ int, flags int, fd *[2]int32) (err error)
//sys recvmsg(s int, msg *Msghdr, flags int) (n int, err error)
//sys sendmsg(s int, msg *Msghdr, flags int) (n int, err error)
//sys recvmmsg(s int, mmsg *Mmsghdr, vlen int, flags int, timeout *Timespec) (n int, err error)

// 64-bit file system and 32-bit uid calls
// (16-bit uid calls are not always supported in newer kernels)
Expand Down
1 change: 1 addition & 0 deletions unix/syscall_linux_arm64.go
Original file line number Diff line number Diff line change
Expand Up @@ -75,6 +75,7 @@ func Ustat(dev int, ubuf *Ustat_t) (err error) {
//sys sendto(s int, buf []byte, flags int, to unsafe.Pointer, addrlen _Socklen) (err error)
//sys recvmsg(s int, msg *Msghdr, flags int) (n int, err error)
//sys sendmsg(s int, msg *Msghdr, flags int) (n int, err error)
//sys recvmmsg(s int, mmsg *Mmsghdr, vlen int, flags int, timeout *Timespec) (n int, err error)
//sys mmap(addr uintptr, length uintptr, prot int, flags int, fd int, offset int64) (xaddr uintptr, err error)

//sysnb Gettimeofday(tv *Timeval) (err error)
Expand Down
1 change: 1 addition & 0 deletions unix/syscall_linux_loong64.go
Original file line number Diff line number Diff line change
Expand Up @@ -107,6 +107,7 @@ func Ustat(dev int, ubuf *Ustat_t) (err error) {
//sys sendto(s int, buf []byte, flags int, to unsafe.Pointer, addrlen _Socklen) (err error)
//sys recvmsg(s int, msg *Msghdr, flags int) (n int, err error)
//sys sendmsg(s int, msg *Msghdr, flags int) (n int, err error)
//sys recvmmsg(s int, mmsg *Mmsghdr, vlen int, flags int, timeout *Timespec) (n int, err error)
//sys mmap(addr uintptr, length uintptr, prot int, flags int, fd int, offset int64) (xaddr uintptr, err error)

//sysnb Gettimeofday(tv *Timeval) (err error)
Expand Down
1 change: 1 addition & 0 deletions unix/syscall_linux_mips64x.go
Original file line number Diff line number Diff line change
Expand Up @@ -55,6 +55,7 @@ func Select(nfd int, r *FdSet, w *FdSet, e *FdSet, timeout *Timeval) (n int, err
//sys sendto(s int, buf []byte, flags int, to unsafe.Pointer, addrlen _Socklen) (err error)
//sys recvmsg(s int, msg *Msghdr, flags int) (n int, err error)
//sys sendmsg(s int, msg *Msghdr, flags int) (n int, err error)
//sys recvmmsg(s int, mmsg *Mmsghdr, vlen int, flags int, timeout *Timespec) (n int, err error)
//sys mmap(addr uintptr, length uintptr, prot int, flags int, fd int, offset int64) (xaddr uintptr, err error)

//sys futimesat(dirfd int, path string, times *[2]Timeval) (err error)
Expand Down
1 change: 1 addition & 0 deletions unix/syscall_linux_mipsx.go
Original file line number Diff line number Diff line change
Expand Up @@ -49,6 +49,7 @@ func Syscall9(trap, a1, a2, a3, a4, a5, a6, a7, a8, a9 uintptr) (r1, r2 uintptr,
//sys sendto(s int, buf []byte, flags int, to unsafe.Pointer, addrlen _Socklen) (err error)
//sys recvmsg(s int, msg *Msghdr, flags int) (n int, err error)
//sys sendmsg(s int, msg *Msghdr, flags int) (n int, err error)
//sys recvmmsg(s int, mmsg *Mmsghdr, vlen int, flags int, timeout *Timespec) (n int, err error)

//sys Ioperm(from int, num int, on int) (err error)
//sys Iopl(level int) (err error)
Expand Down
1 change: 1 addition & 0 deletions unix/syscall_linux_ppc.go
Original file line number Diff line number Diff line change
Expand Up @@ -52,6 +52,7 @@ import (
//sys sendto(s int, buf []byte, flags int, to unsafe.Pointer, addrlen _Socklen) (err error)
//sys recvmsg(s int, msg *Msghdr, flags int) (n int, err error)
//sys sendmsg(s int, msg *Msghdr, flags int) (n int, err error)
//sys recvmmsg(s int, mmsg *Mmsghdr, vlen int, flags int, timeout *Timespec) (n int, err error)

//sys futimesat(dirfd int, path string, times *[2]Timeval) (err error)
//sysnb Gettimeofday(tv *Timeval) (err error)
Expand Down
1 change: 1 addition & 0 deletions unix/syscall_linux_ppc64x.go
Original file line number Diff line number Diff line change
Expand Up @@ -52,6 +52,7 @@ package unix
//sys sendto(s int, buf []byte, flags int, to unsafe.Pointer, addrlen _Socklen) (err error)
//sys recvmsg(s int, msg *Msghdr, flags int) (n int, err error)
//sys sendmsg(s int, msg *Msghdr, flags int) (n int, err error)
//sys recvmmsg(s int, mmsg *Mmsghdr, vlen int, flags int, timeout *Timespec) (n int, err error)
//sys mmap(addr uintptr, length uintptr, prot int, flags int, fd int, offset int64) (xaddr uintptr, err error)

//sys futimesat(dirfd int, path string, times *[2]Timeval) (err error)
Expand Down
1 change: 1 addition & 0 deletions unix/syscall_linux_riscv64.go
Original file line number Diff line number Diff line change
Expand Up @@ -74,6 +74,7 @@ func Ustat(dev int, ubuf *Ustat_t) (err error) {
//sys sendto(s int, buf []byte, flags int, to unsafe.Pointer, addrlen _Socklen) (err error)
//sys recvmsg(s int, msg *Msghdr, flags int) (n int, err error)
//sys sendmsg(s int, msg *Msghdr, flags int) (n int, err error)
//sys recvmmsg(s int, mmsg *Mmsghdr, vlen int, flags int, timeout *Timespec) (n int, err error)
//sys mmap(addr uintptr, length uintptr, prot int, flags int, fd int, offset int64) (xaddr uintptr, err error)

//sysnb Gettimeofday(tv *Timeval) (err error)
Expand Down
1 change: 1 addition & 0 deletions unix/syscall_linux_sparc64.go
Original file line number Diff line number Diff line change
Expand Up @@ -50,6 +50,7 @@ package unix
//sys sendto(s int, buf []byte, flags int, to unsafe.Pointer, addrlen _Socklen) (err error)
//sys recvmsg(s int, msg *Msghdr, flags int) (n int, err error)
//sys sendmsg(s int, msg *Msghdr, flags int) (n int, err error)
//sys recvmmsg(s int, mmsg *Mmsghdr, vlen int, flags int, timeout *Timespec) (n int, err error)
//sys mmap(addr uintptr, length uintptr, prot int, flags int, fd int, offset int64) (xaddr uintptr, err error)

func Ioperm(from int, num int, on int) (err error) {
Expand Down
187 changes: 187 additions & 0 deletions unix/syscall_linux_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -1364,3 +1364,190 @@ func TestSockaddrALG(t *testing.T) {
t.Fatalf("got: %q, want: %q", got, exp)
}
}

func TestRecvmmsg(t *testing.T) {
tests := []struct {
name string
messages int
batchSize int
}{
{
name: "equal_messages_and_batch",
messages: 3,
batchSize: 3,
},
{
name: "fewer_messages_than_batch",
messages: 2,
batchSize: 6,
},
{
name: "more_messages_than_batch",
messages: 5,
batchSize: 2,
},
}

for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
fds, err := unix.Socketpair(unix.AF_UNIX, unix.SOCK_DGRAM, 0)
if err != nil {
t.Fatal(err)
}
defer unix.Close(fds[0])
defer unix.Close(fds[1])

for i := 0; i < tt.messages; i++ {
msg := fmt.Sprintf("msg%d", i+1)
if _, err := unix.Write(fds[1], []byte(msg)); err != nil {
t.Fatalf("Write: %v", err)
}
}

msgs := make([]unix.RecvmmsgData, tt.batchSize)
for i := range msgs {
msgs[i].Data = [][]byte{make([]byte, 64)}
}

read := 0
for read < tt.messages {
n, err := unix.Recvmmsg(fds[0], msgs, unix.MSG_DONTWAIT)
if err != nil {
if errors.Is(err, unix.ENOSYS) {
t.Skipf("recvmmsg not available: %v", err)
}
t.Fatalf("Recvmmsg: %v", err)
}

wantBatchSize := min(tt.messages-read, tt.batchSize)
if n != wantBatchSize {
t.Fatalf("Recvmmsg: got %d messages, want %d", n, wantBatchSize)
}

for i := range n {
got := string(msgs[i].Data[0][:msgs[i].N])
want := fmt.Sprintf("msg%d", read+i+1)
if got != want {
t.Errorf("message %d: got %q, want %q", i, got, want)
}
if msgs[i].N != len(want) {
t.Errorf("message %d: got N=%d, want %d", i, msgs[i].N, len(want))
}
if msgs[i].OOBN != 0 {
t.Errorf("message %d: got OOBN=%d, want 0", i, msgs[i].OOBN)
}
if msgs[i].Flags != 0 {
t.Errorf("message %d: got Flags=%#x, want 0", i, msgs[i].Flags)
}
// socketpair is connected; kernel does not fill sender address
if msgs[i].From.Addr.Family != unix.AF_UNSPEC {
t.Errorf("message %d: got From.Addr.Family=%d, want AF_UNSPEC", i, msgs[i].From.Addr.Family)
}
}

read += n
}
})
}
}

func TestRecvmmsgScatterGather(t *testing.T) {
fds, err := unix.Socketpair(unix.AF_UNIX, unix.SOCK_DGRAM, 0)
if err != nil {
t.Fatal(err)
}
defer unix.Close(fds[0])
defer unix.Close(fds[1])

// Send two messages; each will be received into two separate buffers.
for i := range 2 {
msg := fmt.Sprintf("abcd%d", i)
if _, err := unix.Write(fds[1], []byte(msg)); err != nil {
t.Fatalf("Write: %v", err)
}
}

msgs := []unix.RecvmmsgData{
{Data: [][]byte{make([]byte, 2), make([]byte, 3)}},
{Data: [][]byte{make([]byte, 2), make([]byte, 3)}},
}
n, err := unix.Recvmmsg(fds[0], msgs, unix.MSG_DONTWAIT)
if err != nil {
if errors.Is(err, unix.ENOSYS) {
t.Skipf("recvmmsg not available: %v", err)
}
t.Fatalf("Recvmmsg: %v", err)
}
if n != 2 {
t.Fatalf("got %d messages, want 2", n)
}
for i := range n {
want := fmt.Sprintf("abcd%d", i)
// N is total bytes across both scatter buffers.
if msgs[i].N != len(want) {
t.Errorf("message %d: got N=%d, want %d", i, msgs[i].N, len(want))
}
got := string(msgs[i].Data[0]) + string(msgs[i].Data[1][:msgs[i].N-len(msgs[i].Data[0])])
if got != want {
t.Errorf("message %d: got %q, want %q", i, got, want)
}
}
}

func TestRecvmmsgFrom(t *testing.T) {
// Use an unconnected UDP socket so the kernel fills in the sender address.
srv, err := unix.Socket(unix.AF_INET, unix.SOCK_DGRAM, 0)
if err != nil {
t.Fatal(err)
}
defer unix.Close(srv)
addr := unix.SockaddrInet4{Port: 0, Addr: [4]byte{127, 0, 0, 1}}
if err := unix.Bind(srv, &addr); err != nil {
t.Fatal(err)
}
sa, err := unix.Getsockname(srv)
if err != nil {
t.Fatal(err)
}
port := sa.(*unix.SockaddrInet4).Port

cli, err := unix.Socket(unix.AF_INET, unix.SOCK_DGRAM, 0)
if err != nil {
t.Fatal(err)
}
defer unix.Close(cli)

dst := &unix.SockaddrInet4{Port: port, Addr: [4]byte{127, 0, 0, 1}}
if err := unix.Sendto(cli, []byte("hello"), 0, dst); err != nil {
t.Fatal(err)
}

msgs := []unix.RecvmmsgData{{Data: [][]byte{make([]byte, 16)}}}
n, err := unix.Recvmmsg(srv, msgs, unix.MSG_DONTWAIT)
if err != nil {
if errors.Is(err, unix.ENOSYS) {
t.Skipf("recvmmsg not available: %v", err)
}
t.Fatal(err)
}
if n != 1 {
t.Fatalf("got %d messages, want 1", n)
}
if string(msgs[0].Data[0][:msgs[0].N]) != "hello" {
t.Errorf("got payload %q, want %q", msgs[0].Data[0][:msgs[0].N], "hello")
}
if msgs[0].From.Addr.Family == unix.AF_UNSPEC {
t.Fatal("From is AF_UNSPEC; expected sender address")
}
from, err := unix.AnyToSockaddr(srv, &msgs[0].From)
if err != nil {
t.Fatalf("AnyToSockaddr: %v", err)
}
fromInet, ok := from.(*unix.SockaddrInet4)
if !ok {
t.Fatalf("expected *SockaddrInet4, got %T", from)
}
if fromInet.Addr != ([4]byte{127, 0, 0, 1}) {
t.Errorf("got sender addr %v, want 127.0.0.1", fromInet.Addr)
}
}
Loading