1
0

manager_live_traffic_test.go 7.1 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248
  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. if err := stream.Close(); err != nil {
  94. t.Fatal(err)
  95. }
  96. response, err := p.client.AcceptUniStream(ctx)
  97. if err != nil {
  98. t.Fatal(err)
  99. }
  100. reader = response
  101. } else {
  102. if err := p.client.SendDatagram(frame.Bytes()); err != nil {
  103. t.Fatal(err)
  104. }
  105. response, err := p.client.ReceiveDatagram(ctx)
  106. if err != nil {
  107. t.Fatal(err)
  108. }
  109. reader = bytes.NewReader(response)
  110. }
  111. _, command, err := ReadCommand(reader)
  112. if err != nil || command != CmdPacket {
  113. t.Fatalf("UDP response command=%d error=%v", command, err)
  114. }
  115. hdr, err := ReadPacketHeader(reader)
  116. if err != nil {
  117. t.Fatal(err)
  118. }
  119. payload, err := readPacketPayload(reader, hdr)
  120. if err != nil {
  121. t.Fatal(err)
  122. }
  123. if hdr.AssocID != assoc || !bytes.Equal(payload, message) {
  124. t.Fatalf("UDP echo mismatch association=%d payload=%q", hdr.AssocID, payload)
  125. }
  126. }
  127. for step, controller := range []string{"new_reno", "reno", "bbr", "BBR", "cubic", "CuBiC", "", "invalid"} {
  128. inst.CongestionControl = controller
  129. if err := manager.Ensure(inst); err != nil {
  130. t.Fatal(err)
  131. }
  132. normalized, _ := normalizeCongestionControl(controller)
  133. served := normalized
  134. if served == "cubic" {
  135. served = "new_reno"
  136. }
  137. if server.quicListener != listener || server.packetConn.LocalAddr().String() != address {
  138. t.Fatal("listener changed")
  139. }
  140. client, err := quic.DialAddr(ctx, address, &tls.Config{InsecureSkipVerify: true, NextProtos: []string{"h3"}}, &quic.Config{EnableDatagrams: true, MaxIdleTimeout: 30 * time.Second})
  141. if err != nil {
  142. t.Fatal(err)
  143. }
  144. defer client.CloseWithError(0, "")
  145. tlsState := client.ConnectionState().TLS
  146. token, err := tlsState.ExportKeyingMaterial(string(userID[:]), []byte("secret-reaudit"), 32)
  147. if err != nil {
  148. t.Fatal(err)
  149. }
  150. auth, err := client.OpenUniStreamSync(ctx)
  151. if err != nil {
  152. t.Fatal(err)
  153. }
  154. authBytes := make([]byte, 50)
  155. authBytes[0], authBytes[1] = ProtocolVersion, CmdAuthenticate
  156. copy(authBytes[2:18], userID[:])
  157. copy(authBytes[18:], token)
  158. if _, err := auth.Write(authBytes); err != nil {
  159. t.Fatal(err)
  160. }
  161. if err := auth.Close(); err != nil {
  162. t.Fatal(err)
  163. }
  164. waitForClientCongestionSender(t, server, client, served)
  165. var serverConn *quic.Conn
  166. server.connectionsMu.Lock()
  167. for candidate := range server.connections {
  168. if matchesClientSocket(candidate, client) {
  169. serverConn = candidate
  170. break
  171. }
  172. }
  173. server.connectionsMu.Unlock()
  174. if serverConn == nil {
  175. t.Fatal("server connection missing")
  176. }
  177. actual, sender := reauditActualSender(serverConn)
  178. snap := reauditCCSnapshot{conn: serverConn, chosen: normalized, actual: actual, sender: sender}
  179. if actual != reauditWantedSender(normalized) {
  180. t.Fatalf("wrong sender: %s", actual)
  181. }
  182. tcp, err := client.OpenStreamSync(ctx)
  183. if err != nil {
  184. t.Fatal(err)
  185. }
  186. var connect bytes.Buffer
  187. connect.Write([]byte{ProtocolVersion, CmdConnect})
  188. if err := WriteAddress(&connect, &Address{Type: AddrTypeIPv4, IP: net.ParseIP("1.1.1.1"), Port: 80}); err != nil {
  189. t.Fatal(err)
  190. }
  191. if _, err := tcp.Write(connect.Bytes()); err != nil {
  192. t.Fatal(err)
  193. }
  194. for _, p := range peers {
  195. if p.snapshot.sender == snap.sender {
  196. t.Fatal("sender reused across connections")
  197. }
  198. }
  199. peers = append(peers, &peer{client: client, tcp: tcp, snapshot: snap})
  200. for i, p := range peers {
  201. actual, ptr := reauditActualSender(p.snapshot.conn)
  202. if actual != p.snapshot.actual || ptr != p.snapshot.sender {
  203. t.Fatalf("existing connection sender changed: %s -> %s", p.snapshot.actual, actual)
  204. }
  205. msg := fmt.Appendf(nil, "live-step-%d-peer-%d", step, i)
  206. tcpEcho(p, msg)
  207. udpEcho(p, false, msg)
  208. udpEcho(p, true, msg)
  209. }
  210. t.Logf("step=%d new=%s old peers=%d usable TCP/native UDP/stream UDP; listener preserved", step, snap.actual, len(peers)-1)
  211. }
  212. }
  213. func audit3StartSocksForManager(t *testing.T, expectedUser, expectedPass string, inboundID int) (string, func()) {
  214. ln, err := net.Listen("tcp", fmt.Sprintf("127.0.0.1:%d", SOCKSPortForInbound(inboundID)))
  215. if err != nil {
  216. t.Fatalf("failed to listen: %v", err)
  217. }
  218. stop := make(chan struct{})
  219. go func() {
  220. for {
  221. conn, err := ln.Accept()
  222. if err != nil {
  223. select {
  224. case <-stop:
  225. return
  226. default:
  227. return
  228. }
  229. }
  230. go handleMockSocksConn(conn, expectedUser, expectedPass)
  231. }
  232. }()
  233. return ln.Addr().String(), func() {
  234. close(stop)
  235. _ = ln.Close()
  236. }
  237. }