1
0

manager_shutdown_test.go 3.9 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140
  1. package tuic
  2. import (
  3. "bytes"
  4. "context"
  5. "crypto/tls"
  6. "fmt"
  7. "io"
  8. "net"
  9. "testing"
  10. "time"
  11. "github.com/google/uuid"
  12. clientquic "github.com/quic-go/quic-go"
  13. )
  14. func TestAudit3ManagerStopMustInterruptIdleTCPRelay(t *testing.T) {
  15. var listener net.Listener
  16. var err error
  17. var inboundID int
  18. for p := 64051; p < 64080; p++ {
  19. listener, err = net.Listen("tcp", fmt.Sprintf("127.0.0.1:%d", p))
  20. if err == nil {
  21. inboundID = 500000 + (p - 64000)
  22. break
  23. }
  24. }
  25. if listener == nil {
  26. t.Fatal(err)
  27. }
  28. defer listener.Close()
  29. release := make(chan struct{})
  30. defer close(release)
  31. peerFIN := make(chan struct{})
  32. ready := make(chan struct{})
  33. go func() {
  34. c, err := listener.Accept()
  35. if err != nil {
  36. return
  37. }
  38. defer c.Close()
  39. var greeting [4]byte
  40. if _, err := io.ReadFull(c, greeting[:]); err != nil {
  41. return
  42. }
  43. c.Write([]byte{5, 0})
  44. var req [10]byte
  45. if _, err := io.ReadFull(c, req[:]); err != nil {
  46. return
  47. }
  48. c.Write([]byte{5, 0, 0, 1, 127, 0, 0, 1, 0, 0})
  49. var payload [1]byte
  50. if _, err := io.ReadFull(c, payload[:]); err != nil {
  51. return
  52. }
  53. c.Write(payload[:])
  54. close(ready)
  55. io.Copy(io.Discard, c)
  56. close(peerFIN)
  57. <-release
  58. }()
  59. cert, key := generateTestCert(t)
  60. clientID := uuid.New()
  61. password := "audit3-password"
  62. m := &Manager{servers: make(map[int]*managed), lastStartErr: make(map[int]string), pendingTraffic: make(map[string]ClientTrafficDelta)}
  63. if err := m.Ensure(Instance{Id: inboundID, Tag: "audit3-manager-shutdown", Listen: "127.0.0.1", Port: 0, Certificate: string(cert), PrivateKey: string(key), ALPN: []string{"h3"}, AuthenticationTimeout: 2, MaxIdleTime: 30, Clients: []TuicClientSettings{{UUID: clientID.String(), Password: password, Email: "close-idle@audit3"}}}); err != nil {
  64. t.Fatal(err)
  65. }
  66. t.Cleanup(m.StopAll)
  67. server := m.servers[inboundID].server
  68. dialCtx, dialCancel := context.WithTimeout(context.Background(), 5*time.Second)
  69. defer dialCancel()
  70. conn, err := clientquic.DialAddr(dialCtx, server.packetConn.LocalAddr().String(), &tls.Config{InsecureSkipVerify: true, NextProtos: []string{"h3"}}, &clientquic.Config{EnableDatagrams: true})
  71. if err != nil {
  72. t.Fatal(err)
  73. }
  74. t.Cleanup(func() { conn.CloseWithError(0, "") })
  75. tlsState := conn.ConnectionState().TLS
  76. token, err := tlsState.ExportKeyingMaterial(string(clientID[:]), []byte(password), 32)
  77. if err != nil {
  78. t.Fatal(err)
  79. }
  80. auth, err := conn.OpenUniStreamSync(dialCtx)
  81. if err != nil {
  82. t.Fatal(err)
  83. }
  84. payload := append([]byte{ProtocolVersion, CmdAuthenticate}, clientID[:]...)
  85. payload = append(payload, token...)
  86. if _, err := auth.Write(payload); err != nil {
  87. t.Fatal(err)
  88. }
  89. auth.Close()
  90. ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
  91. defer cancel()
  92. stream, err := conn.OpenStreamSync(ctx)
  93. if err != nil {
  94. t.Fatal(err)
  95. }
  96. var cmd bytes.Buffer
  97. cmd.Write([]byte{ProtocolVersion, CmdConnect})
  98. WriteAddress(&cmd, &Address{Type: AddrTypeIPv4, IP: net.ParseIP("1.1.1.1"), Port: 443})
  99. cmd.WriteByte('x')
  100. if _, err := stream.Write(cmd.Bytes()); err != nil {
  101. t.Fatal(err)
  102. }
  103. var echo [1]byte
  104. if _, err := io.ReadFull(stream, echo[:]); err != nil {
  105. t.Fatal(err)
  106. }
  107. <-ready
  108. closed := make(chan error, 1)
  109. started := time.Now()
  110. go func() { m.StopAll(); closed <- nil }()
  111. queried := make(chan bool, 1)
  112. go func() { queried <- m.HasRunning() }()
  113. select {
  114. case <-closed:
  115. case <-time.After(time.Second):
  116. t.Fatalf("StopAll blocked on idle TCP peer after %s", time.Since(started))
  117. }
  118. select {
  119. case <-queried:
  120. case <-time.After(time.Second):
  121. t.Fatal("manager query blocked after StopAll")
  122. }
  123. select {
  124. case <-peerFIN:
  125. case <-time.After(time.Second):
  126. t.Fatal("upstream connection remained open")
  127. }
  128. _, deltas := m.CollectAllTraffic()
  129. if len(deltas) != 1 || deltas[0].Up != 1 || deltas[0].Down != 1 {
  130. t.Fatalf("final counters: %+v", deltas)
  131. }
  132. _, again := m.CollectAllTraffic()
  133. if len(again) != 0 {
  134. t.Fatalf("repeated final counters: %+v", again)
  135. }
  136. }