1
0

portfwd_test.go 14 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427
  1. package amneziawgnet
  2. import (
  3. "context"
  4. "fmt"
  5. "io"
  6. "net"
  7. "net/netip"
  8. "sync"
  9. "testing"
  10. "time"
  11. awgconn "github.com/amnezia-vpn/amneziawg-go/v3/conn"
  12. "github.com/amnezia-vpn/amneziawg-go/v3/device"
  13. "github.com/amnezia-vpn/amneziawg-go/v3/tun/netstack"
  14. "gvisor.dev/gvisor/pkg/tcpip/stack"
  15. "github.com/mhsanaei/3x-ui/v3/internal/amneziawg"
  16. "github.com/mhsanaei/3x-ui/v3/internal/util/wireguard"
  17. )
  18. func peerWithPortsAndIPs(email, forwardedPorts string, ips ...string) amneziawg.Peer {
  19. return amneziawg.Peer{Email: email, PublicKey: "pub-" + email, AllowedIPs: ips, ForwardedPorts: forwardedPorts}
  20. }
  21. // --- desiredPeerTargets ---
  22. func TestDesiredPeerTargetsPrefersIPv4(t *testing.T) {
  23. inst := amneziawg.Instance{IPv6Enabled: true, Peers: []amneziawg.Peer{
  24. peerWithPortsAndIPs("a@x", "", "10.8.1.2/32", "fd86::2/128"),
  25. }}
  26. got := desiredPeerTargets(inst)
  27. addr, ok := got["a@x"]
  28. if !ok || addr.String() != "10.8.1.2" {
  29. t.Fatalf("desiredPeerTargets = %v, want a@x -> 10.8.1.2", got)
  30. }
  31. }
  32. func TestDesiredPeerTargetsFallsBackToIPv6WhenEnabled(t *testing.T) {
  33. inst := amneziawg.Instance{IPv6Enabled: true, Peers: []amneziawg.Peer{
  34. peerWithPortsAndIPs("a@x", "", "fd86::2/128"),
  35. }}
  36. got := desiredPeerTargets(inst)
  37. addr, ok := got["a@x"]
  38. if !ok || addr.String() != "fd86::2" {
  39. t.Fatalf("desiredPeerTargets = %v, want a@x -> fd86::2", got)
  40. }
  41. }
  42. func TestDesiredPeerTargetsSkipsIPv6OnlyWhenIPv6Disabled(t *testing.T) {
  43. inst := amneziawg.Instance{IPv6Enabled: false, Peers: []amneziawg.Peer{
  44. peerWithPortsAndIPs("a@x", "", "fd86::2/128"),
  45. }}
  46. if got := desiredPeerTargets(inst); len(got) != 0 {
  47. t.Fatalf("desiredPeerTargets = %v, want empty (IPv6-only peer, IPv6 disabled)", got)
  48. }
  49. }
  50. func TestDesiredPeerTargetsSkipsPeerWithoutEmailOrAddress(t *testing.T) {
  51. inst := amneziawg.Instance{IPv6Enabled: true, Peers: []amneziawg.Peer{
  52. peerWithPortsAndIPs("", "", "10.8.1.2/32"), // no email
  53. peerWithPortsAndIPs("b@x", ""), // no AllowedIPs at all
  54. }}
  55. if got := desiredPeerTargets(inst); len(got) != 0 {
  56. t.Fatalf("desiredPeerTargets = %v, want empty", got)
  57. }
  58. }
  59. // --- desiredPortForwardKeys ---
  60. func TestDesiredPortForwardKeysEmptyWhenNoForwardedPorts(t *testing.T) {
  61. inst := amneziawg.Instance{Peers: []amneziawg.Peer{
  62. peerWithPortsAndIPs("a@x", "", "10.8.1.2/32"),
  63. }}
  64. if got := desiredPortForwardKeys(inst); len(got) != 0 {
  65. t.Fatalf("desiredPortForwardKeys = %v, want empty", got)
  66. }
  67. }
  68. func TestDesiredPortForwardKeysEmptyWhenNoResolvableTarget(t *testing.T) {
  69. // ForwardedPorts is set, but the peer has no AllowedIPs to resolve a
  70. // target from -- must not produce keys for a peer nothing can dial.
  71. inst := amneziawg.Instance{Peers: []amneziawg.Peer{
  72. {Email: "a@x", ForwardedPorts: "8080"},
  73. }}
  74. if got := desiredPortForwardKeys(inst); len(got) != 0 {
  75. t.Fatalf("desiredPortForwardKeys = %v, want empty", got)
  76. }
  77. }
  78. func TestDesiredPortForwardKeysOneTCPAndUDPKeyPerPort(t *testing.T) {
  79. inst := amneziawg.Instance{Peers: []amneziawg.Peer{
  80. peerWithPortsAndIPs("a@x", "8080,8081", "10.8.1.2/32"),
  81. }}
  82. got := desiredPortForwardKeys(inst)
  83. if len(got) != 4 {
  84. t.Fatalf("desiredPortForwardKeys = %v, want 4 entries (2 ports x 2 protocols)", got)
  85. }
  86. for _, port := range []int{8080, 8081} {
  87. for _, proto := range []portForwardProto{tcpForward, udpForward} {
  88. key := portForwardKey{email: "a@x", port: port, proto: proto}
  89. if _, ok := got[key]; !ok {
  90. t.Errorf("desiredPortForwardKeys missing %+v", key)
  91. }
  92. }
  93. }
  94. }
  95. func TestDesiredPortForwardKeysMultiplePeersDoNotMix(t *testing.T) {
  96. inst := amneziawg.Instance{Peers: []amneziawg.Peer{
  97. peerWithPortsAndIPs("a@x", "8080", "10.8.1.2/32"),
  98. peerWithPortsAndIPs("b@x", "8080", "10.8.1.3/32"), // same port, different peer
  99. }}
  100. got := desiredPortForwardKeys(inst)
  101. if len(got) != 4 {
  102. t.Fatalf("desiredPortForwardKeys = %v, want 4 entries (2 peers x 2 protocols, same port kept separate per email)", got)
  103. }
  104. }
  105. // --- PortForwardSet.Reconcile: real stack, no handshake needed (dialing
  106. // isn't exercised by these -- only the host-facing listener lifecycle) ---
  107. func newTestStack(t *testing.T, addr string) *stack.Stack {
  108. t.Helper()
  109. tunDev, gstack, err := createNetTUNWithStack([]netip.Addr{netip.MustParseAddr(addr)}, 1420)
  110. if err != nil {
  111. t.Fatalf("createNetTUNWithStack: %v", err)
  112. }
  113. t.Cleanup(func() { tunDev.Close() })
  114. return gstack
  115. }
  116. func dialLoopback(t *testing.T, network string, port int) {
  117. t.Helper()
  118. conn, err := net.DialTimeout(network, fmt.Sprintf("127.0.0.1:%d", port), time.Second)
  119. if err != nil {
  120. t.Fatalf("dial 127.0.0.1:%d (%s): %v", port, network, err)
  121. }
  122. conn.Close()
  123. }
  124. func TestPortForwardSetReconcileOpensAndClosesListeners(t *testing.T) {
  125. gs := newTestStack(t, "10.211.0.1")
  126. set := NewPortForwardSet(gs, 501)
  127. const port = 58910
  128. inst := amneziawg.Instance{Peers: []amneziawg.Peer{
  129. peerWithPortsAndIPs("a@x", fmt.Sprintf("%d", port), "10.211.0.2/32"),
  130. }}
  131. set.Reconcile(inst)
  132. set.mu.Lock()
  133. n := len(set.listeners)
  134. set.mu.Unlock()
  135. if n != 2 {
  136. t.Fatalf("listeners after Reconcile = %d, want 2 (tcp+udp)", n)
  137. }
  138. dialLoopback(t, "tcp", port) // proves a real host listener is actually bound
  139. set.mu.Lock()
  140. tcpBefore := set.listeners[portForwardKey{email: "a@x", port: port, proto: tcpForward}]
  141. set.mu.Unlock()
  142. // Reconciling again with an unchanged instance must not close and
  143. // reopen an unaffected listener.
  144. set.Reconcile(inst)
  145. set.mu.Lock()
  146. tcpAfter := set.listeners[portForwardKey{email: "a@x", port: port, proto: tcpForward}]
  147. set.mu.Unlock()
  148. if tcpBefore != tcpAfter {
  149. t.Error("Reconcile with an unchanged instance replaced an unaffected listener")
  150. }
  151. // Peer removed entirely -> both listeners close.
  152. set.Reconcile(amneziawg.Instance{})
  153. set.mu.Lock()
  154. n = len(set.listeners)
  155. set.mu.Unlock()
  156. if n != 0 {
  157. t.Fatalf("listeners after removal Reconcile = %d, want 0", n)
  158. }
  159. if _, err := net.DialTimeout("tcp", fmt.Sprintf("127.0.0.1:%d", port), time.Second); err == nil {
  160. t.Error("port still accepting connections after the listener should have closed")
  161. }
  162. }
  163. func TestPortForwardSetReconcileSurvivesPreBoundPort(t *testing.T) {
  164. gs := newTestStack(t, "10.211.1.1")
  165. set := NewPortForwardSet(gs, 502)
  166. const collidingPort = 58911
  167. const okPort = 58912
  168. blocker, err := net.Listen("tcp", fmt.Sprintf(":%d", collidingPort))
  169. if err != nil {
  170. t.Fatalf("pre-bind test port: %v", err)
  171. }
  172. defer blocker.Close()
  173. inst := amneziawg.Instance{Peers: []amneziawg.Peer{
  174. peerWithPortsAndIPs("a@x", fmt.Sprintf("%d,%d", collidingPort, okPort), "10.211.1.2/32"),
  175. }}
  176. // Must not panic despite one of the two ports being unbindable, and the
  177. // other port (and its UDP counterpart on the colliding port) must still
  178. // open normally.
  179. set.Reconcile(inst)
  180. set.mu.Lock()
  181. n := len(set.listeners)
  182. _, tcpCollidingOpen := set.listeners[portForwardKey{email: "a@x", port: collidingPort, proto: tcpForward}]
  183. _, udpCollidingOpen := set.listeners[portForwardKey{email: "a@x", port: collidingPort, proto: udpForward}]
  184. set.mu.Unlock()
  185. if n != 3 {
  186. t.Fatalf("listeners after Reconcile with one pre-bound port = %d, want 3 (4 desired minus the 1 that couldn't bind)", n)
  187. }
  188. if tcpCollidingOpen {
  189. t.Error("TCP listener on the pre-bound port opened despite the real bind conflict")
  190. }
  191. if !udpCollidingOpen {
  192. t.Error("UDP listener on the colliding port's own number should still open (TCP and UDP binds are independent)")
  193. }
  194. dialLoopback(t, "tcp", okPort)
  195. set.Close()
  196. }
  197. // --- Real round trip: a genuine amneziawg-go client handshakes against a
  198. // real server Device, PortForwardSet opens a real host listener, and a real
  199. // external-side dial (this test's own process) round-trips bytes through
  200. // the actual encrypted tunnel to a service listening on the client's own
  201. // netstack -- proving the full path, not just the listener bookkeeping
  202. // above. Modeled closely on device_test.go's
  203. // TestNewDeviceHandshakeForwarderAndIdentity.
  204. func TestPortForwardRoundTripTCPAndUDP(t *testing.T) {
  205. serverPriv, serverPub, err := wireguard.GenerateWireguardKeypair()
  206. if err != nil {
  207. t.Fatalf("generate server keypair: %v", err)
  208. }
  209. clientPriv, clientPub, err := wireguard.GenerateWireguardKeypair()
  210. if err != nil {
  211. t.Fatalf("generate client keypair: %v", err)
  212. }
  213. const listenPort = 58920 // fixed loopback test port, matches this package's existing test convention
  214. const tcpPort = 58921
  215. const udpPort = 58922
  216. const clientAddr = "10.202.0.2"
  217. inst := amneziawg.Instance{
  218. Id: 5,
  219. InterfaceName: "awgtest5",
  220. ListenPort: listenPort,
  221. PrivateKey: serverPriv,
  222. PublicKey: serverPub,
  223. Address: []string{"10.202.0.1/24"},
  224. MTU: 1420,
  225. Obfuscation: amneziawg.Obfuscation31{
  226. Jc: 4, Jmin: 40, Jmax: 70,
  227. S1: 20, S2: 30, S3: 20, S4: 20,
  228. },
  229. Peers: []amneziawg.Peer{
  230. {
  231. Email: "client@test",
  232. PublicKey: clientPub,
  233. AllowedIPs: []string{clientAddr + "/32"},
  234. ForwardedPorts: fmt.Sprintf("%d,%d", tcpPort, udpPort),
  235. },
  236. },
  237. }
  238. dev, err := NewDevice(inst, DeviceOptions{})
  239. if err != nil {
  240. t.Fatalf("NewDevice: %v", err)
  241. }
  242. defer dev.Close()
  243. set := NewPortForwardSet(dev.Stack, inst.Id)
  244. set.Reconcile(inst)
  245. defer set.Close()
  246. // Real amneziawg-go client, same recipe as device_test.go.
  247. clientTun, clientNet, err := netstack.CreateNetTUN(
  248. []netip.Addr{netip.MustParseAddr(clientAddr)},
  249. []netip.Addr{netip.MustParseAddr("1.1.1.1")}, 1420)
  250. if err != nil {
  251. t.Fatalf("client CreateNetTUN: %v", err)
  252. }
  253. clientDev := device.NewDevice(clientTun, awgconn.NewDefaultBind(), device.NewLogger(device.LogLevelSilent, ""))
  254. defer clientDev.Close()
  255. // clientDev.Close() closes the tun's packet channel without waiting for
  256. // writers, so every goroutine writing into clientNet must be gone first.
  257. var clientSvc sync.WaitGroup
  258. defer clientSvc.Wait()
  259. clientPrivHex, err := wireguard.KeyToHex(clientPriv)
  260. if err != nil {
  261. t.Fatalf("client key to hex: %v", err)
  262. }
  263. serverPubHex, err := wireguard.KeyToHex(serverPub)
  264. if err != nil {
  265. t.Fatalf("server key to hex: %v", err)
  266. }
  267. clientConf := fmt.Sprintf(
  268. "private_key=%s\njc=4\njmin=40\njmax=70\ns1=20\ns2=30\ns3=20\ns4=20\npublic_key=%s\nendpoint=127.0.0.1:%d\nallowed_ip=0.0.0.0/0\n",
  269. clientPrivHex, serverPubHex, listenPort)
  270. if err := clientDev.IpcSet(clientConf); err != nil {
  271. t.Fatalf("client IpcSet: %v", err)
  272. }
  273. if err := clientDev.Up(); err != nil {
  274. t.Fatalf("client Up: %v", err)
  275. }
  276. // Prime the handshake before exercising the actual port forwards below.
  277. // The server only learns the client's real (roaming) endpoint from a
  278. // packet the client sends it -- buildUAPIConfig never configures an
  279. // endpoint= for a peer server-side (see device.go), and the server has
  280. // no route to initiate a handshake toward an endpoint it doesn't know --
  281. // so without this, relayTCPForward's own dial toward the client races a
  282. // handshake that can never even start server-side and fails outright.
  283. // A throwaway client dial toward nothing in particular is enough:
  284. // queuing any outbound packet triggers amneziawg-go's own automatic
  285. // handshake initiation regardless of whether the dial itself ever
  286. // succeeds (nothing server-side is listening for it), so this loop
  287. // deliberately ignores the dial's own outcome and just gives the
  288. // handshake a few real attempts to complete in the background.
  289. primeCtx, primeCancel := context.WithTimeout(context.Background(), 3*time.Second)
  290. defer primeCancel()
  291. for {
  292. if conn, dialErr := clientNet.DialContext(primeCtx, "tcp", "10.202.9.9:9999"); dialErr == nil {
  293. conn.Close()
  294. }
  295. select {
  296. case <-primeCtx.Done():
  297. goto primed
  298. case <-time.After(200 * time.Millisecond):
  299. }
  300. }
  301. primed:
  302. // A real service on the client's own netstack -- what a real forwarded
  303. // port is ultimately supposed to reach.
  304. tcpSvc, err := clientNet.ListenTCPAddrPort(netip.MustParseAddrPort(fmt.Sprintf("%s:%d", clientAddr, tcpPort)))
  305. if err != nil {
  306. t.Fatalf("client ListenTCP: %v", err)
  307. }
  308. defer tcpSvc.Close()
  309. clientSvc.Add(1)
  310. go func() {
  311. defer clientSvc.Done()
  312. for {
  313. c, err := tcpSvc.Accept()
  314. if err != nil {
  315. return
  316. }
  317. clientSvc.Add(1)
  318. go func() { defer clientSvc.Done(); io.Copy(c, c); c.Close() }()
  319. }
  320. }()
  321. udpSvc, err := clientNet.ListenUDPAddrPort(netip.MustParseAddrPort(fmt.Sprintf("%s:%d", clientAddr, udpPort)))
  322. if err != nil {
  323. t.Fatalf("client ListenUDP: %v", err)
  324. }
  325. defer udpSvc.Close()
  326. clientSvc.Add(1)
  327. go func() {
  328. defer clientSvc.Done()
  329. buf := make([]byte, 1500)
  330. for {
  331. n, addr, err := udpSvc.ReadFrom(buf)
  332. if err != nil {
  333. return
  334. }
  335. udpSvc.WriteTo(buf[:n], addr)
  336. }
  337. }()
  338. // Retry the TCP dial rather than guessing a fixed handshake delay --
  339. // the handshake happens lazily on first real traffic.
  340. const wantTCP = "port-forward tcp round trip"
  341. dialCtx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
  342. defer cancel()
  343. var tcpConn net.Conn
  344. var lastErr error
  345. for {
  346. tcpConn, lastErr = net.DialTimeout("tcp", fmt.Sprintf("127.0.0.1:%d", tcpPort), time.Second)
  347. if lastErr == nil {
  348. break
  349. }
  350. select {
  351. case <-dialCtx.Done():
  352. t.Fatalf("external TCP dial never succeeded: %v", lastErr)
  353. case <-time.After(150 * time.Millisecond):
  354. }
  355. }
  356. defer tcpConn.Close()
  357. if _, err := tcpConn.Write([]byte(wantTCP)); err != nil {
  358. t.Fatalf("write to forwarded TCP port: %v", err)
  359. }
  360. tcpConn.SetReadDeadline(time.Now().Add(5 * time.Second))
  361. gotTCP := make([]byte, len(wantTCP))
  362. if _, err := io.ReadFull(tcpConn, gotTCP); err != nil {
  363. t.Fatalf("read echo from forwarded TCP port: %v", err)
  364. }
  365. if string(gotTCP) != wantTCP {
  366. t.Errorf("TCP round trip = %q, want %q", gotTCP, wantTCP)
  367. }
  368. // UDP: the tunnel is already up (handshake completed above), so this
  369. // can dial straight away.
  370. const wantUDP = "port-forward udp round trip"
  371. udpConn, err := net.DialTimeout("udp", fmt.Sprintf("127.0.0.1:%d", udpPort), time.Second)
  372. if err != nil {
  373. t.Fatalf("external UDP dial: %v", err)
  374. }
  375. defer udpConn.Close()
  376. if _, err := udpConn.Write([]byte(wantUDP)); err != nil {
  377. t.Fatalf("write to forwarded UDP port: %v", err)
  378. }
  379. udpConn.SetReadDeadline(time.Now().Add(5 * time.Second))
  380. gotUDP := make([]byte, len(wantUDP))
  381. if _, err := io.ReadFull(udpConn, gotUDP); err != nil {
  382. t.Fatalf("read echo from forwarded UDP port: %v", err)
  383. }
  384. if string(gotUDP) != wantUDP {
  385. t.Errorf("UDP round trip = %q, want %q", gotUDP, wantUDP)
  386. }
  387. }