udp.go 4.0 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102
  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. // ToSlice already returns an owned copy, so cloning pkt here would only
  31. // strand a pooled packet buffer and its chunks on every datagram.
  32. data := pkt.Data().AsRange().ToSlice()
  33. src := netip.AddrPortFrom(addrFromTcpip(id.RemoteAddress), id.RemotePort)
  34. dst := netip.AddrPortFrom(addrFromTcpip(id.LocalAddress), id.LocalPort)
  35. handler(src, dst, data)
  36. return true
  37. })
  38. }
  39. // WriteUDPReply injects a UDP packet into gstack as if it arrived from
  40. // `from` addressed to `to` -- i.e. a reply travelling back into the tunnel
  41. // toward the client -- constructed by hand since gVisor exposes no
  42. // connected-socket-style Write for an address the stack doesn't itself own.
  43. func WriteUDPReply(gstack *stack.Stack, from, to netip.AddrPort, payload []byte) error {
  44. udpLen := header.UDPMinimumSize + len(payload)
  45. srcIP := tcpip.AddrFromSlice(from.Addr().AsSlice())
  46. dstIP := tcpip.AddrFromSlice(to.Addr().AsSlice())
  47. isIPv4 := from.Addr().Is4()
  48. ipHdrSize := header.IPv6MinimumSize
  49. ipProtocol := header.IPv6ProtocolNumber
  50. if isIPv4 {
  51. ipHdrSize = header.IPv4MinimumSize
  52. ipProtocol = header.IPv4ProtocolNumber
  53. }
  54. pkt := stack.NewPacketBuffer(stack.PacketBufferOptions{
  55. ReserveHeaderBytes: ipHdrSize + header.UDPMinimumSize,
  56. Payload: buffer.MakeWithData(payload),
  57. })
  58. defer pkt.DecRef()
  59. udpHdr := header.UDP(pkt.TransportHeader().Push(header.UDPMinimumSize))
  60. udpHdr.Encode(&header.UDPFields{
  61. SrcPort: from.Port(),
  62. DstPort: to.Port(),
  63. Length: uint16(udpLen),
  64. })
  65. xsum := header.PseudoHeaderChecksum(header.UDPProtocolNumber, srcIP, dstIP, uint16(udpLen))
  66. udpHdr.SetChecksum(^udpHdr.CalculateChecksum(checksum.Checksum(payload, xsum)))
  67. if isIPv4 {
  68. ipHdr := header.IPv4(pkt.NetworkHeader().Push(header.IPv4MinimumSize))
  69. ipHdr.Encode(&header.IPv4Fields{
  70. TotalLength: uint16(header.IPv4MinimumSize + udpLen),
  71. TTL: 64,
  72. Protocol: uint8(header.UDPProtocolNumber),
  73. SrcAddr: srcIP,
  74. DstAddr: dstIP,
  75. })
  76. ipHdr.SetChecksum(^ipHdr.CalculateChecksum())
  77. } else {
  78. ipHdr := header.IPv6(pkt.NetworkHeader().Push(header.IPv6MinimumSize))
  79. ipHdr.Encode(&header.IPv6Fields{
  80. PayloadLength: uint16(udpLen),
  81. TransportProtocol: header.UDPProtocolNumber,
  82. HopLimit: 64,
  83. SrcAddr: srcIP,
  84. DstAddr: dstIP,
  85. })
  86. }
  87. if tcpipErr := gstack.WriteRawPacket(1, ipProtocol, buffer.MakeWithView(pkt.ToView())); tcpipErr != nil {
  88. return fmt.Errorf("amneziawgnet: WriteRawPacket: %s", tcpipErr)
  89. }
  90. return nil
  91. }