diff --git a/forwardproxy.go b/forwardproxy.go index d4a1948..4a16220 100644 --- a/forwardproxy.go +++ b/forwardproxy.go @@ -18,6 +18,7 @@ package forwardproxy import ( "bufio" + "context" "crypto/subtle" "errors" "fmt" @@ -29,6 +30,7 @@ import ( "sync" "time" + "github.com/caddyserver/forwardproxy/httpclient" "github.com/mholt/caddy/caddyhttp/httpserver" ) @@ -52,9 +54,10 @@ type ForwardProxy struct { dialTimeout time.Duration // for initial tcp connection responseTimeout *time.Duration // for getting response (affects GET requests only) - // overridden dial allows to redirect requests to upstream proxy - dial func(network, address string) (net.Conn, error) - upstream string // address of upstream proxy + // overridden dial and dialContext allow to redirect requests to upstream proxy + dial func(network, address string) (net.Conn, error) + dialContext func(ctx context.Context, network, address string) (net.Conn, error) + upstream string // address of upstream proxy aclRules []aclRule whitelistedPorts []int @@ -258,7 +261,20 @@ func (fp *ForwardProxy) dialRequestedAddress(r *http.Request) (net.Conn, error, } if fp.upstream != "" { // if upstreaming -- do not resolve locally nor check acl - conn, err = fp.dial("tcp", hostPort) + if fp.dialContext != nil && !fp.hideIP { + ctxHeader := make(http.Header) + for k, v := range r.Header { + if kL := strings.ToLower(k); kL == "forwarded" || kL == "x-forwarded-for" { + ctxHeader[k] = v + } + } + ctxHeader.Add("Forwarded", "for=\""+r.RemoteAddr+"\"") + ctx := context.WithValue(context.Background(), httpclient.ContextKeyHeader{}, ctxHeader) + + conn, err = fp.dialContext(ctx, "tcp", hostPort) + } else { + conn, err = fp.dial("tcp", hostPort) + } return conn, err, false } diff --git a/httpclient/httpclient.go b/httpclient/httpclient.go new file mode 100644 index 0000000..9fda3fd --- /dev/null +++ b/httpclient/httpclient.go @@ -0,0 +1,272 @@ +// Copyright 2018 Google Inc. +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +// httpclient is used by the upstreaming forwardproxy to establish connections to http(s) upstreams. +// it implements x/net/proxy.Dialer interface +package httpclient + +import ( + "bufio" + "context" + "crypto/tls" + "encoding/base64" + "errors" + "io" + "io/ioutil" + "net" + "net/http" + "net/url" + "sync" + + "golang.org/x/net/http2" +) + +// HTTPConnectDialer allows to configure one-time use HTTP CONNECT client +type HTTPConnectDialer struct { + ProxyUrl url.URL + DefaultHeader http.Header + + // TODO: If spkiFp is set, use it as SPKI fingerprint to confirm identity of the + // proxy, instead of relying on standard PKI CA roots + SpkiFP []byte + + Dialer net.Dialer // overridden dialer allow to control establishment of TCP connection + + // overridden DialTLS allows user to control establishment of TLS connection, + // given TCP connection. returns established TLS connection, ALPN and error + DialTLS func(net.Conn) (net.Conn, string, error) + + EnableH2ConnReuse bool + cacheH2Mu sync.Mutex + cachedH2ClientConn *http2.ClientConn + cachedH2RawConn net.Conn +} + +// NewHTTPClient creates a client to issue CONNECT requests and tunnel traffic via HTTPS proxy. +// proxyUrlStr must provide Scheme and Host, may provide credentials and port. +// Example: https://username:password@golang.org:443 +func NewHTTPConnectDialer(proxyUrlStr string) (*HTTPConnectDialer, error) { + proxyUrl, err := url.Parse(proxyUrlStr) + if err != nil { + return nil, err + } + + if proxyUrl.Host == "" { + return nil, errors.New("misparsed `url=" + proxyUrlStr + + "`, make sure to specify full url like https://username:password@hostname.com:443/") + } + + switch proxyUrl.Scheme { + case "http": + if proxyUrl.Port() == "" { + proxyUrl.Host = net.JoinHostPort(proxyUrl.Host, "80") + } + case "https": + if proxyUrl.Port() == "" { + proxyUrl.Host = net.JoinHostPort(proxyUrl.Host, "443") + } + case "": + return nil, errors.New("specify scheme explicitly (https://)") + default: + return nil, errors.New("scheme " + proxyUrl.Scheme + " is not supported") + } + + client := &HTTPConnectDialer{ + ProxyUrl: *proxyUrl, + DefaultHeader: make(http.Header), + SpkiFP: nil, + EnableH2ConnReuse: true, + } + + if proxyUrl.User != nil { + if proxyUrl.User.Username() != "" { + password, _ := proxyUrl.User.Password() + client.DefaultHeader.Set("Proxy-Authorization", "Basic "+ + base64.StdEncoding.EncodeToString([]byte(proxyUrl.User.Username()+":"+password))) + } + } + return client, nil +} + +func (c *HTTPConnectDialer) Dial(network, address string) (net.Conn, error) { + return c.DialContext(context.Background(), network, address) +} + +// Users of context.WithValue should define their own types for keys +type ContextKeyHeader struct{} + +// ctx.Value will be inspected for optional ContextKeyHeader{} key, with `http.Header` value, +// which will be added to outgoing request headers, overriding any colliding c.DefaultHeader +func (c *HTTPConnectDialer) DialContext(ctx context.Context, network, address string) (net.Conn, error) { + req := (&http.Request{ + Method: "CONNECT", + URL: &url.URL{Host: address}, + Header: make(http.Header), + Host: address, + }).WithContext(ctx) + for k, v := range c.DefaultHeader { + req.Header[k] = v + } + if ctxHeader, ctxHasHeader := ctx.Value(ContextKeyHeader{}).(http.Header); ctxHasHeader { + for k, v := range ctxHeader { + req.Header[k] = v + } + } + + connectHttp2 := func(rawConn net.Conn, h2clientConn *http2.ClientConn) (net.Conn, error) { + req.Proto = "HTTP/2.0" + req.ProtoMajor = 2 + req.ProtoMinor = 0 + pr, pw := io.Pipe() + req.Body = ioutil.NopCloser(pr) + + resp, err := h2clientConn.RoundTrip(req) + if err != nil { + rawConn.Close() + return nil, err + } + + if resp.StatusCode != http.StatusOK { + rawConn.Close() + return nil, errors.New("Proxy responded with non 200 code: " + resp.Status) + } + return NewHttp2Conn(rawConn, pw, resp.Body), nil + } + + connectHttp1 := func(rawConn net.Conn) (net.Conn, error) { + req.Proto = "HTTP/1.1" + req.ProtoMajor = 1 + req.ProtoMinor = 1 + + err := req.Write(rawConn) + if err != nil { + rawConn.Close() + return nil, err + } + + resp, err := http.ReadResponse(bufio.NewReader(rawConn), req) + if err != nil { + rawConn.Close() + return nil, err + } + + if resp.StatusCode != http.StatusOK { + rawConn.Close() + return nil, errors.New("Proxy responded with non 200 code: " + resp.Status) + } + return rawConn, nil + } + + if c.EnableH2ConnReuse { + c.cacheH2Mu.Lock() + if c.cachedH2ClientConn != nil && c.cachedH2RawConn != nil { + if c.cachedH2ClientConn.CanTakeNewRequest() { + proxyConn, err := connectHttp2(c.cachedH2RawConn, c.cachedH2ClientConn) + if err == nil { + c.cacheH2Mu.Unlock() + return proxyConn, err + } + // else: carry on and try again + } + } + c.cacheH2Mu.Unlock() + } + + rawConn, err := c.Dialer.DialContext(ctx, network, c.ProxyUrl.Host) + if err != nil { + return nil, err + } + + negotiatedProtocol := "" + switch c.ProxyUrl.Scheme { + case "http": + case "https": + if c.DialTLS != nil { + rawConn, negotiatedProtocol, err = c.DialTLS(rawConn) + if err != nil { + return nil, err + } + break + } + tlsConf := tls.Config{ + NextProtos: []string{"h2", "http/1.1"}, + ServerName: c.ProxyUrl.Hostname(), + } + + tlsConn := tls.Client(rawConn, &tlsConf) + err = tlsConn.Handshake() + if err != nil { + return nil, err + } + negotiatedProtocol = tlsConn.ConnectionState().NegotiatedProtocol + rawConn = tlsConn + default: + return nil, errors.New("scheme " + c.ProxyUrl.Scheme + " is not supported") + } + + switch negotiatedProtocol { + case "": + fallthrough + case "http/1.1": + return connectHttp1(rawConn) + case "h2": + t := http2.Transport{} + h2clientConn, err := t.NewClientConn(rawConn) + if err != nil { + rawConn.Close() + return nil, err + } + + proxyConn, err := connectHttp2(rawConn, h2clientConn) + if err != nil { + rawConn.Close() + return nil, err + } + if c.EnableH2ConnReuse { + c.cacheH2Mu.Lock() + c.cachedH2ClientConn = h2clientConn + c.cachedH2RawConn = rawConn + c.cacheH2Mu.Unlock() + } + return proxyConn, err + default: + rawConn.Close() + return nil, errors.New("negotiated unsupported application layer protocol: " + + negotiatedProtocol) + } +} + +func NewHttp2Conn(c net.Conn, pipedReqBody *io.PipeWriter, respBody io.ReadCloser) net.Conn { + return &http2Conn{Conn: c, in: pipedReqBody, out: respBody} +} + +type http2Conn struct { + net.Conn + in *io.PipeWriter + out io.ReadCloser +} + +func (h *http2Conn) Read(p []byte) (n int, err error) { + return h.out.Read(p) +} + +func (h *http2Conn) Write(p []byte) (n int, err error) { + return h.in.Write(p) +} + +func (h *http2Conn) Close() error { + h.out.Close() + h.in.Close() + return h.Conn.Close() +} diff --git a/setup.go b/setup.go index c88be97..0610ea6 100644 --- a/setup.go +++ b/setup.go @@ -15,25 +15,24 @@ package forwardproxy import ( + "bufio" + "context" + "crypto/tls" "encoding/base64" "errors" "log" "net" - "net/http" "net/url" + "os" "strconv" "strings" "sync" "time" - "bufio" - "crypto/tls" - "fmt" - + "github.com/caddyserver/forwardproxy/httpclient" "github.com/mholt/caddy" "github.com/mholt/caddy/caddyhttp/httpserver" "golang.org/x/net/proxy" - "os" ) func setup(c *caddy.Controller) error { @@ -272,21 +271,43 @@ func setup(c *caddy.Controller) error { return errors.New("insecure schemes are only allowed to localhost upstreams") } - // TODO: remove homebrewed Dialer when https://go-review.googlesource.com/c/net/+/111135 gets merged - proxy.RegisterDialerType("https", func(u *url.URL, _ proxy.Dialer) (proxy.Dialer, error) { + registerHTTPDialer := func(u *url.URL, _ proxy.Dialer) (proxy.Dialer, error) { // CONNECT request is proxied as-is, so we don't care about target url, but it could be // useful in future to implement policies of choosing between multiple upstream servers. // Given dialer is not used, since it's the same dialer provided by us. - return NewHTTPDialer(dialer, true, upstreamURL), nil - }) - proxy.RegisterDialerType("http", func(u *url.URL, _ proxy.Dialer) (proxy.Dialer, error) { - return NewHTTPDialer(dialer, false, upstreamURL), nil - }) + d, err := httpclient.NewHTTPConnectDialer(upstreamURL.String()) + if err != nil { + return nil, err + } + d.Dialer = *dialer + if isLocalhost(upstreamURL.Hostname()) && upstreamURL.Scheme == "https" { + // disabling verification helps with testing the package and setups + // either way, it's impossible to have a legit TLS certificate for "127.0.0.1" + log.Println("Localhost upstream detected, disabling verification of TLS certificate") + d.DialTLS = func(conn net.Conn) (net.Conn, string, error) { + cl := tls.Client(conn, &tls.Config{InsecureSkipVerify: true}) + err := cl.Handshake() + if err != nil { + return nil, "", err + } + return cl, cl.ConnectionState().NegotiatedProtocol, err + } + } + return d, nil + } + proxy.RegisterDialerType("https", registerHTTPDialer) + proxy.RegisterDialerType("http", registerHTTPDialer) + newDialer, err := proxy.FromURL(upstreamURL, dialer) if err != nil { return errors.New("failed to create proxy to upstream: " + err.Error()) } fp.dial = newDialer.Dial + if ctxDialer, ok := newDialer.(interface { + DialContext(ctx context.Context, network, address string) (net.Conn, error) + }); ok { + fp.dialContext = ctxDialer.DialContext + } } httpserver.GetConfig(c).AddMiddleware(func(next httpserver.Handler) httpserver.Handler { @@ -306,60 +327,6 @@ func init() { }) } -type HTTPDialer struct { - dialer *net.Dialer - upstreamUrl string - tlsConf *tls.Config - extraHeaders string // empty or whole lines together with \r\n\r\n -} - -func NewHTTPDialer(dialer *net.Dialer, useHTTPS bool, upstream *url.URL) *HTTPDialer { - d := &HTTPDialer{ - dialer: dialer, - upstreamUrl: upstream.Host, - tlsConf: nil, - } - if useHTTPS { - d.tlsConf = &tls.Config{ServerName: upstream.Hostname()} - if isLocalhost(upstream.Hostname()) { - log.Println("Localhost upstream detected, disabling verification of TLS certificate") - d.tlsConf.InsecureSkipVerify = true - } - } - if upstream.User != nil { - d.extraHeaders = fmt.Sprintf("Proxy-Authorization: basic %s\r\n", - base64.StdEncoding.EncodeToString([]byte(upstream.User.String()))) - } - - return d -} - -func (d *HTTPDialer) Dial(network, addr string) (net.Conn, error) { - var err error - var c net.Conn - if d.tlsConf == nil { - c, err = d.dialer.Dial(network, d.upstreamUrl) - } else { - c, err = tls.DialWithDialer(d.dialer, network, d.upstreamUrl, d.tlsConf) - } - if err != nil { - return nil, err - } - // TODO: multiplexed http/2 to upstream, also will eventually be added to x/net/proxy - _, err = fmt.Fprintf(c, "CONNECT %s HTTP/1.1\r\nHost: %s\r\n%s\r\n", addr, addr, d.extraHeaders) - if err != nil { - return nil, err - } - resp, err := http.ReadResponse(bufio.NewReader(c), nil) - if err != nil { - return nil, err - } - if resp.StatusCode != http.StatusOK { - return nil, errors.New("Upstream responded with " + resp.Status) - } - return c, nil -} - func isLocalhost(hostname string) bool { if hostname == "localhost" || hostname == "127.0.0.1" || hostname == "::1" { return true diff --git a/setup_test.go b/setup_test.go index a10f71e..cece93c 100644 --- a/setup_test.go +++ b/setup_test.go @@ -122,6 +122,8 @@ func TestSetup(t *testing.T) { testParsing([]string{"upstream https://proxy.site"}, true) testParsing([]string{"upstream https://caddyserver.com", "acl {\nallow all\n}"}, false) testParsing([]string{"upstream https://caddyserver.com", "ports 123"}, false) + testParsing([]string{"upstream https://username:password@caddyserver.com", "ports 123"}, false) + testParsing([]string{"upstream https://username:password@caddyserver.com:90", "ports 123"}, false) testParsing([]string{"acl {\nallow all\n}"}, true) testParsing([]string{"acl {\nallow localhost 128.32.22.1/32 1.1.1.1 caddyserver.com\n deny all\n}"}, true)