relay_test.go 5.7 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216
  1. package ike
  2. import (
  3. "bytes"
  4. "context"
  5. "net"
  6. "sync"
  7. "sync/atomic"
  8. "testing"
  9. "time"
  10. )
  11. type fakeSessionPacket struct {
  12. data []byte
  13. ike bool
  14. err error
  15. }
  16. type fakeSentPacket struct {
  17. data []byte
  18. ike bool
  19. }
  20. type fakeSessionTransport struct {
  21. incoming chan fakeSessionPacket
  22. sent chan fakeSentPacket
  23. closed chan struct{}
  24. once sync.Once
  25. readers atomic.Int32
  26. maxReads atomic.Int32
  27. }
  28. func newFakeSessionTransport() *fakeSessionTransport {
  29. return &fakeSessionTransport{
  30. incoming: make(chan fakeSessionPacket, 16),
  31. sent: make(chan fakeSentPacket, 16),
  32. closed: make(chan struct{}),
  33. }
  34. }
  35. func (transport *fakeSessionTransport) LocalAddr() *net.UDPAddr {
  36. return &net.UDPAddr{IP: net.IPv4(127, 0, 0, 1), Port: 4500}
  37. }
  38. func (transport *fakeSessionTransport) RemoteAddr() *net.UDPAddr {
  39. return &net.UDPAddr{IP: net.IPv4(127, 0, 0, 2), Port: 4500}
  40. }
  41. func (*fakeSessionTransport) Float(context.Context) error { return nil }
  42. func (*fakeSessionTransport) RoundTrip(context.Context, []byte) ([]byte, error) {
  43. return nil, context.DeadlineExceeded
  44. }
  45. func (transport *fakeSessionTransport) SendESP(ctx context.Context, packet []byte) error {
  46. return transport.SendSessionPacket(ctx, packet, false)
  47. }
  48. func (transport *fakeSessionTransport) ReceiveESP(ctx context.Context, buffer []byte) (int, error) {
  49. n, _, err := transport.ReceiveSessionPacket(ctx, buffer)
  50. return n, err
  51. }
  52. func (transport *fakeSessionTransport) SendSessionPacket(
  53. ctx context.Context,
  54. packet []byte,
  55. ike bool,
  56. ) error {
  57. select {
  58. case transport.sent <- fakeSentPacket{data: append([]byte(nil), packet...), ike: ike}:
  59. return nil
  60. case <-ctx.Done():
  61. return ctx.Err()
  62. case <-transport.closed:
  63. return net.ErrClosed
  64. }
  65. }
  66. func (transport *fakeSessionTransport) ReceiveSessionPacket(
  67. ctx context.Context,
  68. buffer []byte,
  69. ) (int, bool, error) {
  70. active := transport.readers.Add(1)
  71. for {
  72. current := transport.maxReads.Load()
  73. if active <= current || transport.maxReads.CompareAndSwap(current, active) {
  74. break
  75. }
  76. }
  77. defer transport.readers.Add(-1)
  78. select {
  79. case packet := <-transport.incoming:
  80. if packet.err != nil {
  81. return 0, false, packet.err
  82. }
  83. copy(buffer, packet.data)
  84. return len(packet.data), packet.ike, nil
  85. case <-time.After(5 * time.Millisecond):
  86. return 0, false, deadlineError{}
  87. case <-ctx.Done():
  88. return 0, false, ctx.Err()
  89. case <-transport.closed:
  90. return 0, false, net.ErrClosed
  91. }
  92. }
  93. func (transport *fakeSessionTransport) Close() error {
  94. transport.once.Do(func() { close(transport.closed) })
  95. return nil
  96. }
  97. func TestSessionRelayDemuxesESPAndAnswersEncryptedDPD(t *testing.T) {
  98. transport := newFakeSessionTransport()
  99. suite := legacyTestSuite()
  100. keys := ikeKeys{
  101. SKai: bytes.Repeat([]byte{0x11}, 20),
  102. SKar: bytes.Repeat([]byte{0x12}, 20),
  103. SKei: bytes.Repeat([]byte{0x13}, 16),
  104. SKer: bytes.Repeat([]byte{0x14}, 16),
  105. }
  106. spii := [8]byte{1}
  107. spir := [8]byte{2}
  108. relay := newSessionRelay(transport, suite, keys, spii, spir, true, time.Hour)
  109. defer relay.Close()
  110. esp := []byte{0, 0, 0, 9, 0, 0, 0, 1, 0xaa}
  111. transport.incoming <- fakeSessionPacket{data: esp, ike: false}
  112. buffer := make([]byte, 64)
  113. n, err := relay.ReceiveESP(context.Background(), buffer)
  114. if err != nil {
  115. t.Fatalf("ReceiveESP() error = %v", err)
  116. }
  117. if !bytes.Equal(buffer[:n], esp) {
  118. t.Fatalf("demuxed ESP = %x, want %x", buffer[:n], esp)
  119. }
  120. dpd, err := encryptPayloads(ikeHeader{
  121. InitiatorSPI: spii,
  122. ResponderSPI: spir,
  123. Exchange: exchangeInformational,
  124. MessageID: 8,
  125. }, nil, suite, keys.SKer, keys.SKar, bytes.NewReader(bytes.Repeat([]byte{0x44}, 64)))
  126. if err != nil {
  127. t.Fatal(err)
  128. }
  129. transport.incoming <- fakeSessionPacket{data: dpd, ike: true}
  130. select {
  131. case response := <-transport.sent:
  132. if !response.ike {
  133. t.Fatal("DPD response was sent as ESP")
  134. }
  135. header, payloads, err := decryptPayloads(response.data, suite, keys.SKei, keys.SKai)
  136. if err != nil {
  137. t.Fatalf("decrypt DPD response: %v", err)
  138. }
  139. if header.Exchange != exchangeInformational || header.MessageID != 8 ||
  140. header.Flags != flagInitiator|flagResponse || len(payloads) != 0 {
  141. t.Fatalf("DPD response header/payloads = %#v %#v", header, payloads)
  142. }
  143. case <-time.After(time.Second):
  144. t.Fatal("relay did not answer DPD")
  145. }
  146. if maximum := transport.maxReads.Load(); maximum != 1 {
  147. t.Fatalf("concurrent socket readers = %d, want exactly one", maximum)
  148. }
  149. }
  150. func TestSessionRelaySendsNATKeepalive(t *testing.T) {
  151. transport := newFakeSessionTransport()
  152. relay := newSessionRelay(
  153. transport,
  154. legacyTestSuite(),
  155. ikeKeys{},
  156. [8]byte{1},
  157. [8]byte{2},
  158. true,
  159. 10*time.Millisecond,
  160. )
  161. defer relay.Close()
  162. select {
  163. case packet := <-transport.sent:
  164. if packet.ike || !bytes.Equal(packet.data, []byte{0xff}) {
  165. t.Fatalf("keepalive = ike:%v data:%x", packet.ike, packet.data)
  166. }
  167. case <-time.After(500 * time.Millisecond):
  168. t.Fatal("relay did not send a NAT-T keepalive")
  169. }
  170. }
  171. func TestSessionRelayDropsDelayedIKEPacketFromPreviousSA(t *testing.T) {
  172. transport := newFakeSessionTransport()
  173. spii := [8]byte{1}
  174. spir := [8]byte{2}
  175. relay := newSessionRelay(
  176. transport,
  177. legacyTestSuite(),
  178. ikeKeys{},
  179. spii,
  180. spir,
  181. true,
  182. time.Hour,
  183. )
  184. defer relay.Close()
  185. transport.incoming <- fakeSessionPacket{
  186. ike: true,
  187. data: ikeHeader{
  188. InitiatorSPI: [8]byte{9},
  189. ResponderSPI: [8]byte{8},
  190. Exchange: exchangeInformational,
  191. }.marshal(nil),
  192. }
  193. wantedESP := []byte{0, 0, 0, 9, 0, 0, 0, 1, 0xaa}
  194. transport.incoming <- fakeSessionPacket{data: wantedESP}
  195. buffer := make([]byte, 64)
  196. count, err := relay.ReceiveESP(context.Background(), buffer)
  197. if err != nil {
  198. t.Fatalf("ReceiveESP() after stale IKE packet = %v", err)
  199. }
  200. if !bytes.Equal(buffer[:count], wantedESP) {
  201. t.Fatalf("ESP after stale IKE packet = %x, want %x", buffer[:count], wantedESP)
  202. }
  203. }