socks_bridge_test.go 11 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416
  1. package tuic
  2. import (
  3. "bytes"
  4. "context"
  5. "encoding/binary"
  6. "errors"
  7. "io"
  8. "net"
  9. "strings"
  10. "sync/atomic"
  11. "testing"
  12. "time"
  13. )
  14. func TestBuildSocks5ConnectRequest(t *testing.T) {
  15. // IPv4
  16. ip4 := net.ParseIP("1.2.3.4")
  17. req4 := buildSocks5ConnectRequest(&Address{Type: AddrTypeIPv4, IP: ip4, Port: 8080})
  18. if len(req4) != 10 || req4[0] != 0x05 || req4[1] != 0x01 || req4[3] != 0x01 {
  19. t.Fatalf("unexpected IPv4 CONNECT request: %x", req4)
  20. }
  21. if binary.BigEndian.Uint16(req4[8:10]) != 8080 {
  22. t.Fatalf("expected port 8080, got %d", binary.BigEndian.Uint16(req4[8:10]))
  23. }
  24. // IPv6
  25. ip6 := net.ParseIP("2001:db8::1")
  26. req6 := buildSocks5ConnectRequest(&Address{Type: AddrTypeIPv6, IP: ip6, Port: 443})
  27. if len(req6) != 22 || req6[3] != 0x04 {
  28. t.Fatalf("unexpected IPv6 CONNECT request: %x", req6)
  29. }
  30. if binary.BigEndian.Uint16(req6[20:22]) != 443 {
  31. t.Fatalf("expected port 443, got %d", binary.BigEndian.Uint16(req6[20:22]))
  32. }
  33. // Domain
  34. reqD := buildSocks5ConnectRequest(&Address{Type: AddrTypeDomain, Host: "example.com", Port: 80})
  35. if reqD == nil || reqD[3] != 0x03 || reqD[4] != byte(len("example.com")) {
  36. t.Fatalf("unexpected Domain CONNECT request: %x", reqD)
  37. }
  38. if binary.BigEndian.Uint16(reqD[len(reqD)-2:]) != 80 {
  39. t.Fatalf("expected port 80, got %d", binary.BigEndian.Uint16(reqD[len(reqD)-2:]))
  40. }
  41. // Nil target
  42. if buildSocks5ConnectRequest(nil) != nil {
  43. t.Fatalf("expected nil for nil target")
  44. }
  45. }
  46. func TestBuildSocks5UDPHeader(t *testing.T) {
  47. // IPv4
  48. ip4 := net.ParseIP("192.168.1.1")
  49. hdr4 := buildSocks5UDPHeader(&Address{Type: AddrTypeIPv4, IP: ip4, Port: 53})
  50. if len(hdr4) != 10 || hdr4[3] != 0x01 || binary.BigEndian.Uint16(hdr4[8:10]) != 53 {
  51. t.Fatalf("unexpected IPv4 UDP header: %x", hdr4)
  52. }
  53. // IPv6
  54. ip6 := net.ParseIP("::1")
  55. hdr6 := buildSocks5UDPHeader(&Address{Type: AddrTypeIPv6, IP: ip6, Port: 5353})
  56. if len(hdr6) != 22 || hdr6[3] != 0x04 || binary.BigEndian.Uint16(hdr6[20:22]) != 5353 {
  57. t.Fatalf("unexpected IPv6 UDP header: %x", hdr6)
  58. }
  59. // Domain
  60. hdrD := buildSocks5UDPHeader(&Address{Type: AddrTypeDomain, Host: "dns.google", Port: 53})
  61. if hdrD == nil || hdrD[3] != 0x03 || hdrD[4] != byte(len("dns.google")) {
  62. t.Fatalf("unexpected Domain UDP header: %x", hdrD)
  63. }
  64. // Nil target
  65. if buildSocks5UDPHeader(nil) != nil {
  66. t.Fatalf("expected nil for nil target")
  67. }
  68. }
  69. func TestBuildSocks5UDPRequestHonorsMaximumForAddressOverhead(t *testing.T) {
  70. domain := &Address{Type: AddrTypeDomain, Host: strings.Repeat("a", 255), Port: 53}
  71. packet, err := buildSocks5UDPRequest(domain, make([]byte, maxSafeUdpRelayPacketSize))
  72. if err != nil {
  73. t.Fatalf("maximum safe payload was rejected: %v", err)
  74. }
  75. if len(packet) != maxSocksUdpDatagramSize {
  76. t.Fatalf("encoded SOCKS datagram = %d bytes, want %d", len(packet), maxSocksUdpDatagramSize)
  77. }
  78. if _, err := buildSocks5UDPRequest(domain, make([]byte, maxSafeUdpRelayPacketSize+1)); !errors.Is(err, ErrUdpPayloadTooLarge) {
  79. t.Fatalf("oversized SOCKS datagram error = %v, want %v", err, ErrUdpPayloadTooLarge)
  80. }
  81. }
  82. func TestCountingConn(t *testing.T) {
  83. serverConn, clientConn := net.Pipe()
  84. defer serverConn.Close()
  85. defer clientConn.Close()
  86. var bytesRead atomic.Int64
  87. var bytesWritten atomic.Int64
  88. c := &CountingConn{
  89. Conn: clientConn,
  90. bytesRead: &bytesRead,
  91. bytesWritten: &bytesWritten,
  92. }
  93. go func() {
  94. buf := make([]byte, 100)
  95. n, _ := serverConn.Read(buf)
  96. _, _ = serverConn.Write(buf[:n])
  97. }()
  98. msg := []byte("hello counting conn")
  99. n, err := c.Write(msg)
  100. if err != nil || n != len(msg) {
  101. t.Fatalf("write failed: %v", err)
  102. }
  103. if bytesWritten.Load() != int64(len(msg)) {
  104. t.Fatalf("expected %d written, got %d", len(msg), bytesWritten.Load())
  105. }
  106. resp := make([]byte, 100)
  107. rn, err := c.Read(resp)
  108. if err != nil || rn != len(msg) {
  109. t.Fatalf("read failed: %v", err)
  110. }
  111. if bytesRead.Load() != int64(len(msg)) {
  112. t.Fatalf("expected %d read, got %d", len(msg), bytesRead.Load())
  113. }
  114. }
  115. func TestPipeBiDirectional(t *testing.T) {
  116. a1, a2 := net.Pipe()
  117. b1, b2 := net.Pipe()
  118. var up, down atomic.Int64
  119. done := make(chan struct{})
  120. go func() {
  121. PipeBiDirectional(a1, b1, &up, &down)
  122. close(done)
  123. }()
  124. // Send from a2 -> a1 -> b1 -> b2 (upload)
  125. testDataUp := []byte("upload stream test")
  126. go func() {
  127. _, _ = a2.Write(testDataUp)
  128. }()
  129. bufUp := make([]byte, len(testDataUp))
  130. _, err := io.ReadFull(b2, bufUp)
  131. if err != nil || !bytes.Equal(bufUp, testDataUp) {
  132. t.Fatalf("upload read failed: %v", err)
  133. }
  134. // Send from b2 -> b1 -> a1 -> a2 (download)
  135. testDataDown := []byte("download stream test")
  136. go func() {
  137. _, _ = b2.Write(testDataDown)
  138. }()
  139. bufDown := make([]byte, len(testDataDown))
  140. _, err = io.ReadFull(a2, bufDown)
  141. if err != nil || !bytes.Equal(bufDown, testDataDown) {
  142. t.Fatalf("download read failed: %v", err)
  143. }
  144. _ = a2.Close()
  145. _ = b2.Close()
  146. select {
  147. case <-done:
  148. case <-time.After(2 * time.Second):
  149. t.Fatal("PipeBiDirectional timed out waiting to finish")
  150. }
  151. if up.Load() < int64(len(testDataUp)) {
  152. t.Fatalf("expected at least %d up, got %d", len(testDataUp), up.Load())
  153. }
  154. if down.Load() < int64(len(testDataDown)) {
  155. t.Fatalf("expected at least %d down, got %d", len(testDataDown), down.Load())
  156. }
  157. }
  158. func startMockSocks5Server(t *testing.T, expectedUser, expectedPass string) (string, func()) {
  159. ln, err := net.Listen("tcp", "127.0.0.1:0")
  160. if err != nil {
  161. t.Fatalf("failed to listen: %v", err)
  162. }
  163. stop := make(chan struct{})
  164. go func() {
  165. for {
  166. conn, err := ln.Accept()
  167. if err != nil {
  168. select {
  169. case <-stop:
  170. return
  171. default:
  172. return
  173. }
  174. }
  175. go handleMockSocksConn(conn, expectedUser, expectedPass)
  176. }
  177. }()
  178. return ln.Addr().String(), func() {
  179. close(stop)
  180. _ = ln.Close()
  181. }
  182. }
  183. func handleMockSocksConn(conn net.Conn, expectedUser, expectedPass string) {
  184. defer conn.Close()
  185. // Read greeting
  186. var greeting [4]byte
  187. if _, err := io.ReadFull(conn, greeting[:]); err != nil {
  188. return
  189. }
  190. // Select user/password auth (0x02)
  191. if _, err := conn.Write([]byte{0x05, 0x02}); err != nil {
  192. return
  193. }
  194. // Auth negotiation
  195. var authVer [2]byte
  196. if _, err := io.ReadFull(conn, authVer[:]); err != nil {
  197. return
  198. }
  199. uLen := int(authVer[1])
  200. user := make([]byte, uLen)
  201. if _, err := io.ReadFull(conn, user); err != nil {
  202. return
  203. }
  204. var pLen [1]byte
  205. if _, err := io.ReadFull(conn, pLen[:]); err != nil {
  206. return
  207. }
  208. pass := make([]byte, int(pLen[0]))
  209. if _, err := io.ReadFull(conn, pass); err != nil {
  210. return
  211. }
  212. if string(user) != expectedUser || string(pass) != expectedPass {
  213. _, _ = conn.Write([]byte{0x01, 0x01}) // auth failure
  214. return
  215. }
  216. _, _ = conn.Write([]byte{0x01, 0x00}) // auth success
  217. // Read command
  218. var cmdHdr [4]byte
  219. if _, err := io.ReadFull(conn, cmdHdr[:]); err != nil {
  220. return
  221. }
  222. cmd := cmdHdr[1]
  223. atyp := cmdHdr[3]
  224. // Read dest address
  225. switch atyp {
  226. case 0x01:
  227. var ip [4]byte
  228. _, _ = io.ReadFull(conn, ip[:])
  229. case 0x04:
  230. var ip [16]byte
  231. _, _ = io.ReadFull(conn, ip[:])
  232. case 0x03:
  233. var dLen [1]byte
  234. _, _ = io.ReadFull(conn, dLen[:])
  235. domain := make([]byte, dLen[0])
  236. _, _ = io.ReadFull(conn, domain)
  237. }
  238. var port [2]byte
  239. _, _ = io.ReadFull(conn, port[:])
  240. switch cmd {
  241. case 0x01: // CONNECT
  242. // Send success reply: 0x05 0x00 0x00 0x01 (IPv4 127.0.0.1:0)
  243. _, _ = conn.Write([]byte{0x05, 0x00, 0x00, 0x01, 127, 0, 0, 1, 0x1f, 0x90})
  244. // Echo server for testing
  245. _, _ = io.Copy(conn, conn)
  246. case 0x03: // UDP ASSOCIATE
  247. // Bind a UDP listener for the mock
  248. u, err := net.ListenUDP("udp", &net.UDPAddr{IP: net.ParseIP("127.0.0.1"), Port: 0})
  249. if err != nil {
  250. return
  251. }
  252. defer u.Close()
  253. bindAddr := u.LocalAddr().(*net.UDPAddr)
  254. bindPort := uint16(bindAddr.Port)
  255. resp := make([]byte, 10)
  256. resp[0] = 0x05
  257. resp[1] = 0x00
  258. resp[2] = 0x00
  259. resp[3] = 0x01
  260. copy(resp[4:8], bindAddr.IP.To4())
  261. binary.BigEndian.PutUint16(resp[8:10], bindPort)
  262. if _, err := conn.Write(resp); err != nil {
  263. return
  264. }
  265. go func() {
  266. buf := make([]byte, maxUdpRelayPacketSize)
  267. for {
  268. n, remoteAddr, err := u.ReadFrom(buf)
  269. if err != nil {
  270. return
  271. }
  272. _, _ = u.WriteTo(buf[:n], remoteAddr)
  273. }
  274. }()
  275. // Keep conn open until closed
  276. buf := make([]byte, 1)
  277. _, _ = conn.Read(buf)
  278. }
  279. }
  280. func TestSocksRelayDialTCP(t *testing.T) {
  281. addr, cleanup := startMockSocks5Server(t, "[email protected]", "secretpass")
  282. defer cleanup()
  283. relay := &SocksRelay{
  284. Addr: addr,
  285. Password: "secretpass",
  286. }
  287. target := &Address{
  288. Type: AddrTypeIPv4,
  289. IP: net.ParseIP("93.184.216.34"),
  290. Port: 80,
  291. }
  292. ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second)
  293. defer cancel()
  294. conn, err := relay.DialTCP(ctx, "[email protected]", target)
  295. if err != nil {
  296. t.Fatalf("DialTCP failed: %v", err)
  297. }
  298. defer conn.Close()
  299. // Send echo payload
  300. msg := []byte("ping through socks")
  301. if _, err := conn.Write(msg); err != nil {
  302. t.Fatalf("write failed: %v", err)
  303. }
  304. reply := make([]byte, len(msg))
  305. if _, err := io.ReadFull(conn, reply); err != nil {
  306. t.Fatalf("read failed: %v", err)
  307. }
  308. if !bytes.Equal(reply, msg) {
  309. t.Fatalf("expected %q, got %q", msg, reply)
  310. }
  311. }
  312. func TestSocksRelayDialUDP(t *testing.T) {
  313. addr, cleanup := startMockSocks5Server(t, "[email protected]", "secretpass")
  314. defer cleanup()
  315. relay := &SocksRelay{
  316. Addr: addr,
  317. Password: "secretpass",
  318. }
  319. ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second)
  320. defer cancel()
  321. session, err := relay.DialUDP(ctx, "[email protected]")
  322. if err != nil {
  323. t.Fatalf("DialUDP failed: %v", err)
  324. }
  325. defer session.Close()
  326. target := &Address{
  327. Type: AddrTypeIPv4,
  328. IP: net.ParseIP("8.8.8.8"),
  329. Port: 53,
  330. }
  331. payload := []byte("dns packet payload")
  332. n, err := session.Send(target, payload)
  333. if err != nil || n == 0 {
  334. t.Fatalf("Send failed: %v", err)
  335. }
  336. buf := make([]byte, 2048)
  337. recvAddr, recvPayload, err := session.Receive(buf)
  338. if err != nil {
  339. t.Fatalf("Receive failed: %v", err)
  340. }
  341. if !bytes.Equal(recvPayload, payload) {
  342. t.Fatalf("expected payload %q, got %q", payload, recvPayload)
  343. }
  344. if recvAddr.IP.String() != "8.8.8.8" || recvAddr.Port != 53 {
  345. t.Fatalf("unexpected addr: %v", recvAddr)
  346. }
  347. }
  348. func TestSOCKSPortForInboundKeepsEverySlotInsideTheWindow(t *testing.T) {
  349. for id := 1; id <= 3000; id++ {
  350. port := SOCKSPortForInbound(id)
  351. if port < 64001 || port > 65000 {
  352. t.Fatalf("id %d derived port %d outside window [64001, 65000]", id, port)
  353. }
  354. }
  355. if got := SOCKSPortForInbound(1); got != 64001 {
  356. t.Fatalf("expected 64001 for id 1, got %d", got)
  357. }
  358. if got := SOCKSPortForInbound(1000); got != 65000 {
  359. t.Fatalf("expected 65000 for id 1000, got %d", got)
  360. }
  361. if got := SOCKSPortForInbound(1001); got != 64001 {
  362. t.Fatalf("expected 64001 for id 1001, got %d", got)
  363. }
  364. if got := SOCKSPortForInbound(0); got != 64001 {
  365. t.Fatalf("expected 64001 for id 0, got %d", got)
  366. }
  367. }