Something went wrong. Try again.
Monorepo for Tangled tangled.org
Something went wrong. Try again.
Go
at sl/gitmirror
123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271package netutil
import ( "context" "errors" "fmt" "io" "net" "net/http" "net/netip" "sync" "time"
"github.com/bluesky-social/indigo/util/ssrf")
var relayDrainGrace = 30 * time.Second
type IPResolver interface { LookupIP(ctx context.Context, host string) ([]netip.Addr, error)}
type defaultResolver struct{}
func (defaultResolver) LookupIP(ctx context.Context, host string) ([]netip.Addr, error) { ips, err := net.DefaultResolver.LookupIP(ctx, "ip", host) if err != nil { return nil, err } addrs := make([]netip.Addr, 0, len(ips)) for _, ip := range ips { if addr, ok := netip.AddrFromSlice(ip); ok { addrs = append(addrs, addr.Unmap()) } } return addrs, nil}
type OutboundDialer interface { DialContext(ctx context.Context, network, address string) (net.Conn, error)}
type defaultDialer struct { dialer net.Dialer}
func (d *defaultDialer) DialContext(ctx context.Context, network, address string) (net.Conn, error) { return d.dialer.DialContext(ctx, network, address)}
type SafeConnectProxy struct { resolver IPResolver dialer OutboundDialer server *http.Server listener net.Listener addr string mu sync.Mutex closed bool active map[net.Conn]struct{}}
func NewSafeConnectProxy(resolver IPResolver, dialer OutboundDialer) *SafeConnectProxy { if resolver == nil { resolver = defaultResolver{} } if dialer == nil { dialer = SSRFDialer(false) } p := &SafeConnectProxy{ resolver: resolver, dialer: dialer, active: make(map[net.Conn]struct{}), } return p}
func (p *SafeConnectProxy) Start() (string, error) { l, err := net.Listen("tcp", "127.0.0.1:0") if err != nil { return "", fmt.Errorf("starting safe connect proxy listener: %w", err) } p.listener = l p.addr = l.Addr().String()
p.server = &http.Server{ Handler: http.HandlerFunc(p.handleHTTP), }
go func() { _ = p.server.Serve(l) }()
return p.addr, nil}
func (p *SafeConnectProxy) Addr() string { return p.addr}
func (p *SafeConnectProxy) Close() error { p.mu.Lock() if p.closed { p.mu.Unlock() return nil } p.closed = true // http.Server drops hijacked conns, so close them explicitly on shutdown for conn := range p.active { _ = conn.Close() } server := p.server p.mu.Unlock()
if server != nil { return server.Close() } return nil}
func (p *SafeConnectProxy) track(conns ...net.Conn) { p.mu.Lock() defer p.mu.Unlock() if p.closed { for _, conn := range conns { _ = conn.Close() } return } for _, conn := range conns { p.active[conn] = struct{}{} }}
func (p *SafeConnectProxy) untrack(conns ...net.Conn) { p.mu.Lock() defer p.mu.Unlock() for _, conn := range conns { delete(p.active, conn) }}
func (p *SafeConnectProxy) handleHTTP(w http.ResponseWriter, r *http.Request) { if r.Method != http.MethodConnect { http.Error(w, "Method Not Allowed: only CONNECT method is permitted", http.StatusMethodNotAllowed) return }
host, port, err := net.SplitHostPort(r.Host) if err != nil { host = r.Host port = "443" }
if port != "443" { http.Error(w, "Forbidden: only port 443 is permitted", http.StatusForbidden) return }
ctx, cancel := context.WithTimeout(r.Context(), 10*time.Second) defer cancel()
addrs, err := p.resolveAndVerify(ctx, host) if err != nil { http.Error(w, fmt.Sprintf("Forbidden: SSRF safety check failed: %v", err), http.StatusForbidden) return }
targetConn, err := p.dialAny(ctx, addrs, port) if err != nil { http.Error(w, fmt.Sprintf("Bad Gateway: dial failed: %v", err), http.StatusBadGateway) return } defer targetConn.Close()
hj, ok := w.(http.Hijacker) if !ok { http.Error(w, "Internal Server Error: hijacking not supported", http.StatusInternalServerError) return }
clientConn, _, err := hj.Hijack() if err != nil { http.Error(w, fmt.Sprintf("Internal Server Error: hijacking failed: %v", err), http.StatusInternalServerError) return } defer clientConn.Close()
if _, err := clientConn.Write([]byte("HTTP/1.1 200 Connection Established\r\n\r\n")); err != nil { return }
p.relay(clientConn, targetConn)}
func (p *SafeConnectProxy) dialAny(ctx context.Context, addrs []netip.Addr, port string) (net.Conn, error) { var errs []error for _, addr := range addrs { conn, err := p.dialer.DialContext(ctx, "tcp", net.JoinHostPort(addr.String(), port)) if err == nil { return conn, nil } errs = append(errs, fmt.Errorf("%s: %w", addr, err)) } return nil, errors.Join(errs...)}
func (p *SafeConnectProxy) relay(clientConn, targetConn net.Conn) { p.track(clientConn, targetConn) defer p.untrack(clientConn, targetConn) defer targetConn.Close() defer clientConn.Close()
var wg sync.WaitGroup wg.Add(2) go func() { defer wg.Done() _, _ = io.Copy(targetConn, clientConn) halfClose(targetConn) }() go func() { defer wg.Done() _, _ = io.Copy(clientConn, targetConn) halfClose(clientConn) }()
drained := make(chan struct{}) go func() { wg.Wait() close(drained) }()
select { case <-drained: case <-time.After(relayDrainGrace): }}
func halfClose(conn net.Conn) { if cw, ok := conn.(interface{ CloseWrite() error }); ok { _ = cw.CloseWrite() }}
func (p *SafeConnectProxy) resolveAndVerify(ctx context.Context, host string) ([]netip.Addr, error) { if addr, err := netip.ParseAddr(host); err == nil { addr = addr.Unmap() if !ssrf.IsPublicIPAddress(net.IP(addr.Unmap().AsSlice())) { return nil, fmt.Errorf("direct IP %s is not a safe public address", addr) } return []netip.Addr{addr}, nil }
resolved, err := p.resolver.LookupIP(ctx, host) if err != nil { return nil, fmt.Errorf("lookup host %q: %w", host, err) } if len(resolved) == 0 { return nil, fmt.Errorf("host %q resolved to no IP addresses", host) }
addrs := make([]netip.Addr, 0, len(resolved)) for _, addr := range resolved { addr = addr.Unmap() if !ssrf.IsPublicIPAddress(net.IP(addr.Unmap().AsSlice())) { return nil, fmt.Errorf("host %q resolved to unsafe IP %s (all answers must be safe)", host, addr) } addrs = append(addrs, addr) } return addrs, nil}