udp_linux.go 8.6 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356
  1. //go:build !android && !e2e_testing
  2. // +build !android,!e2e_testing
  3. package udp
  4. import (
  5. "encoding/binary"
  6. "fmt"
  7. "net"
  8. "net/netip"
  9. "syscall"
  10. "time"
  11. "unsafe"
  12. "github.com/rcrowley/go-metrics"
  13. "github.com/sirupsen/logrus"
  14. "github.com/slackhq/nebula/config"
  15. "github.com/slackhq/nebula/packet"
  16. "golang.org/x/sys/unix"
  17. )
  18. var readTimeout = unix.NsecToTimeval(int64(time.Millisecond * 500))
  19. type StdConn struct {
  20. sysFd int
  21. isV4 bool
  22. l *logrus.Logger
  23. batch int
  24. }
  25. func NewListener(l *logrus.Logger, ip netip.Addr, port int, multi bool, batch int) (Conn, error) {
  26. af := unix.AF_INET6
  27. if ip.Is4() {
  28. af = unix.AF_INET
  29. }
  30. syscall.ForkLock.RLock()
  31. fd, err := unix.Socket(af, unix.SOCK_DGRAM, unix.IPPROTO_UDP)
  32. if err == nil {
  33. unix.CloseOnExec(fd)
  34. }
  35. syscall.ForkLock.RUnlock()
  36. if err != nil {
  37. unix.Close(fd)
  38. return nil, fmt.Errorf("unable to open socket: %s", err)
  39. }
  40. if multi {
  41. if err = unix.SetsockoptInt(fd, unix.SOL_SOCKET, unix.SO_REUSEPORT, 1); err != nil {
  42. return nil, fmt.Errorf("unable to set SO_REUSEPORT: %s", err)
  43. }
  44. }
  45. // Set a read timeout
  46. if err = unix.SetsockoptTimeval(fd, unix.SOL_SOCKET, unix.SO_RCVTIMEO, &readTimeout); err != nil {
  47. return nil, fmt.Errorf("unable to set SO_RCVTIMEO: %s", err)
  48. }
  49. var sa unix.Sockaddr
  50. if ip.Is4() {
  51. sa4 := &unix.SockaddrInet4{Port: port}
  52. sa4.Addr = ip.As4()
  53. sa = sa4
  54. } else {
  55. sa6 := &unix.SockaddrInet6{Port: port}
  56. sa6.Addr = ip.As16()
  57. sa = sa6
  58. }
  59. if err = unix.Bind(fd, sa); err != nil {
  60. return nil, fmt.Errorf("unable to bind to socket: %s", err)
  61. }
  62. return &StdConn{sysFd: fd, isV4: ip.Is4(), l: l, batch: batch}, err
  63. }
  64. func (u *StdConn) Rebind() error {
  65. return nil
  66. }
  67. func (u *StdConn) SetRecvBuffer(n int) error {
  68. return unix.SetsockoptInt(u.sysFd, unix.SOL_SOCKET, unix.SO_RCVBUFFORCE, n)
  69. }
  70. func (u *StdConn) SetSendBuffer(n int) error {
  71. return unix.SetsockoptInt(u.sysFd, unix.SOL_SOCKET, unix.SO_SNDBUFFORCE, n)
  72. }
  73. func (u *StdConn) SetSoMark(mark int) error {
  74. return unix.SetsockoptInt(u.sysFd, unix.SOL_SOCKET, unix.SO_MARK, mark)
  75. }
  76. func (u *StdConn) GetRecvBuffer() (int, error) {
  77. return unix.GetsockoptInt(int(u.sysFd), unix.SOL_SOCKET, unix.SO_RCVBUF)
  78. }
  79. func (u *StdConn) GetSendBuffer() (int, error) {
  80. return unix.GetsockoptInt(int(u.sysFd), unix.SOL_SOCKET, unix.SO_SNDBUF)
  81. }
  82. func (u *StdConn) GetSoMark() (int, error) {
  83. return unix.GetsockoptInt(int(u.sysFd), unix.SOL_SOCKET, unix.SO_MARK)
  84. }
  85. func (u *StdConn) LocalAddr() (netip.AddrPort, error) {
  86. sa, err := unix.Getsockname(u.sysFd)
  87. if err != nil {
  88. return netip.AddrPort{}, err
  89. }
  90. switch sa := sa.(type) {
  91. case *unix.SockaddrInet4:
  92. return netip.AddrPortFrom(netip.AddrFrom4(sa.Addr), uint16(sa.Port)), nil
  93. case *unix.SockaddrInet6:
  94. return netip.AddrPortFrom(netip.AddrFrom16(sa.Addr), uint16(sa.Port)), nil
  95. default:
  96. return netip.AddrPort{}, fmt.Errorf("unsupported sock type: %T", sa)
  97. }
  98. }
  99. func (u *StdConn) ListenOut(pg PacketBufferGetter, pc chan *packet.Packet) error {
  100. var ip netip.Addr
  101. msgs, packets, names := u.PrepareRawMessages(u.batch, pg)
  102. read := u.ReadMulti
  103. if u.batch == 1 {
  104. read = u.ReadSingle
  105. }
  106. for {
  107. n, err := read(msgs)
  108. if err != nil {
  109. return err
  110. }
  111. for i := 0; i < n; i++ {
  112. out := packets[i]
  113. out.Payload = out.Payload[:msgs[i].Len]
  114. // Its ok to skip the ok check here, the slicing is the only error that can occur and it will panic
  115. if u.isV4 {
  116. ip, _ = netip.AddrFromSlice(names[i][4:8])
  117. } else {
  118. ip, _ = netip.AddrFromSlice(names[i][8:24])
  119. }
  120. out.Addr = netip.AddrPortFrom(ip.Unmap(), binary.BigEndian.Uint16(names[i][2:4]))
  121. pc <- out
  122. //rotate this packet out so we don't overwrite it
  123. packets[i] = pg()
  124. msgs[i].Hdr.Iov.Base = &packets[i].Payload[0]
  125. }
  126. }
  127. }
  128. func (u *StdConn) ReadSingle(msgs []rawMessage) (int, error) {
  129. for {
  130. n, _, err := unix.Syscall6(
  131. unix.SYS_RECVMSG,
  132. uintptr(u.sysFd),
  133. uintptr(unsafe.Pointer(&(msgs[0].Hdr))),
  134. 0,
  135. 0,
  136. 0,
  137. 0,
  138. )
  139. if err != 0 {
  140. if err == unix.EAGAIN || err == unix.EINTR {
  141. continue
  142. }
  143. return 0, &net.OpError{Op: "recvmsg", Err: err}
  144. }
  145. msgs[0].Len = uint32(n)
  146. return 1, nil
  147. }
  148. }
  149. func (u *StdConn) ReadMulti(msgs []rawMessage) (int, error) {
  150. for {
  151. n, _, err := unix.Syscall6(
  152. unix.SYS_RECVMMSG,
  153. uintptr(u.sysFd),
  154. uintptr(unsafe.Pointer(&msgs[0])),
  155. uintptr(len(msgs)),
  156. unix.MSG_WAITFORONE,
  157. 0,
  158. 0,
  159. )
  160. if err != 0 {
  161. if err == unix.EAGAIN || err == unix.EINTR {
  162. continue
  163. }
  164. return 0, &net.OpError{Op: "recvmmsg", Err: err}
  165. }
  166. return int(n), nil
  167. }
  168. }
  169. func (u *StdConn) WriteTo(b []byte, ip netip.AddrPort) error {
  170. if u.isV4 {
  171. return u.writeTo4(b, ip)
  172. }
  173. return u.writeTo6(b, ip)
  174. }
  175. func (u *StdConn) writeTo6(b []byte, ip netip.AddrPort) error {
  176. var rsa unix.RawSockaddrInet6
  177. rsa.Family = unix.AF_INET6
  178. rsa.Addr = ip.Addr().As16()
  179. binary.BigEndian.PutUint16((*[2]byte)(unsafe.Pointer(&rsa.Port))[:], ip.Port())
  180. for {
  181. _, _, err := unix.Syscall6(
  182. unix.SYS_SENDTO,
  183. uintptr(u.sysFd),
  184. uintptr(unsafe.Pointer(&b[0])),
  185. uintptr(len(b)),
  186. uintptr(0),
  187. uintptr(unsafe.Pointer(&rsa)),
  188. uintptr(unix.SizeofSockaddrInet6),
  189. )
  190. if err != 0 {
  191. return &net.OpError{Op: "sendto", Err: err}
  192. }
  193. return nil
  194. }
  195. }
  196. func (u *StdConn) writeTo4(b []byte, ip netip.AddrPort) error {
  197. if !ip.Addr().Is4() {
  198. return fmt.Errorf("Listener is IPv4, but writing to IPv6 remote")
  199. }
  200. var rsa unix.RawSockaddrInet4
  201. rsa.Family = unix.AF_INET
  202. rsa.Addr = ip.Addr().As4()
  203. binary.BigEndian.PutUint16((*[2]byte)(unsafe.Pointer(&rsa.Port))[:], ip.Port())
  204. for {
  205. _, _, err := unix.Syscall6(
  206. unix.SYS_SENDTO,
  207. uintptr(u.sysFd),
  208. uintptr(unsafe.Pointer(&b[0])),
  209. uintptr(len(b)),
  210. uintptr(0),
  211. uintptr(unsafe.Pointer(&rsa)),
  212. uintptr(unix.SizeofSockaddrInet4),
  213. )
  214. if err != 0 {
  215. return &net.OpError{Op: "sendto", Err: err}
  216. }
  217. return nil
  218. }
  219. }
  220. func (u *StdConn) ReloadConfig(c *config.C) {
  221. b := c.GetInt("listen.read_buffer", 0)
  222. if b > 0 {
  223. err := u.SetRecvBuffer(b)
  224. if err == nil {
  225. s, err := u.GetRecvBuffer()
  226. if err == nil {
  227. u.l.WithField("size", s).Info("listen.read_buffer was set")
  228. } else {
  229. u.l.WithError(err).Warn("Failed to get listen.read_buffer")
  230. }
  231. } else {
  232. u.l.WithError(err).Error("Failed to set listen.read_buffer")
  233. }
  234. }
  235. b = c.GetInt("listen.write_buffer", 0)
  236. if b > 0 {
  237. err := u.SetSendBuffer(b)
  238. if err == nil {
  239. s, err := u.GetSendBuffer()
  240. if err == nil {
  241. u.l.WithField("size", s).Info("listen.write_buffer was set")
  242. } else {
  243. u.l.WithError(err).Warn("Failed to get listen.write_buffer")
  244. }
  245. } else {
  246. u.l.WithError(err).Error("Failed to set listen.write_buffer")
  247. }
  248. }
  249. b = c.GetInt("listen.so_mark", 0)
  250. s, err := u.GetSoMark()
  251. if b > 0 || (err == nil && s != 0) {
  252. err := u.SetSoMark(b)
  253. if err == nil {
  254. s, err := u.GetSoMark()
  255. if err == nil {
  256. u.l.WithField("mark", s).Info("listen.so_mark was set")
  257. } else {
  258. u.l.WithError(err).Warn("Failed to get listen.so_mark")
  259. }
  260. } else {
  261. u.l.WithError(err).Error("Failed to set listen.so_mark")
  262. }
  263. }
  264. }
  265. func (u *StdConn) getMemInfo(meminfo *[unix.SK_MEMINFO_VARS]uint32) error {
  266. var vallen uint32 = 4 * unix.SK_MEMINFO_VARS
  267. _, _, err := unix.Syscall6(unix.SYS_GETSOCKOPT, uintptr(u.sysFd), uintptr(unix.SOL_SOCKET), uintptr(unix.SO_MEMINFO), uintptr(unsafe.Pointer(meminfo)), uintptr(unsafe.Pointer(&vallen)), 0)
  268. if err != 0 {
  269. return err
  270. }
  271. return nil
  272. }
  273. func (u *StdConn) Close() error {
  274. return syscall.Close(u.sysFd)
  275. }
  276. func NewUDPStatsEmitter(udpConns []Conn) func() {
  277. // Check if our kernel supports SO_MEMINFO before registering the gauges
  278. var udpGauges [][unix.SK_MEMINFO_VARS]metrics.Gauge
  279. var meminfo [unix.SK_MEMINFO_VARS]uint32
  280. if err := udpConns[0].(*StdConn).getMemInfo(&meminfo); err == nil {
  281. udpGauges = make([][unix.SK_MEMINFO_VARS]metrics.Gauge, len(udpConns))
  282. for i := range udpConns {
  283. udpGauges[i] = [unix.SK_MEMINFO_VARS]metrics.Gauge{
  284. metrics.GetOrRegisterGauge(fmt.Sprintf("udp.%d.rmem_alloc", i), nil),
  285. metrics.GetOrRegisterGauge(fmt.Sprintf("udp.%d.rcvbuf", i), nil),
  286. metrics.GetOrRegisterGauge(fmt.Sprintf("udp.%d.wmem_alloc", i), nil),
  287. metrics.GetOrRegisterGauge(fmt.Sprintf("udp.%d.sndbuf", i), nil),
  288. metrics.GetOrRegisterGauge(fmt.Sprintf("udp.%d.fwd_alloc", i), nil),
  289. metrics.GetOrRegisterGauge(fmt.Sprintf("udp.%d.wmem_queued", i), nil),
  290. metrics.GetOrRegisterGauge(fmt.Sprintf("udp.%d.optmem", i), nil),
  291. metrics.GetOrRegisterGauge(fmt.Sprintf("udp.%d.backlog", i), nil),
  292. metrics.GetOrRegisterGauge(fmt.Sprintf("udp.%d.drops", i), nil),
  293. }
  294. }
  295. }
  296. return func() {
  297. for i, gauges := range udpGauges {
  298. if err := udpConns[i].(*StdConn).getMemInfo(&meminfo); err == nil {
  299. for j := 0; j < unix.SK_MEMINFO_VARS; j++ {
  300. gauges[j].Update(int64(meminfo[j]))
  301. }
  302. }
  303. }
  304. }
  305. }