udp.go 3.8 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100
  1. package amneziawgnet
  2. import (
  3. "fmt"
  4. "net/netip"
  5. "gvisor.dev/gvisor/pkg/buffer"
  6. "gvisor.dev/gvisor/pkg/tcpip"
  7. "gvisor.dev/gvisor/pkg/tcpip/checksum"
  8. "gvisor.dev/gvisor/pkg/tcpip/header"
  9. "gvisor.dev/gvisor/pkg/tcpip/stack"
  10. "gvisor.dev/gvisor/pkg/tcpip/transport/udp"
  11. )
  12. // UDPHandler is called for every UDP packet a tunnel client sends, with its
  13. // source (the peer's tunnel-internal address) and its real,
  14. // dynamically-arbitrary destination -- recovered the same way the TCP
  15. // forwarder recovers its destination, from the packet's own transport
  16. // endpoint ID, never from a preconfigured table. The handler owns all flow
  17. // tracking and reply delivery (via WriteUDPReply): gVisor has no
  18. // udp.NewForwarder the way it does for TCP, so unlike AttachTCPForwarder
  19. // this can't just hand back a ready net.Conn.
  20. type UDPHandler func(src, dst netip.AddrPort, payload []byte)
  21. // AttachUDPHandler attaches a raw UDP handler to gstack, independently
  22. // enabling the same promiscuous+spoofing mode AttachTCPForwarder needs --
  23. // safe and idempotent to call regardless of whether AttachTCPForwarder was
  24. // attached to the same stack first, or at all. Adapted from xtls/xray-core's
  25. // proxy/wireguard/tun.go UDP path (MIT), which hand-tracks flows for the
  26. // identical reason: gVisor doesn't provide a UDP forwarder.
  27. func AttachUDPHandler(gstack *stack.Stack, handler UDPHandler) {
  28. enablePromiscuousRouting(gstack)
  29. gstack.SetTransportProtocolHandler(udp.ProtocolNumber, func(id stack.TransportEndpointID, pkt *stack.PacketBuffer) bool {
  30. data := pkt.Clone().Data().AsRange().ToSlice()
  31. src := netip.AddrPortFrom(addrFromTcpip(id.RemoteAddress), id.RemotePort)
  32. dst := netip.AddrPortFrom(addrFromTcpip(id.LocalAddress), id.LocalPort)
  33. handler(src, dst, data)
  34. return true
  35. })
  36. }
  37. // WriteUDPReply injects a UDP packet into gstack as if it arrived from
  38. // `from` addressed to `to` -- i.e. a reply travelling back into the tunnel
  39. // toward the client -- constructed by hand since gVisor exposes no
  40. // connected-socket-style Write for an address the stack doesn't itself own.
  41. func WriteUDPReply(gstack *stack.Stack, from, to netip.AddrPort, payload []byte) error {
  42. udpLen := header.UDPMinimumSize + len(payload)
  43. srcIP := tcpip.AddrFromSlice(from.Addr().AsSlice())
  44. dstIP := tcpip.AddrFromSlice(to.Addr().AsSlice())
  45. isIPv4 := from.Addr().Is4()
  46. ipHdrSize := header.IPv6MinimumSize
  47. ipProtocol := header.IPv6ProtocolNumber
  48. if isIPv4 {
  49. ipHdrSize = header.IPv4MinimumSize
  50. ipProtocol = header.IPv4ProtocolNumber
  51. }
  52. pkt := stack.NewPacketBuffer(stack.PacketBufferOptions{
  53. ReserveHeaderBytes: ipHdrSize + header.UDPMinimumSize,
  54. Payload: buffer.MakeWithData(payload),
  55. })
  56. defer pkt.DecRef()
  57. udpHdr := header.UDP(pkt.TransportHeader().Push(header.UDPMinimumSize))
  58. udpHdr.Encode(&header.UDPFields{
  59. SrcPort: from.Port(),
  60. DstPort: to.Port(),
  61. Length: uint16(udpLen),
  62. })
  63. xsum := header.PseudoHeaderChecksum(header.UDPProtocolNumber, srcIP, dstIP, uint16(udpLen))
  64. udpHdr.SetChecksum(^udpHdr.CalculateChecksum(checksum.Checksum(payload, xsum)))
  65. if isIPv4 {
  66. ipHdr := header.IPv4(pkt.NetworkHeader().Push(header.IPv4MinimumSize))
  67. ipHdr.Encode(&header.IPv4Fields{
  68. TotalLength: uint16(header.IPv4MinimumSize + udpLen),
  69. TTL: 64,
  70. Protocol: uint8(header.UDPProtocolNumber),
  71. SrcAddr: srcIP,
  72. DstAddr: dstIP,
  73. })
  74. ipHdr.SetChecksum(^ipHdr.CalculateChecksum())
  75. } else {
  76. ipHdr := header.IPv6(pkt.NetworkHeader().Push(header.IPv6MinimumSize))
  77. ipHdr.Encode(&header.IPv6Fields{
  78. PayloadLength: uint16(udpLen),
  79. TransportProtocol: header.UDPProtocolNumber,
  80. HopLimit: 64,
  81. SrcAddr: srcIP,
  82. DstAddr: dstIP,
  83. })
  84. }
  85. if tcpipErr := gstack.WriteRawPacket(1, ipProtocol, buffer.MakeWithView(pkt.ToView())); tcpipErr != nil {
  86. return fmt.Errorf("amneziawgnet: WriteRawPacket: %s", tcpipErr)
  87. }
  88. return nil
  89. }