packetinfo_linux.go 2.4 KB

1234567891011121314151617181920212223242526272829303132333435363738394041424344454647484950515253545556575859606162636465666768697071727374757677787980818283848586878889909192
  1. //go:build linux
  2. package discovery
  3. import (
  4. "encoding/binary"
  5. "net"
  6. "syscall"
  7. "unsafe"
  8. )
  9. func enableLocalAddrControl(conn *net.UDPConn) error {
  10. rawConn, err := conn.SyscallConn()
  11. if err != nil {
  12. return err
  13. }
  14. var controlErr error
  15. err = rawConn.Control(func(fd uintptr) {
  16. controlErr = syscall.SetsockoptInt(int(fd), syscall.IPPROTO_IP, syscall.IP_PKTINFO, 1)
  17. if controlErr == nil {
  18. controlErr = syscall.SetsockoptInt(int(fd), syscall.SOL_SOCKET, syscall.SO_BROADCAST, 1)
  19. }
  20. })
  21. if err != nil {
  22. return err
  23. }
  24. return controlErr
  25. }
  26. func readFromUDP(conn *net.UDPConn, buf []byte) (int, *net.UDPAddr, udpPacketInfo, error) {
  27. oob := make([]byte, 128)
  28. n, oobn, _, remote, err := conn.ReadMsgUDP(buf, oob)
  29. if err != nil {
  30. return 0, nil, udpPacketInfo{}, err
  31. }
  32. packetInfo := parsePacketInfo(oob[:oobn])
  33. return n, remote, packetInfo, nil
  34. }
  35. func writeToUDP(conn *net.UDPConn, payload []byte, remote *net.UDPAddr, packetInfo udpPacketInfo, localIP net.IP) (int, error) {
  36. if packetInfo.ifIndex <= 0 {
  37. return conn.WriteToUDP(payload, remote)
  38. }
  39. oob := make([]byte, syscall.CmsgSpace(syscall.SizeofInet4Pktinfo))
  40. header := (*syscall.Cmsghdr)(unsafe.Pointer(&oob[0]))
  41. header.Level = syscall.IPPROTO_IP
  42. header.Type = syscall.IP_PKTINFO
  43. header.SetLen(syscall.CmsgLen(syscall.SizeofInet4Pktinfo))
  44. info := (*syscall.Inet4Pktinfo)(unsafe.Pointer(&oob[syscall.CmsgLen(0)]))
  45. info.Ifindex = int32(packetInfo.ifIndex)
  46. if ipv4 := localIP.To4(); ipv4 != nil {
  47. copy(info.Spec_dst[:], ipv4)
  48. }
  49. n, _, err := conn.WriteMsgUDP(payload, oob, remote)
  50. return n, err
  51. }
  52. func parsePacketInfo(oob []byte) udpPacketInfo {
  53. messages, err := syscall.ParseSocketControlMessage(oob)
  54. if err != nil {
  55. return udpPacketInfo{}
  56. }
  57. for _, message := range messages {
  58. if message.Header.Level != syscall.IPPROTO_IP || message.Header.Type != syscall.IP_PKTINFO || len(message.Data) < 12 {
  59. continue
  60. }
  61. specDst := [4]byte(message.Data[4:8])
  62. addr := [4]byte(message.Data[8:12])
  63. return udpPacketInfo{
  64. localIP: packetInfoIP(specDst, addr),
  65. ifIndex: int(int32(binary.LittleEndian.Uint32(message.Data[0:4]))),
  66. }
  67. }
  68. return udpPacketInfo{}
  69. }
  70. func packetInfoIP(specDst [4]byte, addr [4]byte) net.IP {
  71. if ip := ipv4FromBytes(specDst); ip != nil {
  72. return ip
  73. }
  74. return ipv4FromBytes(addr)
  75. }
  76. func ipv4FromBytes(value [4]byte) net.IP {
  77. if value == [4]byte{} || value == [4]byte{255, 255, 255, 255} {
  78. return nil
  79. }
  80. return net.IPv4(value[0], value[1], value[2], value[3]).To4()
  81. }