packetinfo_linux.go 2.3 KB

12345678910111213141516171819202122232425262728293031323334353637383940414243444546474849505152535455565758596061626364656667686970717273747576777879808182838485868788
  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. })
  18. if err != nil {
  19. return err
  20. }
  21. return controlErr
  22. }
  23. func readFromUDP(conn *net.UDPConn, buf []byte) (int, *net.UDPAddr, udpPacketInfo, error) {
  24. oob := make([]byte, 128)
  25. n, oobn, _, remote, err := conn.ReadMsgUDP(buf, oob)
  26. if err != nil {
  27. return 0, nil, udpPacketInfo{}, err
  28. }
  29. packetInfo := parsePacketInfo(oob[:oobn])
  30. return n, remote, packetInfo, nil
  31. }
  32. func writeToUDP(conn *net.UDPConn, payload []byte, remote *net.UDPAddr, packetInfo udpPacketInfo, localIP net.IP) (int, error) {
  33. ipv4 := localIP.To4()
  34. if packetInfo.ifIndex <= 0 || ipv4 == nil {
  35. return conn.WriteToUDP(payload, remote)
  36. }
  37. oob := make([]byte, syscall.CmsgSpace(syscall.SizeofInet4Pktinfo))
  38. header := (*syscall.Cmsghdr)(unsafe.Pointer(&oob[0]))
  39. header.Level = syscall.IPPROTO_IP
  40. header.Type = syscall.IP_PKTINFO
  41. header.SetLen(syscall.CmsgLen(syscall.SizeofInet4Pktinfo))
  42. info := (*syscall.Inet4Pktinfo)(unsafe.Pointer(&oob[syscall.CmsgLen(0)]))
  43. info.Ifindex = int32(packetInfo.ifIndex)
  44. copy(info.Spec_dst[:], ipv4)
  45. n, _, err := conn.WriteMsgUDP(payload, oob, remote)
  46. return n, err
  47. }
  48. func parsePacketInfo(oob []byte) udpPacketInfo {
  49. messages, err := syscall.ParseSocketControlMessage(oob)
  50. if err != nil {
  51. return udpPacketInfo{}
  52. }
  53. for _, message := range messages {
  54. if message.Header.Level != syscall.IPPROTO_IP || message.Header.Type != syscall.IP_PKTINFO || len(message.Data) < 12 {
  55. continue
  56. }
  57. specDst := [4]byte(message.Data[4:8])
  58. addr := [4]byte(message.Data[8:12])
  59. return udpPacketInfo{
  60. localIP: packetInfoIP(specDst, addr),
  61. ifIndex: int(int32(binary.LittleEndian.Uint32(message.Data[0:4]))),
  62. }
  63. }
  64. return udpPacketInfo{}
  65. }
  66. func packetInfoIP(specDst [4]byte, addr [4]byte) net.IP {
  67. if ip := ipv4FromBytes(specDst); ip != nil {
  68. return ip
  69. }
  70. return ipv4FromBytes(addr)
  71. }
  72. func ipv4FromBytes(value [4]byte) net.IP {
  73. if value == [4]byte{} || value == [4]byte{255, 255, 255, 255} {
  74. return nil
  75. }
  76. return net.IPv4(value[0], value[1], value[2], value[3]).To4()
  77. }