Developer SDK

Go examples

Install the SDK · View source on GitHub ↗

socket/proxy.go

package main

import (
	"context"
	"errors"
	"io"
	"net"
	"net/http"
	"strings"
	"time"
)

// StartProxy adapts clients that cannot accept a userspace stream. The local
// kernel listener is loopback-only; every destination dial uses device.
// It is an application example, not a Device listener/server API.
func StartProxy(parent context.Context, device Dialer) (proxyURL string, closeProxy func() error, err error) {
	listener, err := net.Listen("tcp4", "127.0.0.1:0")
	if err != nil {
		return "", nil, err
	}
	ctx, cancel := context.WithCancel(parent)
	transport := http.DefaultTransport.(*http.Transport).Clone()
	transport.Proxy = nil
	transport.DialContext = device.DialContext
	handler := http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
		if r.Method == http.MethodConnect {
			if _, _, err := net.SplitHostPort(r.Host); err != nil {
				http.Error(w, "CONNECT requires host:port", http.StatusBadRequest)
				return
			}
			dialCtx, cancelDial := context.WithTimeout(r.Context(), 30*time.Second)
			upstream, err := device.DialContext(dialCtx, "tcp", r.Host)
			cancelDial()
			if err != nil {
				http.Error(w, "destination unavailable", http.StatusBadGateway)
				return
			}
			defer upstream.Close()
			hijacker, ok := w.(http.Hijacker)
			if !ok {
				http.Error(w, "CONNECT requires HTTP/1.1", http.StatusHTTPVersionNotSupported)
				return
			}
			client, buffered, err := hijacker.Hijack()
			if err != nil {
				return
			}
			defer client.Close()
			stop := context.AfterFunc(ctx, func() { client.Close(); upstream.Close() })
			defer stop()
			if _, err = buffered.WriteString("HTTP/1.1 200 Connection Established\r\n\r\n"); err != nil {
				return
			}
			if err = buffered.Flush(); err != nil {
				return
			}
			done := make(chan struct{})
			go func() {
				defer close(done)
				_, _ = io.Copy(upstream, buffered)
				// An application closing its write half can still receive a reply.
				if half, ok := upstream.(interface{ CloseWrite() error }); ok {
					_ = half.CloseWrite()
				} else {
					_ = upstream.Close()
				}
			}()
			_, _ = io.Copy(client, upstream)
			_ = client.Close()
			_ = upstream.Close()
			<-done
			return
		}
		if r.URL.Scheme != "http" || r.URL.Host == "" || r.URL.User != nil {
			http.Error(w, "absolute HTTP proxy URL required", http.StatusBadRequest)
			return
		}
		out := r.Clone(r.Context())
		out.RequestURI = ""
		out.Header = r.Header.Clone()
		removeHopHeaders(out.Header)
		resp, err := transport.RoundTrip(out)
		if err != nil {
			http.Error(w, "destination unavailable", http.StatusBadGateway)
			return
		}
		defer resp.Body.Close()
		removeHopHeaders(resp.Header)
		for name, values := range resp.Header {
			for _, value := range values {
				w.Header().Add(name, value)
			}
		}
		w.WriteHeader(resp.StatusCode)
		_, _ = io.Copy(w, resp.Body)
	})
	server := &http.Server{Handler: handler, ReadHeaderTimeout: 10 * time.Second, BaseContext: func(net.Listener) context.Context { return ctx }}
	go func() {
		<-ctx.Done()
		_ = server.Close()
		transport.CloseIdleConnections()
	}()
	go func() {
		if err := server.Serve(listener); err != nil && !errors.Is(err, http.ErrServerClosed) {
			cancel()
		}
	}()
	return "http://" + listener.Addr().String(), func() error { cancel(); return server.Close() }, nil
}

func removeHopHeaders(header http.Header) {
	for _, value := range header.Values("Connection") {
		for _, name := range strings.Split(value, ",") {
			header.Del(strings.TrimSpace(name))
		}
	}
	for _, name := range []string{"Connection", "Proxy-Connection", "Proxy-Authenticate", "Proxy-Authorization", "Keep-Alive", "TE", "Trailer", "Transfer-Encoding", "Upgrade"} {
		header.Del(name)
	}
}