relay_test.go 8.2 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257
  1. package amneziawgnet
  2. import (
  3. "io"
  4. "net"
  5. "net/netip"
  6. "testing"
  7. "time"
  8. )
  9. // newDeadUDPSession builds a socks5UDPSession over real but already-closed
  10. // sockets, so pump's receive fails immediately and its teardown runs at once.
  11. func newDeadUDPSession(t *testing.T) *socks5UDPSession {
  12. t.Helper()
  13. peer, err := net.ListenUDP("udp", &net.UDPAddr{IP: net.IPv4(127, 0, 0, 1)})
  14. if err != nil {
  15. t.Fatalf("listen udp: %v", err)
  16. }
  17. t.Cleanup(func() { _ = peer.Close() })
  18. udpConn, err := net.DialUDP("udp", nil, peer.LocalAddr().(*net.UDPAddr))
  19. if err != nil {
  20. t.Fatalf("dial udp: %v", err)
  21. }
  22. ctrlLn, err := net.Listen("tcp", "127.0.0.1:0")
  23. if err != nil {
  24. t.Fatalf("listen tcp: %v", err)
  25. }
  26. t.Cleanup(func() { _ = ctrlLn.Close() })
  27. ctrl, err := net.Dial("tcp", ctrlLn.Addr().String())
  28. if err != nil {
  29. t.Fatalf("dial tcp: %v", err)
  30. }
  31. _ = udpConn.Close()
  32. return &socks5UDPSession{ctrl: ctrl, udpConn: udpConn}
  33. }
  34. // TestUDPRelayPumpOnlyRetiresItsOwnSession pins the flow that survives a
  35. // duplicate associate: a losing pump must not evict the published session.
  36. func TestUDPRelayPumpOnlyRetiresItsOwnSession(t *testing.T) {
  37. relay := NewUDPRelay(SocksRelay{Addr: "127.0.0.1:1", Password: "x"}, nil)
  38. src := netip.MustParseAddrPort("10.8.1.5:51820")
  39. live := newDeadUDPSession(t)
  40. superseded := newDeadUDPSession(t)
  41. relay.sessions[src] = live
  42. // Returns as soon as receive fails on the closed socket, so no wait is needed.
  43. relay.pump(src, superseded)
  44. got, ok := relay.sessions[src]
  45. if !ok {
  46. t.Fatal("live session was evicted: a retiring pump deleted src's entry regardless of which session held it")
  47. }
  48. if got != live {
  49. t.Fatalf("sessions[%v] = %p, want the live session %p", src, got, live)
  50. }
  51. }
  52. // TestUDPRelayCloseDropsEverySession keeps Close's contract explicit now that
  53. // pump's teardown is conditional on still owning the key.
  54. func TestUDPRelayCloseDropsEverySession(t *testing.T) {
  55. relay := NewUDPRelay(SocksRelay{Addr: "127.0.0.1:1", Password: "x"}, nil)
  56. for _, s := range []string{"10.8.1.5:51820", "10.8.1.6:2000"} {
  57. relay.sessions[netip.MustParseAddrPort(s)] = newDeadUDPSession(t)
  58. }
  59. relay.Close()
  60. if n := len(relay.sessions); n != 0 {
  61. t.Fatalf("Close left %d sessions behind, want 0", n)
  62. }
  63. }
  64. // tcpPair returns a connected pair of real loopback TCP conns; net.Pipe would
  65. // not do, since these tests turn on CloseWrite, which it does not implement.
  66. func tcpPair(t *testing.T) (client, server net.Conn) {
  67. t.Helper()
  68. ln, err := net.Listen("tcp", "127.0.0.1:0")
  69. if err != nil {
  70. t.Fatalf("listen: %v", err)
  71. }
  72. defer ln.Close()
  73. type accepted struct {
  74. conn net.Conn
  75. err error
  76. }
  77. ch := make(chan accepted, 1)
  78. go func() {
  79. c, err := ln.Accept()
  80. ch <- accepted{c, err}
  81. }()
  82. client, err = net.Dial("tcp", ln.Addr().String())
  83. if err != nil {
  84. t.Fatalf("dial: %v", err)
  85. }
  86. got := <-ch
  87. if got.err != nil {
  88. t.Fatalf("accept: %v", got.err)
  89. }
  90. t.Cleanup(func() { _ = client.Close(); _ = got.conn.Close() })
  91. return client, got.conn
  92. }
  93. // TestPipeBothWaysDeliversReplyAfterHalfClose is the half-close regression: a
  94. // client that shuts down its write side must still receive the full response.
  95. func TestPipeBothWaysDeliversReplyAfterHalfClose(t *testing.T) {
  96. const request = "GET / HTTP/1.0\r\n\r\n"
  97. const response = "the reply that arrives only after the request is complete"
  98. client, a := tcpPair(t)
  99. b, server := tcpPair(t)
  100. go pipeBothWays(a, b)
  101. if _, err := client.Write([]byte(request)); err != nil {
  102. t.Fatalf("client write: %v", err)
  103. }
  104. // The half-close the old relay treated as "tear the whole pair down".
  105. if err := client.(*net.TCPConn).CloseWrite(); err != nil {
  106. t.Fatalf("client CloseWrite: %v", err)
  107. }
  108. _ = server.SetReadDeadline(time.Now().Add(10 * time.Second))
  109. gotReq, err := io.ReadAll(server)
  110. if err != nil {
  111. t.Fatalf("server read: %v", err)
  112. }
  113. if string(gotReq) != request {
  114. t.Fatalf("server got request %q, want %q", gotReq, request)
  115. }
  116. if _, err := server.Write([]byte(response)); err != nil {
  117. t.Fatalf("server write: %v", err)
  118. }
  119. if err := server.(*net.TCPConn).CloseWrite(); err != nil {
  120. t.Fatalf("server CloseWrite: %v", err)
  121. }
  122. _ = client.SetReadDeadline(time.Now().Add(10 * time.Second))
  123. gotResp, err := io.ReadAll(client)
  124. if err != nil {
  125. t.Fatalf("client read: %v", err)
  126. }
  127. if string(gotResp) != response {
  128. t.Fatalf("client got response %q, want %q: the reply was cut off by the half-close", gotResp, response)
  129. }
  130. }
  131. // TestPipeBothWaysClosesWhenBothSidesFinish keeps the teardown contract: both
  132. // directions ending must return, not hang on the idle bound.
  133. func TestPipeBothWaysClosesWhenBothSidesFinish(t *testing.T) {
  134. client, a := tcpPair(t)
  135. b, server := tcpPair(t)
  136. done := make(chan struct{})
  137. go func() { defer close(done); pipeBothWays(a, b) }()
  138. _ = client.(*net.TCPConn).CloseWrite()
  139. _, _ = io.ReadAll(server)
  140. _ = server.(*net.TCPConn).CloseWrite()
  141. _, _ = io.ReadAll(client)
  142. select {
  143. case <-done:
  144. case <-time.After(10 * time.Second):
  145. t.Fatal("pipeBothWays did not return after both directions ended")
  146. }
  147. }
  148. // liveUDPSession returns a session whose udpConn is connected to the returned
  149. // peer, so a test can hand receive() one exact reply datagram.
  150. func liveUDPSession(t *testing.T) (*socks5UDPSession, *net.UDPConn) {
  151. t.Helper()
  152. peer, err := net.ListenUDP("udp", &net.UDPAddr{IP: net.IPv4(127, 0, 0, 1)})
  153. if err != nil {
  154. t.Fatalf("listen udp: %v", err)
  155. }
  156. t.Cleanup(func() { _ = peer.Close() })
  157. udpConn, err := net.DialUDP("udp", nil, peer.LocalAddr().(*net.UDPAddr))
  158. if err != nil {
  159. t.Fatalf("dial udp: %v", err)
  160. }
  161. t.Cleanup(func() { _ = udpConn.Close() })
  162. return &socks5UDPSession{udpConn: udpConn}, peer
  163. }
  164. // sendReply delivers one raw datagram to sess's socket.
  165. func sendReply(t *testing.T, sess *socks5UDPSession, peer *net.UDPConn, datagram []byte) {
  166. t.Helper()
  167. if _, err := peer.WriteToUDP(datagram, sess.udpConn.LocalAddr().(*net.UDPAddr)); err != nil {
  168. t.Fatalf("write reply: %v", err)
  169. }
  170. _ = sess.udpConn.SetReadDeadline(time.Now().Add(5 * time.Second))
  171. }
  172. // TestSocks5ReceiveDecodesReplyAddressTypes covers all three ATYP forms; the
  173. // domain form used to misread its own length byte and never skip the name.
  174. func TestSocks5ReceiveDecodesReplyAddressTypes(t *testing.T) {
  175. tests := []struct {
  176. name string
  177. addrPart []byte
  178. wantAddr string
  179. }{
  180. {"IPv4", []byte{0x01, 10, 0, 0, 7}, "10.0.0.7"},
  181. {"IPv6", append([]byte{0x04}, netip.MustParseAddr("2001:db8::5").AsSlice()...), "2001:db8::5"},
  182. {"domain holding a literal", append([]byte{0x03, 8}, []byte("10.0.0.9")...), "10.0.0.9"},
  183. }
  184. for _, tt := range tests {
  185. t.Run(tt.name, func(t *testing.T) {
  186. sess, peer := liveUDPSession(t)
  187. payload := []byte("the-actual-datagram-payload")
  188. datagram := append([]byte{0x00, 0x00, 0x00}, tt.addrPart...)
  189. datagram = append(datagram, 0x1f, 0x90) // port 8080
  190. datagram = append(datagram, payload...)
  191. sendReply(t, sess, peer, datagram)
  192. buf := make([]byte, 4096)
  193. from, got, err := sess.receive(buf)
  194. if err != nil {
  195. t.Fatalf("receive: %v", err)
  196. }
  197. want := netip.AddrPortFrom(netip.MustParseAddr(tt.wantAddr), 8080)
  198. if from != want {
  199. t.Errorf("source = %v, want %v", from, want)
  200. }
  201. if string(got) != string(payload) {
  202. t.Errorf("payload = %q, want %q", got, payload)
  203. }
  204. })
  205. }
  206. }
  207. // TestSocks5ReceiveRejectsTruncatedReplies pins that a short datagram is an
  208. // error, not a slice-bounds panic in the relay's own pump goroutine.
  209. func TestSocks5ReceiveRejectsTruncatedReplies(t *testing.T) {
  210. tests := []struct {
  211. name string
  212. datagram []byte
  213. }{
  214. {"header only, IPv4 announced", []byte{0x00, 0x00, 0x00, 0x01}},
  215. {"IPv4 address cut short", []byte{0x00, 0x00, 0x00, 0x01, 10, 0}},
  216. {"IPv6 address cut short", []byte{0x00, 0x00, 0x00, 0x04, 0x20, 0x01}},
  217. {"domain length past the end", []byte{0x00, 0x00, 0x00, 0x03, 40, 'a', 'b'}},
  218. {"address complete but port missing", []byte{0x00, 0x00, 0x00, 0x01, 10, 0, 0, 7}},
  219. {"unsupported address type", []byte{0x00, 0x00, 0x00, 0x09, 1, 2, 3, 4, 0, 80}},
  220. }
  221. for _, tt := range tests {
  222. t.Run(tt.name, func(t *testing.T) {
  223. sess, peer := liveUDPSession(t)
  224. sendReply(t, sess, peer, tt.datagram)
  225. buf := make([]byte, 4096)
  226. if _, _, err := sess.receive(buf); err == nil {
  227. t.Fatal("receive accepted a malformed datagram instead of returning an error")
  228. }
  229. })
  230. }
  231. }