package 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 }