//go:build linux || freebsd // +build linux freebsd package transport import ( "bytes" "context" "crypto/tls" "errors" "fmt" "io" "net" "net/http" "net/url" "os" "strings" "sync" "testing" "time" ) // --------------------------------------------------------------------------- // Dialer tests // --------------------------------------------------------------------------- func TestDefaultDialerConfig(t *testing.T) { t.Parallel() cfg := DefaultDialerConfig() if cfg.Timeout != 30*time.Second { t.Errorf("expected Timeout=30s, got %v", cfg.Timeout) } if cfg.IPv6Timeout != 5*time.Second { t.Errorf("expected IPv6Timeout=5s, got %v", cfg.IPv6Timeout) } if !cfg.IPv4Fallback { t.Error("expected IPv4Fallback=true") } if cfg.Debug { t.Error("expected Debug=false") } if cfg.OnConnect != nil { t.Error("expected OnConnect=nil") } } func TestNewDialer(t *testing.T) { t.Parallel() t.Run("with nil config uses defaults", func(t *testing.T) { d := NewDialer(nil) if d == nil { t.Fatal("expected non-nil dialer") } if d.config == nil { t.Fatal("expected non-nil config") } if d.config.Timeout != 30*time.Second { t.Errorf("expected default timeout, got %v", d.config.Timeout) } }) t.Run("with custom config", func(t *testing.T) { cfg := &DialerConfig{ Timeout: 10 * time.Second, IPv6Timeout: 2 * time.Second, IPv4Fallback: false, Debug: true, } d := NewDialer(cfg) if d == nil { t.Fatal("expected non-nil dialer") } if d.config.Timeout != 10*time.Second { t.Errorf("expected 10s, got %v", d.config.Timeout) } if d.config.IPv6Timeout != 2*time.Second { t.Errorf("expected 2s, got %v", d.config.IPv6Timeout) } if d.config.IPv4Fallback { t.Error("expected IPv4Fallback=false") } if !d.config.Debug { t.Error("expected Debug=true") } }) t.Run("resolver is initialized", func(t *testing.T) { d := NewDialer(DefaultDialerConfig()) if d.resolver == nil { t.Fatal("expected resolver to be initialized") } }) } func TestNewDialerPreservesOnConnectCallback(t *testing.T) { t.Parallel() var ( mu sync.Mutex called bool gotIP string gotVer int ) cfg := &DialerConfig{ Timeout: 5 * time.Second, IPv6Timeout: 1 * time.Second, IPv4Fallback: true, OnConnect: func(ip string, version int) { mu.Lock() called = true gotIP = ip gotVer = version mu.Unlock() }, } d := NewDialer(cfg) if d == nil { t.Fatal("expected non-nil dialer") } // Verify the callback is stored (actual connection is not tested here) if d.config.OnConnect == nil { t.Fatal("expected OnConnect callback to be preserved") } // Call it manually to verify it works d.config.OnConnect("192.0.2.1", 4) mu.Lock() if !called { t.Fatal("expected OnConnect to have been called") } if gotIP != "192.0.2.1" { t.Errorf("expected IP 192.0.2.1, got %s", gotIP) } if gotVer != 4 { t.Errorf("expected version 4, got %d", gotVer) } mu.Unlock() } func TestParseProxyURL(t *testing.T) { t.Parallel() t.Run("empty string returns nil, nil", func(t *testing.T) { u, err := ParseProxyURL("") if err != nil { t.Errorf("expected nil error, got %v", err) } if u != nil { t.Errorf("expected nil URL, got %v", u) } }) t.Run("valid HTTP proxy URL", func(t *testing.T) { u, err := ParseProxyURL("http://proxy.example.com:8080") if err != nil { t.Fatalf("unexpected error: %v", err) } if u == nil { t.Fatal("expected non-nil URL") } if u.Scheme != "http" { t.Errorf("expected scheme http, got %s", u.Scheme) } if u.Host != "proxy.example.com:8080" { t.Errorf("expected host proxy.example.com:8080, got %s", u.Host) } }) t.Run("valid HTTPS proxy URL with auth", func(t *testing.T) { u, err := ParseProxyURL("https://user:pass@proxy.example.com:443") if err != nil { t.Fatalf("unexpected error: %v", err) } if u == nil { t.Fatal("expected non-nil URL") } if u.Scheme != "https" { t.Errorf("expected scheme https, got %s", u.Scheme) } if u.User == nil { t.Fatal("expected user info") } password, _ := u.User.Password() if u.User.Username() != "user" || password != "pass" { t.Errorf("expected user:pass, got %s:%s", u.User.Username(), password) } }) t.Run("invalid URL returns error", func(t *testing.T) { _, err := ParseProxyURL("://invalid") if err == nil { t.Error("expected error for invalid URL") } }) t.Run("URL without host returns error", func(t *testing.T) { _, err := ParseProxyURL("http://") if err == nil { t.Error("expected error for URL without host") } }) t.Run("socks5 proxy URL", func(t *testing.T) { u, err := ParseProxyURL("socks5://127.0.0.1:1080") if err != nil { t.Fatalf("unexpected error: %v", err) } if u == nil { t.Fatal("expected non-nil URL") } if u.Scheme != "socks5" { t.Errorf("expected scheme socks5, got %s", u.Scheme) } if u.Host != "127.0.0.1:1080" { t.Errorf("expected host 127.0.0.1:1080, got %s", u.Host) } }) } // --------------------------------------------------------------------------- // Proxy tests // --------------------------------------------------------------------------- func TestParseProxyConfig(t *testing.T) { t.Parallel() t.Run("empty string returns nil, nil", func(t *testing.T) { cfg, err := ParseProxyConfig("") if err != nil { t.Errorf("expected nil error, got %v", err) } if cfg != nil { t.Errorf("expected nil config, got %+v", cfg) } }) t.Run("HTTP proxy URL", func(t *testing.T) { cfg, err := ParseProxyConfig("http://proxy.example.com:3128") if err != nil { t.Fatalf("unexpected error: %v", err) } if cfg == nil { t.Fatal("expected non-nil config") } if cfg.Type != ProxyHTTP { t.Errorf("expected ProxyHTTP, got %s", cfg.Type) } if cfg.URL.Host != "proxy.example.com:3128" { t.Errorf("expected host proxy.example.com:3128, got %s", cfg.URL.Host) } }) t.Run("HTTPS proxy URL maps to ProxyHTTP type", func(t *testing.T) { cfg, err := ParseProxyConfig("https://secure-proxy.example.com:443") if err != nil { t.Fatalf("unexpected error: %v", err) } if cfg == nil { t.Fatal("expected non-nil config") } if cfg.Type != ProxyHTTP { t.Errorf("expected ProxyHTTP for https scheme, got %s", cfg.Type) } }) t.Run("SOCKS5 proxy URL", func(t *testing.T) { cfg, err := ParseProxyConfig("socks5://127.0.0.1:1080") if err != nil { t.Fatalf("unexpected error: %v", err) } if cfg == nil { t.Fatal("expected non-nil config") } if cfg.Type != ProxySOCKS5 { t.Errorf("expected ProxySOCKS5, got %s", cfg.Type) } if !cfg.SOCKS5LocalResolve { t.Error("expected SOCKS5LocalResolve=true for socks5:// scheme") } }) t.Run("SOCKS5h proxy URL", func(t *testing.T) { cfg, err := ParseProxyConfig("socks5h://127.0.0.1:1080") if err != nil { t.Fatalf("unexpected error: %v", err) } if cfg == nil { t.Fatal("expected non-nil config") } if cfg.Type != ProxySOCKS5 { t.Errorf("expected ProxySOCKS5 for socks5h scheme, got %s", cfg.Type) } if cfg.SOCKS5LocalResolve { t.Error("expected SOCKS5LocalResolve=false for socks5h:// scheme") } }) t.Run("unknown scheme defaults to ProxyHTTP", func(t *testing.T) { cfg, err := ParseProxyConfig("ftp://proxy.example.com:21") if err != nil { t.Fatalf("unexpected error: %v", err) } if cfg == nil { t.Fatal("expected non-nil config") } if cfg.Type != ProxyHTTP { t.Errorf("expected ProxyHTTP for unknown scheme, got %s", cfg.Type) } }) t.Run("with authentication", func(t *testing.T) { cfg, err := ParseProxyConfig("http://alice:secret@proxy.example.com:8080") if err != nil { t.Fatalf("unexpected error: %v", err) } if cfg == nil { t.Fatal("expected non-nil config") } if cfg.Username != "alice" { t.Errorf("expected username alice, got %s", cfg.Username) } if cfg.Password != "secret" { t.Errorf("expected password secret, got %s", cfg.Password) } // Auth should be stripped from URL if cfg.URL.User != nil { t.Error("expected User to be nil in stored URL") } }) t.Run("with username only (no password)", func(t *testing.T) { cfg, err := ParseProxyConfig("http://alice@proxy.example.com:8080") if err != nil { t.Fatalf("unexpected error: %v", err) } if cfg == nil { t.Fatal("expected non-nil config") } if cfg.Username != "alice" { t.Errorf("expected username alice, got %s", cfg.Username) } if cfg.Password != "" { t.Errorf("expected empty password, got %s", cfg.Password) } }) t.Run("invalid URL returns error", func(t *testing.T) { _, err := ParseProxyConfig("://bad") if err == nil { t.Error("expected error for invalid URL") } }) } func TestGetProxyFromEnv(t *testing.T) { // Start with a clean environment origHTTPS := os.Getenv("HTTPS_PROXY") origHTTP := os.Getenv("HTTP_PROXY") origHTTPLower := os.Getenv("http_proxy") origHTTPSSecure := os.Getenv("https_proxy") defer func() { os.Setenv("HTTPS_PROXY", origHTTPS) os.Setenv("HTTP_PROXY", origHTTP) os.Setenv("http_proxy", origHTTPLower) os.Setenv("https_proxy", origHTTPSSecure) }() t.Run("no proxy env vars returns empty", func(t *testing.T) { os.Unsetenv("HTTPS_PROXY") os.Unsetenv("https_proxy") os.Unsetenv("HTTP_PROXY") os.Unsetenv("http_proxy") result := GetProxyFromEnv() if result != "" { t.Errorf("expected empty string, got %s", result) } }) t.Run("reads HTTPS_PROXY first", func(t *testing.T) { os.Unsetenv("HTTPS_PROXY") os.Unsetenv("https_proxy") os.Unsetenv("HTTP_PROXY") os.Unsetenv("http_proxy") os.Setenv("HTTPS_PROXY", "https://secure-proxy:443") os.Setenv("HTTP_PROXY", "http://proxy:8080") result := GetProxyFromEnv() if result != "https://secure-proxy:443" { t.Errorf("expected https://secure-proxy:443, got %s", result) } }) t.Run("reads HTTP_PROXY when HTTPS_PROXY is empty", func(t *testing.T) { os.Unsetenv("HTTPS_PROXY") os.Unsetenv("https_proxy") os.Unsetenv("HTTP_PROXY") os.Unsetenv("http_proxy") os.Setenv("HTTP_PROXY", "http://proxy:8080") result := GetProxyFromEnv() if result != "http://proxy:8080" { t.Errorf("expected http://proxy:8080, got %s", result) } }) t.Run("reads http_proxy (lowercase) when others are empty", func(t *testing.T) { os.Unsetenv("HTTPS_PROXY") os.Unsetenv("https_proxy") os.Unsetenv("HTTP_PROXY") os.Unsetenv("http_proxy") os.Setenv("http_proxy", "http://lower-proxy:3128") result := GetProxyFromEnv() if result != "http://lower-proxy:3128" { t.Errorf("expected http://lower-proxy:3128, got %s", result) } }) t.Run("trims whitespace from proxy value", func(t *testing.T) { os.Unsetenv("HTTPS_PROXY") os.Unsetenv("https_proxy") os.Unsetenv("HTTP_PROXY") os.Unsetenv("http_proxy") os.Setenv("HTTP_PROXY", " http://proxy:8080 ") result := GetProxyFromEnv() if result != "http://proxy:8080" { t.Errorf("expected http://proxy:8080 (trimmed), got %q", result) } }) t.Run("priority: HTTPS_PROXY > https_proxy > HTTP_PROXY > http_proxy", func(t *testing.T) { os.Unsetenv("HTTPS_PROXY") os.Unsetenv("https_proxy") os.Unsetenv("HTTP_PROXY") os.Unsetenv("http_proxy") os.Setenv("https_proxy", "https://lower-secure:443") os.Setenv("HTTP_PROXY", "http://proxy:8080") result := GetProxyFromEnv() if result != "https://lower-secure:443" { t.Errorf("expected https://lower-secure:443, got %s", result) } }) } func TestNewProxyTransport(t *testing.T) { t.Parallel() t.Run("nil config returns base transport unchanged", func(t *testing.T) { base := http.DefaultTransport.(*http.Transport).Clone() result, err := NewProxyTransport(nil, base) if err != nil { t.Fatalf("unexpected error: %v", err) } if result != base { t.Error("expected same transport back") } }) t.Run("sets HTTP proxy", func(t *testing.T) { base := http.DefaultTransport.(*http.Transport).Clone() cfg := &ProxyConfig{ URL: &url.URL{Scheme: "http", Host: "proxy.example.com:8080"}, Type: ProxyHTTP, } result, err := NewProxyTransport(cfg, base) if err != nil { t.Fatalf("unexpected error: %v", err) } if result == nil { t.Fatal("expected non-nil transport") } // Verify proxy function works req := &http.Request{ URL: &url.URL{Scheme: "http", Host: "example.com"}, } proxyURL, err := result.Proxy(req) if err != nil { t.Fatalf("proxy function error: %v", err) } if proxyURL == nil { t.Fatal("expected non-nil proxy URL") } if proxyURL.Host != "proxy.example.com:8080" { t.Errorf("expected proxy host proxy.example.com:8080, got %s", proxyURL.Host) } }) t.Run("sets proxy with authentication", func(t *testing.T) { base := http.DefaultTransport.(*http.Transport).Clone() cfg := &ProxyConfig{ URL: &url.URL{Scheme: "http", Host: "proxy.example.com:8080"}, Type: ProxyHTTP, Username: "user", Password: "pass", } result, err := NewProxyTransport(cfg, base) if err != nil { t.Fatalf("unexpected error: %v", err) } if result == nil { t.Fatal("expected non-nil transport") } req := &http.Request{ URL: &url.URL{Scheme: "http", Host: "example.com"}, } proxyURL, err := result.Proxy(req) if err != nil { t.Fatalf("proxy function error: %v", err) } if proxyURL == nil { t.Fatal("expected non-nil proxy URL") } if proxyURL.User == nil { t.Fatal("expected user info in proxy URL") } password, _ := proxyURL.User.Password() if proxyURL.User.Username() != "user" || password != "pass" { t.Errorf("expected user:pass, got %s:%s", proxyURL.User.Username(), password) } }) t.Run("sets proxy with username only (no password)", func(t *testing.T) { base := http.DefaultTransport.(*http.Transport).Clone() cfg := &ProxyConfig{ URL: &url.URL{Scheme: "http", Host: "proxy.example.com:8080"}, Type: ProxyHTTP, Username: "alice", } result, err := NewProxyTransport(cfg, base) if err != nil { t.Fatalf("unexpected error: %v", err) } if result == nil { t.Fatal("expected non-nil transport") } req := &http.Request{ URL: &url.URL{Scheme: "http", Host: "example.com"}, } proxyURL, err := result.Proxy(req) if err != nil { t.Fatalf("proxy function error: %v", err) } if proxyURL == nil { t.Fatal("expected non-nil proxy URL") } if proxyURL.User == nil { t.Fatal("expected user info in proxy URL") } if proxyURL.User.Username() != "alice" { t.Errorf("expected username alice, got %s", proxyURL.User.Username()) } }) t.Run("proxy URL parse failure returns error", func(t *testing.T) { // This is hard to trigger with normal inputs; skip edge case. // The url.Parse function is very lenient, so this test is for completeness. cfg := &ProxyConfig{ URL: &url.URL{Scheme: "http", Host: "proxy:8080"}, Type: ProxyHTTP, Username: "user\x00", // null byte may cause issues } base := http.DefaultTransport.(*http.Transport).Clone() _, err := NewProxyTransport(cfg, base) // We allow error or no error here; just make sure it doesn't panic _ = err }) t.Run("HTTP proxy preserves TLSClientConfig", func(t *testing.T) { // Regression guard for the BACKLOG entry "TLS config lost through // SOCKS5 proxy transport". The original concern was that // NewProxyTransport might construct a fresh transport without // copying the base's TLSClientConfig, silently dropping // InsecureSkipVerify, custom CAs, and pinned certs. Verify that // is *not* the case for HTTP proxy too. base := http.DefaultTransport.(*http.Transport).Clone() base.TLSClientConfig = &tls.Config{ InsecureSkipVerify: true, //nolint:gosec // test-only MinVersion: tls.VersionTLS12, } cfg := &ProxyConfig{ URL: &url.URL{Scheme: "http", Host: "proxy.example.com:8080"}, Type: ProxyHTTP, } result, err := NewProxyTransport(cfg, base) if err != nil { t.Fatalf("unexpected error: %v", err) } if result.TLSClientConfig == nil { t.Fatal("TLSClientConfig was dropped from the resulting transport") } if !result.TLSClientConfig.InsecureSkipVerify { t.Error("InsecureSkipVerify was dropped from TLSClientConfig") } if result.TLSClientConfig.MinVersion != tls.VersionTLS12 { t.Errorf("MinVersion = %d, want %d", result.TLSClientConfig.MinVersion, tls.VersionTLS12) } }) } func TestProxyTypes(t *testing.T) { t.Parallel() if string(ProxyHTTP) != "http" { t.Errorf("expected ProxyHTTP to be 'http', got %q", string(ProxyHTTP)) } if string(ProxyHTTPS) != "https" { t.Errorf("expected ProxyHTTPS to be 'https', got %q", string(ProxyHTTPS)) } if string(ProxySOCKS5) != "socks5" { t.Errorf("expected ProxySOCKS5 to be 'socks5', got %q", string(ProxySOCKS5)) } } func TestDialProxy(t *testing.T) { t.Parallel() t.Run("nil config uses direct dialer", func(t *testing.T) { ctx, cancel := context.WithTimeout(context.Background(), 100*time.Millisecond) defer cancel() // This will fail quickly because there's no server, but it should // attempt a direct connection rather than returning an unsupported error. _, err := DialProxy(ctx, "tcp", "127.0.0.1:1", nil) // Should NOT get "unsupported proxy type" if err != nil && strings.Contains(err.Error(), "unsupported proxy type") { t.Errorf("unexpected 'unsupported proxy type' error: %v", err) } }) t.Run("unsupported network returns error", func(t *testing.T) { ctx := context.Background() cfg := &ProxyConfig{ URL: &url.URL{Scheme: "http", Host: "127.0.0.1:8080"}, Type: ProxyHTTP, } _, err := DialProxy(ctx, "udp", "127.0.0.1:80", cfg) if err == nil { t.Error("expected error for unsupported network") } }) } // --------------------------------------------------------------------------- // SOCKS5 tests // --------------------------------------------------------------------------- // mockSOCKS5Server accepts a single SOCKS5 connection, performs (or skips) // the auth negotiation, then verifies the CONNECT request and replies with // success. The connection is then handed off to the caller via the returned // channel. type mockSOCKS5Result struct { host string port uint16 addrType byte gotUsername string gotPassword string requestedAuth byte conn net.Conn } func startMockSOCKS5(t *testing.T, requireAuth bool) (addr string, results <-chan mockSOCKS5Result, stop func()) { t.Helper() ln, err := net.Listen("tcp", "127.0.0.1:0") if err != nil { t.Fatalf("listen: %v", err) } out := make(chan mockSOCKS5Result, 1) done := make(chan struct{}) go func() { defer close(done) conn, err := ln.Accept() if err != nil { return } defer conn.Close() res := mockSOCKS5Result{conn: conn} // Greeting: read VER + NMETHODS, then NMETHODS method bytes. // The client offers at least no-auth (0x00), and additionally // user/pass (0x02) when a username is configured. greet := make([]byte, 2) if _, err := io.ReadFull(conn, greet); err != nil { return } if greet[0] != 0x05 || int(greet[1]) > 4 { return } methods := make([]byte, greet[1]) if _, err := io.ReadFull(conn, methods); err != nil { return } // Pick the strongest method we support. requireAuth means we // require user/pass; otherwise we accept no-auth if offered. offered := make(map[byte]bool, len(methods)) for _, m := range methods { offered[m] = true } if requireAuth { if !offered[0x02] { conn.Write([]byte{0x05, 0xFF}) // no acceptable method return } conn.Write([]byte{0x05, 0x02}) res.requestedAuth = 0x02 } else { if !offered[0x00] { conn.Write([]byte{0x05, 0xFF}) return } conn.Write([]byte{0x05, 0x00}) res.requestedAuth = 0x00 } // User/pass sub-negotiation (only when 0x02 was selected). if res.requestedAuth == 0x02 { ver := make([]byte, 1) if _, err := io.ReadFull(conn, ver); err != nil { return } if ver[0] != 0x01 { return } ulen := make([]byte, 1) if _, err := io.ReadFull(conn, ulen); err != nil { return } uname := make([]byte, ulen[0]) if _, err := io.ReadFull(conn, uname); err != nil { return } plen := make([]byte, 1) if _, err := io.ReadFull(conn, plen); err != nil { return } pass := make([]byte, plen[0]) if _, err := io.ReadFull(conn, pass); err != nil { return } res.gotUsername = string(uname) res.gotPassword = string(pass) conn.Write([]byte{0x01, 0x00}) } // CONNECT request: VER, CMD, RSV, ATYP hdr := make([]byte, 4) if _, err := io.ReadFull(conn, hdr); err != nil { return } if hdr[0] != 0x05 || hdr[1] != 0x01 { return } res.addrType = hdr[3] switch hdr[3] { case 0x01: ip := make([]byte, 4) if _, err := io.ReadFull(conn, ip); err != nil { return } res.host = net.IP(ip).String() case 0x04: ip := make([]byte, 16) if _, err := io.ReadFull(conn, ip); err != nil { return } res.host = net.IP(ip).String() case 0x03: l := make([]byte, 1) if _, err := io.ReadFull(conn, l); err != nil { return } name := make([]byte, l[0]) if _, err := io.ReadFull(conn, name); err != nil { return } res.host = string(name) } portBuf := make([]byte, 2) if _, err := io.ReadFull(conn, portBuf); err != nil { return } res.port = uint16(portBuf[0])<<8 | uint16(portBuf[1]) // Reply success with 0.0.0.0:0 bind address. conn.Write([]byte{0x05, 0x00, 0x00, 0x01, 0, 0, 0, 0, 0, 0}) out <- res // Keep the connection open until the test closes it. // Block here so the conn is not garbage-collected. buf := make([]byte, 1) conn.Read(buf) }() stop = func() { ln.Close() <-done } return ln.Addr().String(), out, stop } func TestNewProxyTransport_SOCKS5(t *testing.T) { t.Parallel() t.Run("SOCKS5 sets DialContext and clears Proxy field", func(t *testing.T) { base := http.DefaultTransport.(*http.Transport).Clone() cfg := &ProxyConfig{ URL: &url.URL{Scheme: "socks5", Host: "127.0.0.1:1080"}, Type: ProxySOCKS5, } result, err := NewProxyTransport(cfg, base) if err != nil { t.Fatalf("unexpected error: %v", err) } if result == nil { t.Fatal("expected non-nil transport") } if result.Proxy != nil { t.Errorf("expected Proxy=nil for SOCKS5 (handled via DialContext), got %T", result.Proxy) } if result.DialContext == nil { t.Error("expected DialContext to be set for SOCKS5") } }) t.Run("SOCKS5 with auth", func(t *testing.T) { base := http.DefaultTransport.(*http.Transport).Clone() cfg := &ProxyConfig{ URL: &url.URL{Scheme: "socks5", Host: "127.0.0.1:1080"}, Type: ProxySOCKS5, Username: "alice", Password: "secret", } result, err := NewProxyTransport(cfg, base) if err != nil { t.Fatalf("unexpected error: %v", err) } if result.DialContext == nil { t.Error("expected DialContext to be set for SOCKS5 with auth") } }) t.Run("SOCKS5 proxy preserves TLSClientConfig", func(t *testing.T) { // Regression guard for the BACKLOG entry "TLS config lost through // SOCKS5 proxy transport". After Clone(), the new transport // must still carry the caller's TLSClientConfig (with // InsecureSkipVerify, custom CAs, pinned certs, etc.). base := http.DefaultTransport.(*http.Transport).Clone() base.TLSClientConfig = &tls.Config{ InsecureSkipVerify: true, //nolint:gosec // test-only MinVersion: tls.VersionTLS12, } cfg := &ProxyConfig{ URL: &url.URL{Scheme: "socks5", Host: "127.0.0.1:1080"}, Type: ProxySOCKS5, } result, err := NewProxyTransport(cfg, base) if err != nil { t.Fatalf("unexpected error: %v", err) } if result.TLSClientConfig == nil { t.Fatal("TLSClientConfig was dropped from the SOCKS5 transport") } if !result.TLSClientConfig.InsecureSkipVerify { t.Error("InsecureSkipVerify was dropped from TLSClientConfig") } if result.TLSClientConfig.MinVersion != tls.VersionTLS12 { t.Errorf("MinVersion = %d, want %d", result.TLSClientConfig.MinVersion, tls.VersionTLS12) } }) } func TestDialProxy_SOCKS5NoAuth(t *testing.T) { t.Parallel() addr, results, stop := startMockSOCKS5(t, false) defer stop() host, port, err := net.SplitHostPort(addr) if err != nil { t.Fatalf("split addr: %v", err) } cfg := &ProxyConfig{ URL: &url.URL{Scheme: "socks5", Host: addr}, Type: ProxySOCKS5, } ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) defer cancel() conn, err := DialProxy(ctx, "tcp", net.JoinHostPort("192.0.2.1", "8080"), cfg) if err != nil { t.Fatalf("DialProxy failed: %v", err) } defer conn.Close() select { case res := <-results: if res.addrType != 0x01 { t.Errorf("expected IPv4 ATYP=0x01, got 0x%02x", res.addrType) } if res.host != "192.0.2.1" { t.Errorf("expected host 192.0.2.1, got %s", res.host) } if res.port != 8080 { t.Errorf("expected port 8080, got %d", res.port) } if res.requestedAuth != 0x00 { t.Errorf("expected no-auth selection, got 0x%02x", res.requestedAuth) } _ = host _ = port case <-time.After(2 * time.Second): t.Fatal("timed out waiting for SOCKS5 server to receive request") } } func TestDialProxy_SOCKS5WithAuth(t *testing.T) { t.Parallel() addr, results, stop := startMockSOCKS5(t, true) defer stop() cfg := &ProxyConfig{ URL: &url.URL{Scheme: "socks5", Host: addr}, Type: ProxySOCKS5, Username: "alice", Password: "s3cret", } ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) defer cancel() conn, err := DialProxy(ctx, "tcp", "192.0.2.1:8443", cfg) if err != nil { t.Fatalf("DialProxy failed: %v", err) } defer conn.Close() select { case res := <-results: if res.requestedAuth != 0x02 { t.Errorf("expected user/pass auth selection, got 0x%02x", res.requestedAuth) } if res.gotUsername != "alice" { t.Errorf("expected username alice, got %s", res.gotUsername) } if res.gotPassword != "s3cret" { t.Errorf("expected password s3cret, got %s", res.gotPassword) } if res.port != 8443 { t.Errorf("expected port 8443, got %d", res.port) } case <-time.After(2 * time.Second): t.Fatal("timed out waiting for SOCKS5 server to receive request") } } func TestDialProxy_SOCKS5HostnameRemoteResolve(t *testing.T) { t.Parallel() addr, results, stop := startMockSOCKS5(t, false) defer stop() cfg := &ProxyConfig{ URL: &url.URL{Scheme: "socks5h", Host: addr}, Type: ProxySOCKS5, SOCKS5LocalResolve: false, } ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) defer cancel() conn, err := DialProxy(ctx, "tcp", "example.com:443", cfg) if err != nil { t.Fatalf("DialProxy failed: %v", err) } defer conn.Close() select { case res := <-results: if res.addrType != 0x03 { t.Errorf("expected domain ATYP=0x03, got 0x%02x", res.addrType) } if res.host != "example.com" { t.Errorf("expected host example.com, got %s", res.host) } if res.port != 443 { t.Errorf("expected port 443, got %d", res.port) } case <-time.After(2 * time.Second): t.Fatal("timed out waiting for SOCKS5 server to receive request") } } func TestDialProxy_SOCKS5LocalResolve(t *testing.T) { t.Parallel() addr, results, stop := startMockSOCKS5(t, false) defer stop() cfg := &ProxyConfig{ URL: &url.URL{Scheme: "socks5", Host: addr}, Type: ProxySOCKS5, SOCKS5LocalResolve: true, } ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) defer cancel() // Use a hostname that always resolves to a real IP. localhost works. conn, err := DialProxy(ctx, "tcp", "localhost:8080", cfg) if err != nil { t.Fatalf("DialProxy failed: %v", err) } defer conn.Close() select { case res := <-results: // The hostname is resolved locally, so the SOCKS5 server must // receive an IP literal — never a bare hostname like "localhost". // The exact address family depends on the host resolver: most // systems return 127.0.0.1, but IPv6-preferring environments // (some CI runners) return ::1 instead. Accept either as long // as it is a loopback IP and the ATYP field matches. if res.host == "localhost" { t.Fatal("expected hostname to be resolved to IP literal, got hostname") } ip := net.ParseIP(res.host) if ip == nil { t.Fatalf("expected IP literal after local resolve, got %q", res.host) } if !ip.IsLoopback() { t.Errorf("expected loopback IP, got %s", res.host) } wantATYP := byte(0x01) // IPv4 if ip.To4() == nil { wantATYP = 0x04 // IPv6 } if res.addrType != wantATYP { t.Errorf("ATYP = 0x%02x, want 0x%02x for host %s", res.addrType, wantATYP, res.host) } case <-time.After(2 * time.Second): t.Fatal("timed out waiting for SOCKS5 server to receive request") } } // --------------------------------------------------------------------------- // Retry tests // --------------------------------------------------------------------------- func TestDefaultRetryConfig(t *testing.T) { t.Parallel() cfg := DefaultRetryConfig() if cfg.MaxRetries != 3 { t.Errorf("expected MaxRetries=3, got %d", cfg.MaxRetries) } if cfg.InitialDelay != 1*time.Second { t.Errorf("expected InitialDelay=1s, got %v", cfg.InitialDelay) } if cfg.MaxDelay != 30*time.Second { t.Errorf("expected MaxDelay=30s, got %v", cfg.MaxDelay) } if cfg.Multiplier != 2.0 { t.Errorf("expected Multiplier=2.0, got %f", cfg.Multiplier) } if len(cfg.RetryableStatusCodes) != 6 { t.Errorf("expected 6 retryable status codes, got %d", len(cfg.RetryableStatusCodes)) } if cfg.OnRetry != nil { t.Error("expected OnRetry=nil") } // Verify specific status codes expectedCodes := []int{408, 429, 500, 502, 503, 504} for i, code := range expectedCodes { if cfg.RetryableStatusCodes[i] != code { t.Errorf("expected status code %d at index %d, got %d", code, i, cfg.RetryableStatusCodes[i]) } } } func TestRetrySuccess(t *testing.T) { t.Parallel() t.Run("succeeds on first attempt", func(t *testing.T) { ctx := context.Background() attempts := 0 err := Retry(ctx, func(ctx context.Context, attempt int) error { attempts++ return nil }, &RetryConfig{ MaxRetries: 3, InitialDelay: 10 * time.Millisecond, MaxDelay: 100 * time.Millisecond, Multiplier: 2.0, }) if err != nil { t.Errorf("expected no error, got %v", err) } if attempts != 1 { t.Errorf("expected 1 attempt, got %d", attempts) } }) t.Run("succeeds after retries", func(t *testing.T) { ctx := context.Background() var attempts int err := Retry(ctx, func(ctx context.Context, attempt int) error { attempts++ if attempt < 3 { return fmt.Errorf("attempt %d failed", attempt) } return nil }, &RetryConfig{ MaxRetries: 3, InitialDelay: 5 * time.Millisecond, MaxDelay: 50 * time.Millisecond, Multiplier: 1.5, }) if err != nil { t.Errorf("expected no error, got %v", err) } if attempts != 3 { t.Errorf("expected 3 attempts, got %d", attempts) } }) t.Run("uses default config when nil", func(t *testing.T) { ctx := context.Background() err := Retry(ctx, func(ctx context.Context, attempt int) error { return nil }, nil) if err != nil { t.Errorf("expected no error, got %v", err) } }) } func TestRetryFailAll(t *testing.T) { t.Parallel() ctx := context.Background() var attempts int expectedFailures := 4 // maxRetries(3) + initial(1) = 4 err := Retry(ctx, func(ctx context.Context, attempt int) error { attempts++ return fmt.Errorf("persistent error on attempt %d", attempt) }, &RetryConfig{ MaxRetries: 3, InitialDelay: 5 * time.Millisecond, MaxDelay: 50 * time.Millisecond, Multiplier: 1.5, }) if err == nil { t.Fatal("expected an error") } if !strings.Contains(err.Error(), "all 4 attempts failed") { t.Errorf("expected 'all 4 attempts failed' in error, got %q", err.Error()) } if attempts != expectedFailures { t.Errorf("expected %d attempts, got %d", expectedFailures, attempts) } } func TestRetryWithContextCancel(t *testing.T) { t.Parallel() t.Run("cancel before retry begins", func(t *testing.T) { ctx, cancel := context.WithCancel(context.Background()) cancel() // cancel immediately err := Retry(ctx, func(ctx context.Context, attempt int) error { return fmt.Errorf("should not be called") }, &RetryConfig{ MaxRetries: 3, InitialDelay: 5 * time.Millisecond, MaxDelay: 50 * time.Millisecond, Multiplier: 1.5, }) if err == nil { t.Fatal("expected context canceled error") } if !errors.Is(err, context.Canceled) { t.Errorf("expected context.Canceled, got %v", err) } }) t.Run("cancel during retry delay", func(t *testing.T) { ctx, cancel := context.WithCancel(context.Background()) var attempts int errCh := make(chan error, 1) go func() { errCh <- Retry(ctx, func(ctx context.Context, attempt int) error { attempts++ // Cancel on second attempt if attempt == 2 { cancel() } return fmt.Errorf("failing attempt %d", attempt) }, &RetryConfig{ MaxRetries: 5, InitialDelay: 1 * time.Second, // long enough for cancel to take effect MaxDelay: 2 * time.Second, Multiplier: 1.0, }) }() select { case err := <-errCh: if err == nil { t.Fatal("expected an error") } if !errors.Is(err, context.Canceled) { t.Errorf("expected context.Canceled, got %v", err) } case <-time.After(3 * time.Second): t.Fatal("test timed out") } }) } func TestIsRetryableError(t *testing.T) { t.Parallel() t.Run("nil error is not retryable", func(t *testing.T) { if IsRetryableError(nil) { t.Error("expected false for nil error") } }) t.Run("connection refused", func(t *testing.T) { if !IsRetryableError(fmt.Errorf("connection refused")) { t.Error("expected true for 'connection refused'") } }) t.Run("connection reset", func(t *testing.T) { if !IsRetryableError(fmt.Errorf("connection reset by peer")) { t.Error("expected true for 'connection reset'") } }) t.Run("connection timed out", func(t *testing.T) { if !IsRetryableError(fmt.Errorf("connection timed out")) { t.Error("expected true for 'connection timed out'") } }) t.Run("i/o timeout", func(t *testing.T) { if !IsRetryableError(fmt.Errorf("i/o timeout")) { t.Error("expected true for 'i/o timeout'") } }) t.Run("temporary failure", func(t *testing.T) { if !IsRetryableError(fmt.Errorf("temporary failure in name resolution")) { t.Error("expected true for 'temporary failure'") } }) t.Run("no route to host", func(t *testing.T) { if !IsRetryableError(fmt.Errorf("no route to host")) { t.Error("expected true for 'no route to host'") } }) t.Run("case insensitive matching", func(t *testing.T) { if !IsRetryableError(fmt.Errorf("connection refused")) { t.Error("expected true for 'connection refused' (case insensitive)") } if !IsRetryableError(fmt.Errorf("connection timed out")) { t.Error("expected true for 'connection timed out' (case insensitive)") } }) t.Run("non-retryable error", func(t *testing.T) { if IsRetryableError(fmt.Errorf("404 not found")) { t.Error("expected false for non-retryable error") } if IsRetryableError(fmt.Errorf("permission denied")) { t.Error("expected false for permission denied") } }) t.Run("retryable substring in longer message", func(t *testing.T) { if !IsRetryableError(fmt.Errorf("dial tcp 192.168.1.1:80: connect: connection refused")) { t.Error("expected true for longer message containing 'connection refused'") } }) } func TestIsRetryableStatusCode(t *testing.T) { t.Parallel() t.Run("uses default status codes with nil config", func(t *testing.T) { if !IsRetryableStatusCode(408, nil) { t.Error("expected 408 to be retryable with nil config") } if !IsRetryableStatusCode(429, nil) { t.Error("expected 429 to be retryable with nil config") } if !IsRetryableStatusCode(500, nil) { t.Error("expected 500 to be retryable with nil config") } if !IsRetryableStatusCode(502, nil) { t.Error("expected 502 to be retryable with nil config") } if !IsRetryableStatusCode(503, nil) { t.Error("expected 503 to be retryable with nil config") } if !IsRetryableStatusCode(504, nil) { t.Error("expected 504 to be retryable with nil config") } }) t.Run("non-retryable status codes", func(t *testing.T) { if IsRetryableStatusCode(200, nil) { t.Error("expected 200 not to be retryable") } if IsRetryableStatusCode(301, nil) { t.Error("expected 301 not to be retryable") } if IsRetryableStatusCode(400, nil) { t.Error("expected 400 not to be retryable") } if IsRetryableStatusCode(401, nil) { t.Error("expected 401 not to be retryable") } if IsRetryableStatusCode(403, nil) { t.Error("expected 403 not to be retryable") } if IsRetryableStatusCode(404, nil) { t.Error("expected 404 not to be retryable") } }) t.Run("custom retryable codes", func(t *testing.T) { cfg := &RetryConfig{ RetryableStatusCodes: []int{429, 503}, } if !IsRetryableStatusCode(429, cfg) { t.Error("expected 429 to be retryable with custom config") } if !IsRetryableStatusCode(503, cfg) { t.Error("expected 503 to be retryable with custom config") } if IsRetryableStatusCode(500, cfg) { t.Error("expected 500 not to be retryable with custom config") } if IsRetryableStatusCode(502, cfg) { t.Error("expected 502 not to be retryable with custom config") } }) t.Run("empty retryable codes list", func(t *testing.T) { cfg := &RetryConfig{ RetryableStatusCodes: []int{}, } if IsRetryableStatusCode(503, cfg) { t.Error("expected 503 not to be retryable with empty codes list") } }) } func TestOnRetryCallback(t *testing.T) { t.Parallel() var ( mu sync.Mutex callbackCount int lastAttempt int lastDelay time.Duration ) cfg := &RetryConfig{ MaxRetries: 2, InitialDelay: 5 * time.Millisecond, MaxDelay: 50 * time.Millisecond, Multiplier: 1.0, OnRetry: func(attempt int, err error, delay time.Duration) { mu.Lock() callbackCount++ lastAttempt = attempt lastDelay = delay mu.Unlock() }, } ctx := context.Background() err := Retry(ctx, func(ctx context.Context, attempt int) error { return fmt.Errorf("failing attempt %d", attempt) }, cfg) if err == nil { t.Fatal("expected an error") } mu.Lock() if callbackCount != 2 { t.Errorf("expected 2 OnRetry callbacks, got %d", callbackCount) } // Last callback should be for attempt 2 (the retry before final attempt) if lastAttempt != 2 { t.Errorf("expected last callback attempt=2, got %d", lastAttempt) } if lastDelay != 5*time.Millisecond { t.Errorf("expected last delay=5ms, got %v", lastDelay) } mu.Unlock() } // --------------------------------------------------------------------------- // TLS tests // --------------------------------------------------------------------------- func TestDefaultTLSConfig(t *testing.T) { t.Parallel() cfg := DefaultTLSConfig() if cfg.MinVersion != tls.VersionTLS12 { t.Errorf("expected MinVersion=TLS 1.2, got 0x%04X", cfg.MinVersion) } if cfg.MaxVersion != tls.VersionTLS13 { t.Errorf("expected MaxVersion=TLS 1.3, got 0x%04X", cfg.MaxVersion) } if cfg.InsecureSkipVerify { t.Error("expected InsecureSkipVerify=false") } if cfg.CACertFile != "" { t.Errorf("expected empty CACertFile, got %s", cfg.CACertFile) } if cfg.PinnedCertHash != "" { t.Errorf("expected empty PinnedCertHash, got %s", cfg.PinnedCertHash) } if cfg.TLSServerNameOverride != "" { t.Errorf("expected empty TLSServerNameOverride, got %s", cfg.TLSServerNameOverride) } } func TestToTLSConfig(t *testing.T) { t.Parallel() t.Run("basic conversion", func(t *testing.T) { cfg := &TLSConfig{ MinVersion: tls.VersionTLS12, MaxVersion: tls.VersionTLS13, } tlsCfg, err := cfg.ToTLSConfig() if err != nil { t.Fatalf("unexpected error: %v", err) } if tlsCfg == nil { t.Fatal("expected non-nil tls.Config") } if tlsCfg.MinVersion != tls.VersionTLS12 { t.Errorf("expected MinVersion=TLS 1.2, got 0x%04X", tlsCfg.MinVersion) } if tlsCfg.MaxVersion != tls.VersionTLS13 { t.Errorf("expected MaxVersion=TLS 1.3, got 0x%04X", tlsCfg.MaxVersion) } if tlsCfg.InsecureSkipVerify { t.Error("expected InsecureSkipVerify=false") } if tlsCfg.ServerName != "" { t.Errorf("expected empty ServerName, got %s", tlsCfg.ServerName) } }) t.Run("with InsecureSkipVerify set", func(t *testing.T) { cfg := &TLSConfig{ InsecureSkipVerify: true, } tlsCfg, err := cfg.ToTLSConfig() if err != nil { t.Fatalf("unexpected error: %v", err) } if !tlsCfg.InsecureSkipVerify { t.Error("expected InsecureSkipVerify=true") } }) t.Run("with TLSServerNameOverride set", func(t *testing.T) { cfg := &TLSConfig{ TLSServerNameOverride: "example.com", } tlsCfg, err := cfg.ToTLSConfig() if err != nil { t.Fatalf("unexpected error: %v", err) } if tlsCfg.ServerName != "example.com" { t.Errorf("expected ServerName=example.com, got %s", tlsCfg.ServerName) } }) t.Run("with empty TLSServerNameOverride leaves ServerName empty", func(t *testing.T) { // Sanity check: an empty override must not be propagated into the // tls.Config, otherwise Go's stdlib would treat it as an explicit // (empty) SNI value and certificate verification would fail. cfg := &TLSConfig{} tlsCfg, err := cfg.ToTLSConfig() if err != nil { t.Fatalf("unexpected error: %v", err) } if tlsCfg.ServerName != "" { t.Errorf("expected empty ServerName, got %s", tlsCfg.ServerName) } }) t.Run("with custom TLS versions", func(t *testing.T) { cfg := &TLSConfig{ MinVersion: tls.VersionTLS10, MaxVersion: tls.VersionTLS12, } tlsCfg, err := cfg.ToTLSConfig() if err != nil { t.Fatalf("unexpected error: %v", err) } if tlsCfg.MinVersion != tls.VersionTLS10 { t.Errorf("expected MinVersion=TLS 1.0, got 0x%04X", tlsCfg.MinVersion) } if tlsCfg.MaxVersion != tls.VersionTLS12 { t.Errorf("expected MaxVersion=TLS 1.2, got 0x%04X", tlsCfg.MaxVersion) } }) t.Run("non-existent CA cert file returns error", func(t *testing.T) { cfg := &TLSConfig{ CACertFile: "/nonexistent/path/ca.pem", } _, err := cfg.ToTLSConfig() if err == nil { t.Error("expected error for non-existent CA cert file") } }) t.Run("with PinnedCertHash returns config without error", func(t *testing.T) { // PinnedCertHash is not yet implemented, but it should not cause an error cfg := &TLSConfig{ PinnedCertHash: "sha256-AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA=", } tlsCfg, err := cfg.ToTLSConfig() if err != nil { t.Fatalf("unexpected error: %v", err) } if tlsCfg == nil { t.Fatal("expected non-nil tls.Config") } }) t.Run("default TLSConfig via ToTLSConfig", func(t *testing.T) { cfg := DefaultTLSConfig() tlsCfg, err := cfg.ToTLSConfig() if err != nil { t.Fatalf("unexpected error: %v", err) } if tlsCfg.MinVersion != tls.VersionTLS12 { t.Errorf("expected MinVersion=TLS 1.2, got 0x%04X", tlsCfg.MinVersion) } if tlsCfg.MaxVersion != tls.VersionTLS13 { t.Errorf("expected MaxVersion=TLS 1.3, got 0x%04X", tlsCfg.MaxVersion) } }) } func TestGetTLSVersionName(t *testing.T) { t.Parallel() tests := []struct { version uint16 want string }{ {tls.VersionTLS10, "TLS 1.0"}, {tls.VersionTLS11, "TLS 1.1"}, {tls.VersionTLS12, "TLS 1.2"}, {tls.VersionTLS13, "TLS 1.3"}, {0x0000, "TLS 0x0000"}, {0x0300, "TLS 0x0300"}, {0xFFFF, "TLS 0xFFFF"}, } for _, tt := range tests { t.Run(fmt.Sprintf("version_0x%04X", tt.version), func(t *testing.T) { got := GetTLSVersionName(tt.version) if got != tt.want { t.Errorf("expected %q for version 0x%04X, got %q", tt.want, tt.version, got) } }) } } // --------------------------------------------------------------------------- // Token bucket tests // --------------------------------------------------------------------------- func TestNewTokenBucket(t *testing.T) { t.Parallel() t.Run("valid parameters", func(t *testing.T) { tb := NewTokenBucket(1024, 2048) if tb == nil { t.Fatal("expected non-nil token bucket") } // Check initial state via Allow and subsequent behavior if tb.capacity != 2048 { t.Errorf("expected capacity=2048, got %d", tb.capacity) } if tb.tokens != float64(2048) { t.Errorf("expected tokens=2048, got %f", tb.tokens) } if tb.refillRate != float64(1024) { t.Errorf("expected refillRate=1024, got %f", tb.refillRate) } }) t.Run("equal rate and burst", func(t *testing.T) { tb := NewTokenBucket(5000, 5000) if tb == nil { t.Fatal("expected non-nil token bucket") } if tb.capacity != 5000 { t.Errorf("expected capacity=5000, got %d", tb.capacity) } }) } func TestNewTokenBucketNil(t *testing.T) { t.Parallel() t.Run("zero rate returns nil", func(t *testing.T) { tb := NewTokenBucket(0, 1024) if tb != nil { t.Error("expected nil for zero rate") } }) t.Run("negative rate returns nil", func(t *testing.T) { tb := NewTokenBucket(-100, 1024) if tb != nil { t.Error("expected nil for negative rate") } }) } func TestTokenBucketAllow(t *testing.T) { t.Parallel() t.Run("nil bucket Allow is no-op", func(t *testing.T) { // Should not panic var tb *TokenBucket tb.Allow(100) // No assertion needed - if it didn't panic, we're good }) t.Run("allows up to capacity immediately", func(t *testing.T) { tb := NewTokenBucket(1000, 500) if tb == nil { t.Fatal("expected non-nil bucket") } // Should allow up to burst size without blocking start := time.Now() tb.Allow(300) tb.Allow(200) elapsed := time.Since(start) if elapsed > 100*time.Millisecond { t.Errorf("expected near-instant Allow for <= capacity, took %v", elapsed) } }) t.Run("blocks when exceeding rate", func(t *testing.T) { // Rate: 100 bytes/sec, burst: 100 bytes tb := NewTokenBucket(100, 100) if tb == nil { t.Fatal("expected non-nil bucket") } // Consume all tokens tb.Allow(100) // Next Allow should block for some time start := time.Now() tb.Allow(50) // needs ~500ms to refill at 100 bytes/sec elapsed := time.Since(start) if elapsed < 200*time.Millisecond { t.Errorf("expected Allow(50) to block for ~500ms, took %v", elapsed) } }) t.Run("allows after waiting for refill", func(t *testing.T) { tb := NewTokenBucket(1000, 500) if tb == nil { t.Fatal("expected non-nil bucket") } // Consume all tokens tb.Allow(500) // Wait for refill time.Sleep(200 * time.Millisecond) // ~200 tokens refilled start := time.Now() tb.Allow(150) // should have enough tokens elapsed := time.Since(start) if elapsed > 50*time.Millisecond { t.Errorf("expected Allow(150) to be near-instant after waiting, took %v", elapsed) } }) } func TestTokenBucketSmallBurst(t *testing.T) { t.Parallel() t.Run("burst of 1 works with rate limiting", func(t *testing.T) { // Rate: 1000 bytes/sec, burst: 1 byte — each byte needs ~1ms refill tb := NewTokenBucket(1000, 1) if tb == nil { t.Fatal("expected non-nil bucket") } if tb.capacity != 1 { t.Errorf("expected capacity=1, got %d", tb.capacity) } // First byte should be allowed immediately (burst capacity) start := time.Now() tb.Allow(1) elapsed := time.Since(start) if elapsed > 50*time.Millisecond { t.Errorf("expected Allow(1) to be instant, took %v", elapsed) } // Second byte should block for ~1ms at 1000 bytes/sec start = time.Now() tb.Allow(1) elapsed = time.Since(start) if elapsed < 500*time.Microsecond { t.Errorf("expected Allow(1) to block briefly for refill, took %v", elapsed) } }) t.Run("zero burst with valid rate returns non-nil but Allow blocks indefinitely", func(t *testing.T) { // capacity=0 means tokens are always capped to 0, so no token can ever be consumed. // This is an inherent limitation — the test verifies the bucket is created without // panic and that Allow(1) would block (but we don't actually call Allow to avoid timeout). tb := NewTokenBucket(1000, 0) if tb == nil { t.Fatal("expected non-nil bucket") } if tb.capacity != 0 { t.Errorf("expected capacity=0, got %d", tb.capacity) } // Note: Allow(1) with capacity=0 would loop forever because tokens // are always capped back to 0 after refill. }) } func TestWrapReader(t *testing.T) { t.Parallel() t.Run("nil bucket returns original reader", func(t *testing.T) { var tb *TokenBucket r := strings.NewReader("hello") wrapped := tb.WrapReader(r) if wrapped != r { t.Error("expected original reader when bucket is nil") } }) t.Run("wraps reader with rate limiter", func(t *testing.T) { tb := NewTokenBucket(10000, 10000) if tb == nil { t.Fatal("expected non-nil bucket") } r := strings.NewReader("hello world") wrapped := tb.WrapReader(r) if wrapped == nil { t.Fatal("expected non-nil wrapped reader") } if _, ok := wrapped.(*limitedReader); !ok { t.Error("expected *limitedReader type") } }) } func TestLimitedReader(t *testing.T) { t.Parallel() t.Run("reads all data through limiter", func(t *testing.T) { tb := NewTokenBucket(100000, 100000) if tb == nil { t.Fatal("expected non-nil bucket") } input := []byte("the quick brown fox jumps over the lazy dog") r := tb.WrapReader(bytes.NewReader(input)) output, err := io.ReadAll(r) if err != nil { t.Fatalf("unexpected error: %v", err) } if !bytes.Equal(output, input) { t.Errorf("output doesn't match input: got %q, want %q", string(output), string(input)) } }) t.Run("reads with small buffer", func(t *testing.T) { tb := NewTokenBucket(50000, 50000) if tb == nil { t.Fatal("expected non-nil bucket") } input := []byte("hello world") r := tb.WrapReader(bytes.NewReader(input)) // Read byte by byte var output []byte buf := make([]byte, 1) for { n, err := r.Read(buf) if n > 0 { output = append(output, buf[:n]...) } if err == io.EOF { break } if err != nil { t.Fatalf("unexpected error: %v", err) } } if !bytes.Equal(output, input) { t.Errorf("output doesn't match input: got %q, want %q", string(output), string(input)) } }) t.Run("empty reader returns immediately", func(t *testing.T) { tb := NewTokenBucket(1000, 1000) if tb == nil { t.Fatal("expected non-nil bucket") } r := tb.WrapReader(bytes.NewReader(nil)) n, err := r.Read(make([]byte, 10)) if n != 0 || err != io.EOF { t.Errorf("expected (0, EOF), got (%d, %v)", n, err) } }) } // --------------------------------------------------------------------------- // Integration test: retry with OnRetry counts // --------------------------------------------------------------------------- func TestRetryExponentialBackoff(t *testing.T) { t.Parallel() var delays []time.Duration cfg := &RetryConfig{ MaxRetries: 3, InitialDelay: 10 * time.Millisecond, MaxDelay: 100 * time.Millisecond, Multiplier: 3.0, // aggressive multiplier for testing OnRetry: func(attempt int, err error, delay time.Duration) { delays = append(delays, delay) }, } ctx := context.Background() Retry(ctx, func(ctx context.Context, attempt int) error { return fmt.Errorf("fail") }, cfg) // We should have had 3 retries: delays: 10ms, 30ms, 90ms if len(delays) != 3 { t.Fatalf("expected 3 delays, got %d: %v", len(delays), delays) } if delays[0] != 10*time.Millisecond { t.Errorf("expected first delay=10ms, got %v", delays[0]) } if delays[1] != 30*time.Millisecond { t.Errorf("expected second delay=30ms, got %v", delays[1]) } if delays[2] != 90*time.Millisecond { t.Errorf("expected third delay=90ms, got %v", delays[2]) } } func TestRetryExponentialBackoffMaxDelay(t *testing.T) { t.Parallel() var delays []time.Duration cfg := &RetryConfig{ MaxRetries: 4, InitialDelay: 10 * time.Millisecond, MaxDelay: 25 * time.Millisecond, Multiplier: 3.0, OnRetry: func(attempt int, err error, delay time.Duration) { delays = append(delays, delay) }, } ctx := context.Background() Retry(ctx, func(ctx context.Context, attempt int) error { return fmt.Errorf("fail") }, cfg) // delays: 10ms, 30ms (but capped at 25ms), 75ms (capped at 25ms), ... if len(delays) != 4 { t.Fatalf("expected 4 delays, got %d: %v", len(delays), delays) } if delays[1] != 25*time.Millisecond { t.Errorf("expected second delay capped at 25ms, got %v", delays[1]) } if delays[2] != 25*time.Millisecond { t.Errorf("expected third delay capped at 25ms, got %v", delays[2]) } if delays[3] != 25*time.Millisecond { t.Errorf("expected fourth delay capped at 25ms, got %v", delays[3]) } } // --------------------------------------------------------------------------- // Dialer edge cases (unit-level, no network required) // --------------------------------------------------------------------------- func TestDialerLookupIPAddr(t *testing.T) { t.Parallel() d := NewDialer(DefaultDialerConfig()) if d == nil { t.Fatal("expected non-nil dialer") } // LookupIPAddr sorts IPv6 first - we can't easily test with real DNS // without network, but we can verify the method exists and returns // appropriate error for invalid hostname. ctx, cancel := context.WithTimeout(context.Background(), 100*time.Millisecond) defer cancel() _, err := d.LookupIPAddr(ctx, "invalid-host-name-that-does-not-exist.example.com") // Should get some kind of DNS error if err == nil { t.Log("DNS lookup succeeded (unexpected on offline test)") } } func TestDialerDialContextInvalidAddress(t *testing.T) { t.Parallel() d := NewDialer(DefaultDialerConfig()) if d == nil { t.Fatal("expected non-nil dialer") } ctx := context.Background() _, err := d.DialContext(ctx, "tcp", "invalid-address") if err == nil { t.Error("expected error for invalid address") } if !strings.Contains(err.Error(), "invalid address") { t.Errorf("expected 'invalid address' in error, got %v", err) } } func TestDialerDialContextEmptyHost(t *testing.T) { t.Parallel() d := NewDialer(DefaultDialerConfig()) if d == nil { t.Fatal("expected non-nil dialer") } ctx := context.Background() // ":80" will fail SplitHostPort with "missing host" or similar _, err := d.DialContext(ctx, "tcp", ":80") if err == nil { t.Error("expected error for empty host with port") } } func TestContainsIgnoreCase(t *testing.T) { t.Parallel() tests := []struct { s, substr string want bool }{ {"Hello World", "world", true}, {"Hello World", "hello", true}, {"Hello World", "WORLD", true}, {"Hello World", "xyz", false}, {"", "", true}, {"abc", "", true}, {"", "abc", false}, } for _, tt := range tests { t.Run("", func(t *testing.T) { got := containsIgnoreCase(tt.s, tt.substr) if got != tt.want { t.Errorf("containsIgnoreCase(%q, %q) = %v, want %v", tt.s, tt.substr, got, tt.want) } }) } }