transport_test.go 9.5 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360
  1. package ike
  2. import (
  3. "context"
  4. "encoding/binary"
  5. "errors"
  6. "net"
  7. "testing"
  8. "time"
  9. )
  10. type deadlineError struct{}
  11. func (deadlineError) Error() string { return "deadline" }
  12. func (deadlineError) Timeout() bool { return true }
  13. func (deadlineError) Temporary() bool { return true }
  14. func TestRoundTripDatagramWaitsBeyondFirst500Milliseconds(t *testing.T) {
  15. available := make(chan struct{})
  16. go func() {
  17. time.Sleep(700 * time.Millisecond)
  18. close(available)
  19. }()
  20. writes := 0
  21. started := time.Now()
  22. response, err := roundTripDatagram(
  23. context.Background(),
  24. 2*time.Second,
  25. func([]byte) error {
  26. writes++
  27. return nil
  28. },
  29. func(buffer []byte, deadline time.Time) (int, error) {
  30. select {
  31. case <-available:
  32. copy(buffer, []byte("response"))
  33. return len("response"), nil
  34. case <-time.After(time.Until(deadline)):
  35. return 0, deadlineError{}
  36. }
  37. },
  38. []byte("request"),
  39. )
  40. if err != nil {
  41. t.Fatalf("roundTripDatagram() error = %v", err)
  42. }
  43. if string(response) != "response" || writes < 2 {
  44. t.Fatalf("response=%q writes=%d", response, writes)
  45. }
  46. if elapsed := time.Since(started); elapsed < 650*time.Millisecond {
  47. t.Fatalf("round trip returned too early after %v", elapsed)
  48. }
  49. }
  50. func TestRoundTripDatagramHonorsTotalTimeout(t *testing.T) {
  51. started := time.Now()
  52. _, err := roundTripDatagram(
  53. context.Background(),
  54. 120*time.Millisecond,
  55. func([]byte) error { return nil },
  56. func(_ []byte, deadline time.Time) (int, error) {
  57. time.Sleep(time.Until(deadline))
  58. return 0, deadlineError{}
  59. },
  60. []byte("request"),
  61. )
  62. if err == nil {
  63. t.Fatal("roundTripDatagram() accepted a missing response")
  64. }
  65. elapsed := time.Since(started)
  66. if elapsed < 100*time.Millisecond || elapsed > 400*time.Millisecond {
  67. t.Fatalf("total timeout elapsed = %v, want approximately 120ms", elapsed)
  68. }
  69. }
  70. func TestSOCKS5UDPDatagramRoundTrip(t *testing.T) {
  71. remote := &net.UDPAddr{IP: net.IPv4(203, 0, 113, 7), Port: 4500}
  72. encoded, err := marshalSOCKS5Datagram(remote, []byte{1, 2, 3, 4})
  73. if err != nil {
  74. t.Fatal(err)
  75. }
  76. payload, decoded, err := parseSOCKS5Datagram(encoded)
  77. if err != nil {
  78. t.Fatal(err)
  79. }
  80. if !decoded.IP.Equal(remote.IP) || decoded.Port != remote.Port || string(payload) != string([]byte{1, 2, 3, 4}) {
  81. t.Fatalf("decoded SOCKS datagram = %v %v %x", decoded, remote, payload)
  82. }
  83. fragmented := append([]byte(nil), encoded...)
  84. fragmented[2] = 1
  85. if _, _, err := parseSOCKS5Datagram(fragmented); err == nil {
  86. t.Fatal("fragmented SOCKS5 UDP datagram was accepted")
  87. }
  88. }
  89. func TestSOCKS5UDPAssociateDomainReplyIsResolved(t *testing.T) {
  90. client, server := net.Pipe()
  91. defer client.Close()
  92. defer server.Close()
  93. go func() {
  94. reply := []byte{5, 0, 0, 3, byte(len("localhost"))}
  95. reply = append(reply, "localhost"...)
  96. var port [2]byte
  97. binary.BigEndian.PutUint16(port[:], 7897)
  98. reply = append(reply, port[:]...)
  99. _, _ = server.Write(reply)
  100. }()
  101. address, err := readSOCKS5Reply(context.Background(), client, net.DefaultResolver)
  102. if err != nil {
  103. t.Fatalf("readSOCKS5Reply() error = %v", err)
  104. }
  105. if address.IP == nil || address.Port != 7897 {
  106. t.Fatalf("resolved relay = %v", address)
  107. }
  108. }
  109. func TestSOCKS5InitialExchangeFallsBackAcrossResolvedEPDGAddresses(t *testing.T) {
  110. relay, err := net.ListenUDP("udp", &net.UDPAddr{IP: net.IPv4(127, 0, 0, 1), Port: 0})
  111. if err != nil {
  112. t.Fatal(err)
  113. }
  114. defer relay.Close()
  115. connection, err := net.DialUDP("udp", nil, relay.LocalAddr().(*net.UDPAddr))
  116. if err != nil {
  117. t.Fatal(err)
  118. }
  119. first := &net.UDPAddr{IP: net.IPv4(192, 0, 2, 10), Port: 500}
  120. second := &net.UDPAddr{IP: net.IPv4(192, 0, 2, 20), Port: 500}
  121. transport := &socks5UDP{
  122. config: transportConfig{Timeout: 80 * time.Millisecond},
  123. udp: connection,
  124. remote: cloneUDPAddr(first),
  125. remotes: cloneUDPAddrs([]*net.UDPAddr{first, second}),
  126. }
  127. defer transport.Close()
  128. requestHeader := ikeHeader{
  129. InitiatorSPI: [8]byte{1, 2, 3, 4, 5, 6, 7, 8},
  130. Exchange: exchangeIKEInit,
  131. Flags: flagInitiator,
  132. }
  133. request := requestHeader.marshal([]byte("request"))
  134. response := ikeHeader{
  135. InitiatorSPI: requestHeader.InitiatorSPI,
  136. ResponderSPI: [8]byte{8, 7, 6, 5, 4, 3, 2, 1},
  137. Exchange: exchangeIKEInit,
  138. Flags: flagResponse,
  139. }.marshal([]byte("response"))
  140. serverDone := make(chan error, 1)
  141. go func() {
  142. buffer := make([]byte, 2048)
  143. for {
  144. n, peer, readErr := relay.ReadFromUDP(buffer)
  145. if readErr != nil {
  146. serverDone <- readErr
  147. return
  148. }
  149. _, destination, parseErr := parseSOCKS5Datagram(buffer[:n])
  150. if parseErr != nil {
  151. serverDone <- parseErr
  152. return
  153. }
  154. if !destination.IP.Equal(second.IP) {
  155. continue
  156. }
  157. wire, marshalErr := marshalSOCKS5Datagram(second, response)
  158. if marshalErr == nil {
  159. _, marshalErr = relay.WriteToUDP(wire, peer)
  160. }
  161. serverDone <- marshalErr
  162. return
  163. }
  164. }()
  165. got, err := transport.RoundTrip(context.Background(), request)
  166. if err != nil {
  167. t.Fatalf("RoundTrip() error = %v", err)
  168. }
  169. if string(got) != string(response) {
  170. t.Fatalf("RoundTrip() response = %x", got)
  171. }
  172. if !transport.RemoteAddr().IP.Equal(second.IP) {
  173. t.Fatalf("selected ePDG = %v, want %v", transport.RemoteAddr(), second)
  174. }
  175. if err := <-serverDone; err != nil {
  176. t.Fatalf("relay: %v", err)
  177. }
  178. }
  179. func TestSOCKS5RoundTripSkipsStaleAndESPDatagrams(t *testing.T) {
  180. relay, err := net.ListenUDP(
  181. "udp",
  182. &net.UDPAddr{IP: net.IPv4(127, 0, 0, 1), Port: 0},
  183. )
  184. if err != nil {
  185. t.Fatal(err)
  186. }
  187. defer relay.Close()
  188. connection, err := net.DialUDP(
  189. "udp",
  190. nil,
  191. relay.LocalAddr().(*net.UDPAddr),
  192. )
  193. if err != nil {
  194. t.Fatal(err)
  195. }
  196. remote := &net.UDPAddr{IP: net.IPv4(203, 0, 113, 7), Port: 4500}
  197. transport := &socks5UDP{
  198. config: transportConfig{Timeout: time.Second},
  199. udp: connection,
  200. remote: cloneUDPAddr(remote),
  201. floated: true,
  202. }
  203. defer transport.Close()
  204. requestHeader := ikeHeader{
  205. InitiatorSPI: [8]byte{1, 2, 3, 4, 5, 6, 7, 8},
  206. ResponderSPI: [8]byte{8, 7, 6, 5, 4, 3, 2, 1},
  207. Exchange: exchangeIKEAuth,
  208. Flags: flagInitiator,
  209. MessageID: 3,
  210. }
  211. request := requestHeader.marshal([]byte("request"))
  212. validResponse := ikeHeader{
  213. InitiatorSPI: requestHeader.InitiatorSPI,
  214. ResponderSPI: requestHeader.ResponderSPI,
  215. Exchange: requestHeader.Exchange,
  216. Flags: flagResponse,
  217. MessageID: requestHeader.MessageID,
  218. }.marshal([]byte("response"))
  219. serverDone := make(chan error, 1)
  220. go func() {
  221. buffer := make([]byte, 2048)
  222. _, peer, err := relay.ReadFromUDP(buffer)
  223. if err != nil {
  224. serverDone <- err
  225. return
  226. }
  227. stale, err := marshalSOCKS5Datagram(
  228. &net.UDPAddr{IP: remote.IP, Port: 500},
  229. append([]byte{0, 0, 0, 0}, []byte("stale")...),
  230. )
  231. if err != nil {
  232. serverDone <- err
  233. return
  234. }
  235. if _, err := relay.WriteToUDP(stale, peer); err != nil {
  236. serverDone <- err
  237. return
  238. }
  239. esp, err := marshalSOCKS5Datagram(
  240. remote,
  241. []byte{1, 2, 3, 4, 5, 6, 7, 8},
  242. )
  243. if err != nil {
  244. serverDone <- err
  245. return
  246. }
  247. if _, err := relay.WriteToUDP(esp, peer); err != nil {
  248. serverDone <- err
  249. return
  250. }
  251. staleIKE := ikeHeader{
  252. InitiatorSPI: requestHeader.InitiatorSPI,
  253. ResponderSPI: requestHeader.ResponderSPI,
  254. Exchange: requestHeader.Exchange,
  255. Flags: flagResponse,
  256. MessageID: requestHeader.MessageID - 1,
  257. }.marshal([]byte("stale IKE"))
  258. staleIKE, err = marshalSOCKS5Datagram(
  259. remote,
  260. append([]byte{0, 0, 0, 0}, staleIKE...),
  261. )
  262. if err != nil {
  263. serverDone <- err
  264. return
  265. }
  266. if _, err := relay.WriteToUDP(staleIKE, peer); err != nil {
  267. serverDone <- err
  268. return
  269. }
  270. valid, err := marshalSOCKS5Datagram(
  271. remote,
  272. append([]byte{0, 0, 0, 0}, validResponse...),
  273. )
  274. if err == nil {
  275. _, err = relay.WriteToUDP(valid, peer)
  276. }
  277. serverDone <- err
  278. }()
  279. response, err := transport.RoundTrip(
  280. context.Background(),
  281. request,
  282. )
  283. if err != nil {
  284. t.Fatalf("RoundTrip() error = %v", err)
  285. }
  286. if string(response) != string(validResponse) {
  287. t.Fatalf("RoundTrip() response = %x", response)
  288. }
  289. if err := <-serverDone; err != nil {
  290. t.Fatalf("relay: %v", err)
  291. }
  292. }
  293. func TestSessionReadDoesNotBlockIndependentWrite(t *testing.T) {
  294. server, err := net.ListenUDP("udp", &net.UDPAddr{IP: net.IPv4(127, 0, 0, 1), Port: 0})
  295. if err != nil {
  296. t.Fatal(err)
  297. }
  298. defer server.Close()
  299. connection, err := net.DialUDP("udp", nil, server.LocalAddr().(*net.UDPAddr))
  300. if err != nil {
  301. t.Fatal(err)
  302. }
  303. transport := &directUDP{
  304. config: transportConfig{Timeout: time.Second},
  305. conn: connection,
  306. remote: cloneUDPAddr(server.LocalAddr().(*net.UDPAddr)),
  307. floated: true,
  308. }
  309. defer transport.Close()
  310. readDone := make(chan error, 1)
  311. go func() {
  312. buffer := make([]byte, 64)
  313. _, _, err := transport.ReceiveSessionPacket(context.Background(), buffer)
  314. readDone <- err
  315. }()
  316. time.Sleep(30 * time.Millisecond)
  317. started := time.Now()
  318. if err := transport.SendSessionPacket(context.Background(), []byte{1, 2, 3, 4, 5, 6, 7, 8}, false); err != nil {
  319. t.Fatalf("SendSessionPacket() error = %v", err)
  320. }
  321. if elapsed := time.Since(started); elapsed > 200*time.Millisecond {
  322. t.Fatalf("session write blocked behind reader for %v", elapsed)
  323. }
  324. buffer := make([]byte, 64)
  325. _ = server.SetReadDeadline(time.Now().Add(time.Second))
  326. n, _, err := server.ReadFromUDP(buffer)
  327. if err != nil {
  328. t.Fatal(err)
  329. }
  330. if n != 8 {
  331. t.Fatalf("server received %d bytes, want 8", n)
  332. }
  333. _ = transport.Close()
  334. select {
  335. case err := <-readDone:
  336. if err == nil || (!errors.Is(err, net.ErrClosed) && !isNetworkClose(err)) {
  337. t.Fatalf("reader close error = %v", err)
  338. }
  339. case <-time.After(2 * time.Second):
  340. t.Fatal("session reader did not wake after Close")
  341. }
  342. }
  343. func isNetworkClose(err error) bool {
  344. var networkError net.Error
  345. return errors.As(err, &networkError)
  346. }