manager_live_traffic_test.go 7.1 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244
  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/apernet/quic-go"
  12. "github.com/google/uuid"
  13. )
  14. type reauditCCSnapshot struct {
  15. conn *quic.Conn
  16. chosen string
  17. actual string
  18. sender uintptr
  19. }
  20. func reauditActualSender(conn *quic.Conn) (string, uintptr) {
  21. cc, unlock := lockedCongestion(conn)
  22. defer unlock()
  23. ptr := cc.Pointer()
  24. if cc.Type().String() == "*ackhandler.ccAdapterEx" || cc.Type().String() == "*ackhandler.ccAdapter" {
  25. sender := cc.Elem().FieldByName("CC").Elem()
  26. return sender.Type().String(), ptr
  27. }
  28. return fmt.Sprintf("%s reno=%t", cc.Type(), cc.Elem().FieldByName("reno").Bool()), ptr
  29. }
  30. func reauditWantedSender(controller string) string {
  31. if controller == "bbr" {
  32. return "*bbr.bbrSender"
  33. }
  34. return "*congestion.cubicSender reno=true"
  35. }
  36. func TestAudit3ManagerEnsureActualSendersWithPersistentTraffic(t *testing.T) {
  37. cert, key := generateTestCert(t)
  38. _, cleanup := audit3StartSocksForManager(t, "[email protected]", SocksPassword(), 99115)
  39. defer cleanup()
  40. userID := uuid.MustParse("a0000000-0000-0000-0000-000000000015")
  41. inst := Instance{Id: 99115, Tag: "reaudit-cc", Listen: "127.0.0.1", Certificate: string(cert), PrivateKey: string(key), CongestionControl: "new_reno", AuthenticationTimeout: 3, MaxIdleTime: 30, Clients: []TuicClientSettings{{UUID: userID.String(), Password: "secret-reaudit", Email: "[email protected]"}}}
  42. manager := &Manager{servers: map[int]*managed{}, lastStartErr: map[int]string{}}
  43. if err := manager.Ensure(inst); err != nil {
  44. t.Fatal(err)
  45. }
  46. defer manager.StopAll()
  47. server := manager.servers[inst.Id].server
  48. listener := server.quicListener
  49. address := server.packetConn.LocalAddr().String()
  50. ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second)
  51. defer cancel()
  52. type peer struct {
  53. client *quic.Conn
  54. tcp *quic.Stream
  55. snapshot reauditCCSnapshot
  56. packetID uint16
  57. }
  58. var peers []*peer
  59. tcpEcho := func(p *peer, message []byte) {
  60. t.Helper()
  61. _ = p.tcp.SetDeadline(time.Now().Add(2 * time.Second))
  62. if _, err := p.tcp.Write(message); err != nil {
  63. t.Fatal(err)
  64. }
  65. reply := make([]byte, len(message))
  66. if _, err := io.ReadFull(p.tcp, reply); err != nil {
  67. t.Fatal(err)
  68. }
  69. if !bytes.Equal(reply, message) {
  70. t.Fatalf("TCP echo mismatch: %q", reply)
  71. }
  72. }
  73. udpEcho := func(p *peer, streamMode bool, message []byte) {
  74. t.Helper()
  75. p.packetID++
  76. assoc := uint16(100)
  77. if streamMode {
  78. assoc = 200
  79. }
  80. var frame bytes.Buffer
  81. if err := WritePacket(&frame, assoc, p.packetID, 1, 0, &Address{Type: AddrTypeIPv4, IP: net.ParseIP("8.8.8.8"), Port: 53}, message); err != nil {
  82. t.Fatal(err)
  83. }
  84. var reader io.Reader
  85. if streamMode {
  86. stream, err := p.client.OpenUniStreamSync(ctx)
  87. if err != nil {
  88. t.Fatal(err)
  89. }
  90. if _, err := stream.Write(frame.Bytes()); err != nil {
  91. t.Fatal(err)
  92. }
  93. closeUniStream(t, stream)
  94. response, err := p.client.AcceptUniStream(ctx)
  95. if err != nil {
  96. t.Fatal(err)
  97. }
  98. reader = response
  99. } else {
  100. if err := p.client.SendDatagram(frame.Bytes()); err != nil {
  101. t.Fatal(err)
  102. }
  103. response, err := p.client.ReceiveDatagram(ctx)
  104. if err != nil {
  105. t.Fatal(err)
  106. }
  107. reader = bytes.NewReader(response)
  108. }
  109. _, command, err := ReadCommand(reader)
  110. if err != nil || command != CmdPacket {
  111. t.Fatalf("UDP response command=%d error=%v", command, err)
  112. }
  113. hdr, err := ReadPacketHeader(reader)
  114. if err != nil {
  115. t.Fatal(err)
  116. }
  117. payload, err := readPacketPayload(reader, hdr)
  118. if err != nil {
  119. t.Fatal(err)
  120. }
  121. if hdr.AssocID != assoc || !bytes.Equal(payload, message) {
  122. t.Fatalf("UDP echo mismatch association=%d payload=%q", hdr.AssocID, payload)
  123. }
  124. }
  125. for step, controller := range []string{"new_reno", "reno", "bbr", "BBR", "cubic", "CuBiC", "", "invalid"} {
  126. inst.CongestionControl = controller
  127. if err := manager.Ensure(inst); err != nil {
  128. t.Fatal(err)
  129. }
  130. normalized, _ := normalizeCongestionControl(controller)
  131. served := normalized
  132. if served == "cubic" {
  133. served = "new_reno"
  134. }
  135. if server.quicListener != listener || server.packetConn.LocalAddr().String() != address {
  136. t.Fatal("listener changed")
  137. }
  138. client, err := quic.DialAddr(ctx, address, &tls.Config{InsecureSkipVerify: true, NextProtos: []string{"h3"}}, &quic.Config{EnableDatagrams: true, MaxIdleTimeout: 30 * time.Second})
  139. if err != nil {
  140. t.Fatal(err)
  141. }
  142. defer client.CloseWithError(0, "")
  143. tlsState := client.ConnectionState().TLS
  144. token, err := tlsState.ExportKeyingMaterial(string(userID[:]), []byte("secret-reaudit"), 32)
  145. if err != nil {
  146. t.Fatal(err)
  147. }
  148. auth, err := client.OpenUniStreamSync(ctx)
  149. if err != nil {
  150. t.Fatal(err)
  151. }
  152. authBytes := make([]byte, 50)
  153. authBytes[0], authBytes[1] = ProtocolVersion, CmdAuthenticate
  154. copy(authBytes[2:18], userID[:])
  155. copy(authBytes[18:], token)
  156. if _, err := auth.Write(authBytes); err != nil {
  157. t.Fatal(err)
  158. }
  159. closeUniStream(t, auth)
  160. waitForClientCongestionSender(t, server, client, served)
  161. var serverConn *quic.Conn
  162. server.connectionsMu.Lock()
  163. for candidate := range server.connections {
  164. if matchesClientSocket(candidate, client) {
  165. serverConn = candidate
  166. break
  167. }
  168. }
  169. server.connectionsMu.Unlock()
  170. if serverConn == nil {
  171. t.Fatal("server connection missing")
  172. }
  173. actual, sender := reauditActualSender(serverConn)
  174. snap := reauditCCSnapshot{conn: serverConn, chosen: normalized, actual: actual, sender: sender}
  175. if actual != reauditWantedSender(normalized) {
  176. t.Fatalf("wrong sender: %s", actual)
  177. }
  178. tcp, err := client.OpenStreamSync(ctx)
  179. if err != nil {
  180. t.Fatal(err)
  181. }
  182. var connect bytes.Buffer
  183. connect.Write([]byte{ProtocolVersion, CmdConnect})
  184. if err := WriteAddress(&connect, &Address{Type: AddrTypeIPv4, IP: net.ParseIP("1.1.1.1"), Port: 80}); err != nil {
  185. t.Fatal(err)
  186. }
  187. if _, err := tcp.Write(connect.Bytes()); err != nil {
  188. t.Fatal(err)
  189. }
  190. for _, p := range peers {
  191. if p.snapshot.sender == snap.sender {
  192. t.Fatal("sender reused across connections")
  193. }
  194. }
  195. peers = append(peers, &peer{client: client, tcp: tcp, snapshot: snap})
  196. for i, p := range peers {
  197. actual, ptr := reauditActualSender(p.snapshot.conn)
  198. if actual != p.snapshot.actual || ptr != p.snapshot.sender {
  199. t.Fatalf("existing connection sender changed: %s -> %s", p.snapshot.actual, actual)
  200. }
  201. msg := fmt.Appendf(nil, "live-step-%d-peer-%d", step, i)
  202. tcpEcho(p, msg)
  203. udpEcho(p, false, msg)
  204. udpEcho(p, true, msg)
  205. }
  206. t.Logf("step=%d new=%s old peers=%d usable TCP/native UDP/stream UDP; listener preserved", step, snap.actual, len(peers)-1)
  207. }
  208. }
  209. func audit3StartSocksForManager(t *testing.T, expectedUser, expectedPass string, inboundID int) (string, func()) {
  210. ln, err := net.Listen("tcp", fmt.Sprintf("127.0.0.1:%d", SOCKSPortForInbound(inboundID)))
  211. if err != nil {
  212. t.Fatalf("failed to listen: %v", err)
  213. }
  214. stop := make(chan struct{})
  215. go func() {
  216. for {
  217. conn, err := ln.Accept()
  218. if err != nil {
  219. select {
  220. case <-stop:
  221. return
  222. default:
  223. return
  224. }
  225. }
  226. go handleMockSocksConn(conn, expectedUser, expectedPass)
  227. }
  228. }()
  229. return ln.Addr().String(), func() {
  230. close(stop)
  231. _ = ln.Close()
  232. }
  233. }