3
0

udp_linux.go 8.6 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363
  1. //go:build !android && !e2e_testing
  2. // +build !android,!e2e_testing
  3. package udp
  4. import (
  5. "encoding/binary"
  6. "fmt"
  7. "net"
  8. "syscall"
  9. "unsafe"
  10. "github.com/rcrowley/go-metrics"
  11. "github.com/sirupsen/logrus"
  12. "github.com/slackhq/nebula/config"
  13. "github.com/slackhq/nebula/firewall"
  14. "github.com/slackhq/nebula/header"
  15. "golang.org/x/sys/unix"
  16. )
  17. //TODO: make it support reload as best you can!
  18. type StdConn struct {
  19. sysFd int
  20. isV4 bool
  21. l *logrus.Logger
  22. batch int
  23. }
  24. var x int
  25. // From linux/sock_diag.h
  26. const (
  27. _SK_MEMINFO_RMEM_ALLOC = iota
  28. _SK_MEMINFO_RCVBUF
  29. _SK_MEMINFO_WMEM_ALLOC
  30. _SK_MEMINFO_SNDBUF
  31. _SK_MEMINFO_FWD_ALLOC
  32. _SK_MEMINFO_WMEM_QUEUED
  33. _SK_MEMINFO_OPTMEM
  34. _SK_MEMINFO_BACKLOG
  35. _SK_MEMINFO_DROPS
  36. _SK_MEMINFO_VARS
  37. )
  38. type _SK_MEMINFO [_SK_MEMINFO_VARS]uint32
  39. func maybeIPV4(ip net.IP) (net.IP, bool) {
  40. ip4 := ip.To4()
  41. if ip4 != nil {
  42. return ip4, true
  43. }
  44. return ip, false
  45. }
  46. func NewListener(l *logrus.Logger, ip net.IP, port int, multi bool, batch int) (Conn, error) {
  47. ipV4, isV4 := maybeIPV4(ip)
  48. af := unix.AF_INET6
  49. if isV4 {
  50. af = unix.AF_INET
  51. }
  52. syscall.ForkLock.RLock()
  53. fd, err := unix.Socket(af, unix.SOCK_DGRAM, unix.IPPROTO_UDP)
  54. if err == nil {
  55. unix.CloseOnExec(fd)
  56. }
  57. syscall.ForkLock.RUnlock()
  58. if err != nil {
  59. unix.Close(fd)
  60. return nil, fmt.Errorf("unable to open socket: %s", err)
  61. }
  62. if multi {
  63. if err = unix.SetsockoptInt(fd, unix.SOL_SOCKET, unix.SO_REUSEPORT, 1); err != nil {
  64. return nil, fmt.Errorf("unable to set SO_REUSEPORT: %s", err)
  65. }
  66. }
  67. //TODO: support multiple listening IPs (for limiting ipv6)
  68. var sa unix.Sockaddr
  69. if isV4 {
  70. sa4 := &unix.SockaddrInet4{Port: port}
  71. copy(sa4.Addr[:], ipV4)
  72. sa = sa4
  73. } else {
  74. sa6 := &unix.SockaddrInet6{Port: port}
  75. copy(sa6.Addr[:], ip.To16())
  76. sa = sa6
  77. }
  78. if err = unix.Bind(fd, sa); err != nil {
  79. return nil, fmt.Errorf("unable to bind to socket: %s", err)
  80. }
  81. //TODO: this may be useful for forcing threads into specific cores
  82. //unix.SetsockoptInt(fd, unix.SOL_SOCKET, unix.SO_INCOMING_CPU, x)
  83. //v, err := unix.GetsockoptInt(fd, unix.SOL_SOCKET, unix.SO_INCOMING_CPU)
  84. //l.Println(v, err)
  85. return &StdConn{sysFd: fd, isV4: isV4, l: l, batch: batch}, err
  86. }
  87. func (u *StdConn) Rebind() error {
  88. return nil
  89. }
  90. func (u *StdConn) SetRecvBuffer(n int) error {
  91. return unix.SetsockoptInt(u.sysFd, unix.SOL_SOCKET, unix.SO_RCVBUFFORCE, n)
  92. }
  93. func (u *StdConn) SetSendBuffer(n int) error {
  94. return unix.SetsockoptInt(u.sysFd, unix.SOL_SOCKET, unix.SO_SNDBUFFORCE, n)
  95. }
  96. func (u *StdConn) GetRecvBuffer() (int, error) {
  97. return unix.GetsockoptInt(int(u.sysFd), unix.SOL_SOCKET, unix.SO_RCVBUF)
  98. }
  99. func (u *StdConn) GetSendBuffer() (int, error) {
  100. return unix.GetsockoptInt(int(u.sysFd), unix.SOL_SOCKET, unix.SO_SNDBUF)
  101. }
  102. func (u *StdConn) LocalAddr() (*Addr, error) {
  103. sa, err := unix.Getsockname(u.sysFd)
  104. if err != nil {
  105. return nil, err
  106. }
  107. addr := &Addr{}
  108. switch sa := sa.(type) {
  109. case *unix.SockaddrInet4:
  110. addr.IP = net.IP{sa.Addr[0], sa.Addr[1], sa.Addr[2], sa.Addr[3]}.To16()
  111. addr.Port = uint16(sa.Port)
  112. case *unix.SockaddrInet6:
  113. addr.IP = sa.Addr[0:]
  114. addr.Port = uint16(sa.Port)
  115. }
  116. return addr, nil
  117. }
  118. func (u *StdConn) ListenOut(r EncReader, lhf LightHouseHandlerFunc, cache *firewall.ConntrackCacheTicker, q int) {
  119. plaintext := make([]byte, MTU)
  120. h := &header.H{}
  121. fwPacket := &firewall.Packet{}
  122. udpAddr := &Addr{}
  123. nb := make([]byte, 12, 12)
  124. //TODO: should we track this?
  125. //metric := metrics.GetOrRegisterHistogram("test.batch_read", nil, metrics.NewExpDecaySample(1028, 0.015))
  126. msgs, buffers, names := u.PrepareRawMessages(u.batch)
  127. read := u.ReadMulti
  128. if u.batch == 1 {
  129. read = u.ReadSingle
  130. }
  131. for {
  132. n, err := read(msgs)
  133. if err != nil {
  134. u.l.WithError(err).Debug("udp socket is closed, exiting read loop")
  135. return
  136. }
  137. //metric.Update(int64(n))
  138. for i := 0; i < n; i++ {
  139. if u.isV4 {
  140. udpAddr.IP = names[i][4:8]
  141. } else {
  142. udpAddr.IP = names[i][8:24]
  143. }
  144. udpAddr.Port = binary.BigEndian.Uint16(names[i][2:4])
  145. r(udpAddr, plaintext[:0], buffers[i][:msgs[i].Len], h, fwPacket, lhf, nb, q, cache.Get(u.l))
  146. }
  147. }
  148. }
  149. func (u *StdConn) ReadSingle(msgs []rawMessage) (int, error) {
  150. for {
  151. n, _, err := unix.Syscall6(
  152. unix.SYS_RECVMSG,
  153. uintptr(u.sysFd),
  154. uintptr(unsafe.Pointer(&(msgs[0].Hdr))),
  155. 0,
  156. 0,
  157. 0,
  158. 0,
  159. )
  160. if err != 0 {
  161. return 0, &net.OpError{Op: "recvmsg", Err: err}
  162. }
  163. msgs[0].Len = uint32(n)
  164. return 1, nil
  165. }
  166. }
  167. func (u *StdConn) ReadMulti(msgs []rawMessage) (int, error) {
  168. for {
  169. n, _, err := unix.Syscall6(
  170. unix.SYS_RECVMMSG,
  171. uintptr(u.sysFd),
  172. uintptr(unsafe.Pointer(&msgs[0])),
  173. uintptr(len(msgs)),
  174. unix.MSG_WAITFORONE,
  175. 0,
  176. 0,
  177. )
  178. if err != 0 {
  179. return 0, &net.OpError{Op: "recvmmsg", Err: err}
  180. }
  181. return int(n), nil
  182. }
  183. }
  184. func (u *StdConn) WriteTo(b []byte, addr *Addr) error {
  185. if u.isV4 {
  186. return u.writeTo4(b, addr)
  187. }
  188. return u.writeTo6(b, addr)
  189. }
  190. func (u *StdConn) writeTo6(b []byte, addr *Addr) error {
  191. var rsa unix.RawSockaddrInet6
  192. rsa.Family = unix.AF_INET6
  193. // Little Endian -> Network Endian
  194. rsa.Port = (addr.Port >> 8) | ((addr.Port & 0xff) << 8)
  195. copy(rsa.Addr[:], addr.IP.To16())
  196. for {
  197. _, _, err := unix.Syscall6(
  198. unix.SYS_SENDTO,
  199. uintptr(u.sysFd),
  200. uintptr(unsafe.Pointer(&b[0])),
  201. uintptr(len(b)),
  202. uintptr(0),
  203. uintptr(unsafe.Pointer(&rsa)),
  204. uintptr(unix.SizeofSockaddrInet6),
  205. )
  206. if err != 0 {
  207. return &net.OpError{Op: "sendto", Err: err}
  208. }
  209. //TODO: handle incomplete writes
  210. return nil
  211. }
  212. }
  213. func (u *StdConn) writeTo4(b []byte, addr *Addr) error {
  214. addrV4, isAddrV4 := maybeIPV4(addr.IP)
  215. if !isAddrV4 {
  216. return fmt.Errorf("Listener is IPv4, but writing to IPv6 remote")
  217. }
  218. var rsa unix.RawSockaddrInet4
  219. rsa.Family = unix.AF_INET
  220. // Little Endian -> Network Endian
  221. rsa.Port = (addr.Port >> 8) | ((addr.Port & 0xff) << 8)
  222. copy(rsa.Addr[:], addrV4)
  223. for {
  224. _, _, err := unix.Syscall6(
  225. unix.SYS_SENDTO,
  226. uintptr(u.sysFd),
  227. uintptr(unsafe.Pointer(&b[0])),
  228. uintptr(len(b)),
  229. uintptr(0),
  230. uintptr(unsafe.Pointer(&rsa)),
  231. uintptr(unix.SizeofSockaddrInet4),
  232. )
  233. if err != 0 {
  234. return &net.OpError{Op: "sendto", Err: err}
  235. }
  236. //TODO: handle incomplete writes
  237. return nil
  238. }
  239. }
  240. func (u *StdConn) ReloadConfig(c *config.C) {
  241. b := c.GetInt("listen.read_buffer", 0)
  242. if b > 0 {
  243. err := u.SetRecvBuffer(b)
  244. if err == nil {
  245. s, err := u.GetRecvBuffer()
  246. if err == nil {
  247. u.l.WithField("size", s).Info("listen.read_buffer was set")
  248. } else {
  249. u.l.WithError(err).Warn("Failed to get listen.read_buffer")
  250. }
  251. } else {
  252. u.l.WithError(err).Error("Failed to set listen.read_buffer")
  253. }
  254. }
  255. b = c.GetInt("listen.write_buffer", 0)
  256. if b > 0 {
  257. err := u.SetSendBuffer(b)
  258. if err == nil {
  259. s, err := u.GetSendBuffer()
  260. if err == nil {
  261. u.l.WithField("size", s).Info("listen.write_buffer was set")
  262. } else {
  263. u.l.WithError(err).Warn("Failed to get listen.write_buffer")
  264. }
  265. } else {
  266. u.l.WithError(err).Error("Failed to set listen.write_buffer")
  267. }
  268. }
  269. }
  270. func (u *StdConn) getMemInfo(meminfo *_SK_MEMINFO) error {
  271. var vallen uint32 = 4 * _SK_MEMINFO_VARS
  272. _, _, 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)
  273. if err != 0 {
  274. return err
  275. }
  276. return nil
  277. }
  278. func (u *StdConn) Close() error {
  279. //TODO: this will not interrupt the read loop
  280. return syscall.Close(u.sysFd)
  281. }
  282. func NewUDPStatsEmitter(udpConns []Conn) func() {
  283. // Check if our kernel supports SO_MEMINFO before registering the gauges
  284. var udpGauges [][_SK_MEMINFO_VARS]metrics.Gauge
  285. var meminfo _SK_MEMINFO
  286. if err := udpConns[0].(*StdConn).getMemInfo(&meminfo); err == nil {
  287. udpGauges = make([][_SK_MEMINFO_VARS]metrics.Gauge, len(udpConns))
  288. for i := range udpConns {
  289. udpGauges[i] = [_SK_MEMINFO_VARS]metrics.Gauge{
  290. metrics.GetOrRegisterGauge(fmt.Sprintf("udp.%d.rmem_alloc", i), nil),
  291. metrics.GetOrRegisterGauge(fmt.Sprintf("udp.%d.rcvbuf", i), nil),
  292. metrics.GetOrRegisterGauge(fmt.Sprintf("udp.%d.wmem_alloc", i), nil),
  293. metrics.GetOrRegisterGauge(fmt.Sprintf("udp.%d.sndbuf", i), nil),
  294. metrics.GetOrRegisterGauge(fmt.Sprintf("udp.%d.fwd_alloc", i), nil),
  295. metrics.GetOrRegisterGauge(fmt.Sprintf("udp.%d.wmem_queued", i), nil),
  296. metrics.GetOrRegisterGauge(fmt.Sprintf("udp.%d.optmem", i), nil),
  297. metrics.GetOrRegisterGauge(fmt.Sprintf("udp.%d.backlog", i), nil),
  298. metrics.GetOrRegisterGauge(fmt.Sprintf("udp.%d.drops", i), nil),
  299. }
  300. }
  301. }
  302. return func() {
  303. for i, gauges := range udpGauges {
  304. if err := udpConns[i].(*StdConn).getMemInfo(&meminfo); err == nil {
  305. for j := 0; j < _SK_MEMINFO_VARS; j++ {
  306. gauges[j].Update(int64(meminfo[j]))
  307. }
  308. }
  309. }
  310. }
  311. }