netstack.go 6.4 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226
  1. // Package amneziawgnet embeds amneziawg-go and gVisor netstack in-process
  2. // as a userspace alternative to kernel wireguard / awg-quick.
  3. package amneziawgnet
  4. import (
  5. "fmt"
  6. "net/netip"
  7. "os"
  8. "sync"
  9. "syscall"
  10. awgtun "github.com/amnezia-vpn/amneziawg-go/v3/tun"
  11. "gvisor.dev/gvisor/pkg/buffer"
  12. "gvisor.dev/gvisor/pkg/tcpip"
  13. "gvisor.dev/gvisor/pkg/tcpip/header"
  14. "gvisor.dev/gvisor/pkg/tcpip/link/channel"
  15. "gvisor.dev/gvisor/pkg/tcpip/network/ipv4"
  16. "gvisor.dev/gvisor/pkg/tcpip/network/ipv6"
  17. "gvisor.dev/gvisor/pkg/tcpip/stack"
  18. "gvisor.dev/gvisor/pkg/tcpip/transport/icmp"
  19. "gvisor.dev/gvisor/pkg/tcpip/transport/tcp"
  20. "gvisor.dev/gvisor/pkg/tcpip/transport/udp"
  21. )
  22. // tunQueueDepth is the outbound queue depth for channel endpoint and handoff.
  23. // 1024 starved simultaneous TCP slow-starts; channel.Endpoint drops silently when full.
  24. const tunQueueDepth = 8192
  25. // stackTun implements amneziawg-go tun.Device over a gVisor channel endpoint,
  26. // exposing *stack.Stack for forwarder attachment.
  27. type stackTun struct {
  28. ep *channel.Endpoint
  29. stack *stack.Stack
  30. events chan awgtun.Event
  31. notifyHandle *channel.NotificationHandle
  32. incomingPacket chan *buffer.View
  33. done chan struct{}
  34. closeMu sync.Mutex
  35. closed bool
  36. mtu int
  37. }
  38. // createNetTUNWithStack builds a gVisor-backed tun.Device for localAddresses
  39. // and returns underlying *stack.Stack to attach forwarders.
  40. func createNetTUNWithStack(localAddresses []netip.Addr, mtu int) (awgtun.Device, *stack.Stack, error) {
  41. opts := stack.Options{
  42. NetworkProtocols: []stack.NetworkProtocolFactory{ipv4.NewProtocol, ipv6.NewProtocol},
  43. TransportProtocols: []stack.TransportProtocolFactory{tcp.NewProtocol, udp.NewProtocol, icmp.NewProtocol6, icmp.NewProtocol4},
  44. // HandleLocal stays false so non-local destinations reach forwarder.
  45. HandleLocal: false,
  46. }
  47. dev := &stackTun{
  48. // tunQueueDepth buffers channel.New and incomingPacket for pipelining.
  49. ep: channel.New(tunQueueDepth, uint32(mtu), ""),
  50. stack: stack.New(opts),
  51. events: make(chan awgtun.Event, 10),
  52. incomingPacket: make(chan *buffer.View, tunQueueDepth),
  53. done: make(chan struct{}),
  54. mtu: mtu,
  55. }
  56. sackEnabledOpt := tcpip.TCPSACKEnabled(true)
  57. if err := dev.stack.SetTransportProtocolOption(tcp.ProtocolNumber, &sackEnabledOpt); err != nil {
  58. return nil, nil, fmt.Errorf("amneziawgnet: enable TCP SACK: %s", err)
  59. }
  60. dev.notifyHandle = dev.ep.AddNotify(dev)
  61. if err := dev.stack.CreateNIC(1, dev.ep); err != nil {
  62. return nil, nil, fmt.Errorf("amneziawgnet: CreateNIC: %s", err)
  63. }
  64. var hasV4, hasV6 bool
  65. for _, ip := range localAddresses {
  66. var protoNumber tcpip.NetworkProtocolNumber
  67. switch {
  68. case ip.Is4():
  69. protoNumber = ipv4.ProtocolNumber
  70. hasV4 = true
  71. case ip.Is6():
  72. protoNumber = ipv6.ProtocolNumber
  73. hasV6 = true
  74. default:
  75. continue
  76. }
  77. protoAddr := tcpip.ProtocolAddress{
  78. Protocol: protoNumber,
  79. AddressWithPrefix: tcpip.AddrFromSlice(ip.AsSlice()).WithPrefix(),
  80. }
  81. if err := dev.stack.AddProtocolAddress(1, protoAddr, stack.AddressProperties{}); err != nil {
  82. return nil, nil, fmt.Errorf("amneziawgnet: AddProtocolAddress(%v): %s", ip, err)
  83. }
  84. }
  85. if hasV4 {
  86. dev.stack.AddRoute(tcpip.Route{Destination: header.IPv4EmptySubnet, NIC: 1})
  87. }
  88. if hasV6 {
  89. dev.stack.AddRoute(tcpip.Route{Destination: header.IPv6EmptySubnet, NIC: 1})
  90. }
  91. dev.events <- awgtun.EventUp
  92. return dev, dev.stack, nil
  93. }
  94. func (t *stackTun) Name() (string, error) { return "amneziawgnet", nil }
  95. func (t *stackTun) File() *os.File { return nil }
  96. func (t *stackTun) Events() <-chan awgtun.Event { return t.events }
  97. func (t *stackTun) MTU() (int, error) { return t.mtu, nil }
  98. func (t *stackTun) BatchSize() int { return 1 }
  99. // Read drains incomingPacket into buf, supporting batched reads. Each view is
  100. // released once copied out, so the download path reuses gVisor's pooled chunks.
  101. func (t *stackTun) Read(buf [][]byte, sizes []int, offset int) (int, error) {
  102. var view *buffer.View
  103. select {
  104. case <-t.done:
  105. return 0, os.ErrClosed
  106. case view = <-t.incomingPacket:
  107. }
  108. n, err := view.Read(buf[0][offset:])
  109. view.Release()
  110. if err != nil {
  111. return 0, err
  112. }
  113. sizes[0] = n
  114. count := 1
  115. for count < len(buf) {
  116. select {
  117. case view = <-t.incomingPacket:
  118. n, err := view.Read(buf[count][offset:])
  119. view.Release()
  120. if err != nil {
  121. return count, nil
  122. }
  123. sizes[count] = n
  124. count++
  125. default:
  126. return count, nil
  127. }
  128. }
  129. return count, nil
  130. }
  131. // Write injects each packet into the stack. The injector owns the packet
  132. // buffer -- DecRef returns it and its chunk to gVisor's pools (see loopback.go).
  133. func (t *stackTun) Write(buf [][]byte, offset int) (int, error) {
  134. for _, b := range buf {
  135. packet := b[offset:]
  136. if len(packet) == 0 {
  137. continue
  138. }
  139. pkb := stack.NewPacketBuffer(stack.PacketBufferOptions{Payload: buffer.MakeWithData(packet)})
  140. switch packet[0] >> 4 {
  141. case 4:
  142. t.ep.InjectInbound(header.IPv4ProtocolNumber, pkb)
  143. case 6:
  144. t.ep.InjectInbound(header.IPv6ProtocolNumber, pkb)
  145. default:
  146. pkb.DecRef()
  147. return 0, syscall.EAFNOSUPPORT
  148. }
  149. pkb.DecRef()
  150. }
  151. return len(buf), nil
  152. }
  153. // WriteNotify runs on gVisor dispatch while Close tears the endpoint down,
  154. // so it must never block on closeMu across ep.Read or stack teardown.
  155. func (t *stackTun) WriteNotify() {
  156. t.closeMu.Lock()
  157. if t.closed {
  158. t.closeMu.Unlock()
  159. return
  160. }
  161. t.closeMu.Unlock()
  162. pkt := t.ep.Read()
  163. if pkt == nil {
  164. return
  165. }
  166. view := pkt.ToView()
  167. pkt.DecRef()
  168. // Select against done so racing dispatch abandons packet on close
  169. // without blocking Close or panicking on closed channel.
  170. select {
  171. case t.incomingPacket <- view:
  172. case <-t.done:
  173. view.Release()
  174. }
  175. }
  176. func (t *stackTun) Close() error {
  177. t.closeMu.Lock()
  178. if t.closed {
  179. t.closeMu.Unlock()
  180. return nil
  181. }
  182. t.closed = true
  183. close(t.done)
  184. t.closeMu.Unlock()
  185. t.stack.RemoveNIC(1)
  186. t.stack.Close()
  187. t.ep.RemoveNotify(t.notifyHandle)
  188. t.ep.Close()
  189. if t.events != nil {
  190. close(t.events)
  191. }
  192. return nil
  193. }
  194. // enablePromiscuousRouting configures NIC promiscuous and spoofing modes.
  195. func enablePromiscuousRouting(gstack *stack.Stack) {
  196. gstack.SetPromiscuousMode(1, true)
  197. gstack.SetSpoofing(1, true)
  198. }
  199. // addrFromTcpip converts a gVisor tcpip.Address to netip.Addr.
  200. func addrFromTcpip(a tcpip.Address) netip.Addr {
  201. if a.Len() == 4 {
  202. var b [4]byte
  203. copy(b[:], a.AsSlice())
  204. return netip.AddrFrom4(b)
  205. }
  206. var b [16]byte
  207. copy(b[:], a.AsSlice())
  208. return netip.AddrFrom16(b)
  209. }