package proxy import ( "bufio" "context" "io" "net" "net/http" "testing" "time" ) // startStubConnectProxy runs a minimal HTTP CONNECT proxy for the // duration of the test: it accepts one CONNECT request, replies with the // given status, and if status is 200 splices the tunnel through to a // real dial of the requested host. Returns the proxy's address. func startStubConnectProxy(t *testing.T, status int) string { t.Helper() ln, err := net.Listen("tcp", "127.0.0.1:0") if err != nil { t.Fatalf("listen: %v", err) } t.Cleanup(func() { ln.Close() }) go func() { c, err := ln.Accept() if err != nil { return } defer c.Close() br := bufio.NewReader(c) req, err := http.ReadRequest(br) if err != nil { return } if req.Method != http.MethodConnect { c.Write([]byte("HTTP/1.1 405 Method Not Allowed\r\n\r\n")) return } if status != http.StatusOK { c.Write([]byte("HTTP/1.1 403 Forbidden\r\n\r\n")) return } target, err := net.Dial("tcp", req.Host) if err != nil { c.Write([]byte("HTTP/1.1 502 Bad Gateway\r\n\r\n")) return } defer target.Close() c.Write([]byte("HTTP/1.1 200 Connection Established\r\n\r\n")) done := make(chan struct{}, 2) go func() { io.Copy(target, br); done <- struct{}{} }() go func() { io.Copy(c, target); done <- struct{}{} }() <-done }() return ln.Addr().String() } // startEchoServer runs a TCP server that echoes back whatever it reads, // standing in for "the origin" on the far side of a CONNECT tunnel. func startEchoServer(t *testing.T) string { t.Helper() ln, err := net.Listen("tcp", "127.0.0.1:0") if err != nil { t.Fatalf("listen: %v", err) } t.Cleanup(func() { ln.Close() }) go func() { for { c, err := ln.Accept() if err != nil { return } go func(c net.Conn) { defer c.Close() io.Copy(c, c) }(c) } }() return ln.Addr().String() } func TestDialViaProxyDirect(t *testing.T) { echoAddr := startEchoServer(t) ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second) defer cancel() conn, err := dialViaProxy(ctx, echoAddr, "") if err != nil { t.Fatalf("dialViaProxy direct: %v", err) } defer conn.Close() if _, err := conn.Write([]byte("hello")); err != nil { t.Fatalf("write: %v", err) } buf := make([]byte, 5) if _, err := io.ReadFull(conn, buf); err != nil { t.Fatalf("read: %v", err) } if string(buf) != "hello" { t.Errorf("got %q, want %q", buf, "hello") } } func TestDialViaProxyTunneled(t *testing.T) { echoAddr := startEchoServer(t) proxyAddr := startStubConnectProxy(t, http.StatusOK) ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second) defer cancel() conn, err := dialViaProxy(ctx, echoAddr, proxyAddr) if err != nil { t.Fatalf("dialViaProxy via proxy: %v", err) } defer conn.Close() if _, err := conn.Write([]byte("world")); err != nil { t.Fatalf("write: %v", err) } buf := make([]byte, 5) if _, err := io.ReadFull(conn, buf); err != nil { t.Fatalf("read: %v", err) } if string(buf) != "world" { t.Errorf("got %q, want %q", buf, "world") } } func TestDialViaProxyRejected(t *testing.T) { proxyAddr := startStubConnectProxy(t, http.StatusForbidden) ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second) defer cancel() _, err := dialViaProxy(ctx, "example.invalid:443", proxyAddr) if err == nil { t.Fatal("expected an error when the upstream proxy refuses CONNECT, got nil") } }