diff --git a/netproxy/dialer.go b/netproxy/dialer.go index 44faf40..a031153 100644 --- a/netproxy/dialer.go +++ b/netproxy/dialer.go @@ -36,14 +36,26 @@ const DefaultTimeout = 30 * time.Second // never resolved locally, so it works when local DNS for the target is // broken or poisoned. Callers that want TLS wrap the returned conn with // tls.Client themselves; TLS then runs end-to-end with the real server. +// +// When the proxy rejects the credentials (407 Proxy Authentication +// Required, or a CONNECT response so garbled it cannot be parsed, which is +// how some proxies' rejections arrive), the Dialer refreshes the Resolver +// (see Resolver.Refresh) and, if that produced different credentials, +// retries once on a new connection. Tunnels opened earlier are never +// touched. A Resolver with nothing to refresh gets no retry. +// +// Errors never quote the proxy's response: a 407 or other refusal is a +// *ConnectError, and a response that cannot be parsed is reported by what +// was wrong with it ("malformed HTTP status code", say) without its text. type Dialer struct { // Resolver picks the proxy per target. nil never proxies. Resolver *Resolver // Timeout bounds the whole establishment: the TCP connect to the proxy // (or to the target when dialing directly), the TLS handshake with an - // https:// proxy, and the CONNECT exchange. Zero means DefaultTimeout. - // An earlier deadline on the DialContext context wins. + // https:// proxy, and the CONNECT exchange, including a credential + // refresh and the one retry after a rejection. Zero means + // DefaultTimeout. An earlier deadline on the DialContext context wins. Timeout time.Duration // Forward dials the proxy itself, and the target when no proxy applies. @@ -97,6 +109,9 @@ func (d *Dialer) DialContext(ctx context.Context, network, addr string) (net.Con ctx, cancel := context.WithTimeout(ctx, timeout) defer cancel() + // Noted before the pick, so a timed refresh the pick itself starts + // counts as having happened after it (see Resolver.reauth). + start := resolver.attemptCount() proxyURL, err := resolver.ProxyForAddr(addr) if err != nil { return nil, err @@ -109,7 +124,88 @@ func (d *Dialer) DialContext(ctx context.Context, network, addr string) (net.Con default: return nil, fmt.Errorf("netproxy: cannot tunnel network %q through proxy %s", network, Redact(proxyURL)) } - return d.dialConnect(ctx, proxyURL, addr) + conn, err := d.dialConnect(ctx, proxyURL, addr) + if err == nil || !credentialsRejected(err) { + return conn, err + } + next, retry, refreshErr := resolver.reauth(ctx, start, proxyURL, func(st *proxyState) *url.URL { + u, _ := resolver.pickAddr(st, addr) // addr already parsed once + return u + }) + if !retry { + return nil, withRefreshError(err, refreshErr) + } + if next == nil { + // The refreshed settings no longer proxy this target. + return d.forward(ctx, network, addr) + } + return d.dialConnect(ctx, next, addr) +} + +// credentialsRejected reports whether a CONNECT failed in a way stale +// credentials explain: a 407, or a response that could not be parsed. +func credentialsRejected(err error) bool { + var ce *ConnectError + if errors.As(err, &ce) { + return ce.StatusCode == http.StatusProxyAuthRequired + } + var bad *badConnectResponse + return errors.As(err, &bad) +} + +// badConnectResponse is a CONNECT response that could not be parsed (as +// opposed to an I/O error while reading it). It keeps only what was wrong +// with the response: net/http's own message quotes the offending bytes, +// and a proxy can fill those with whatever it likes, including the +// Proxy-Authorization it was sent. +type badConnectResponse struct { + // what is a fixed description, e.g. "malformed HTTP status code". + what string +} + +func (e *badConnectResponse) Error() string { + return "read CONNECT response: " + e.what + " (response text withheld)" +} + +// responseFaults are net/http's (and net/textproto's) complaints about a +// response http.ReadResponse cannot parse, most specific first. Their +// messages go on to quote the offending text, which is why only these +// fixed descriptions are ever kept. +var responseFaults = []string{ + "malformed HTTP status code", + "malformed HTTP response", + "malformed HTTP version", + "malformed MIME header initial line", + "malformed MIME header line", + "malformed MIME header", + "invalid empty Content-Length", + "bad Content-Length", + "multiple Content-Length headers", + "too many transfer encodings", + "unsupported transfer encoding", + "invalid Trailer key", +} + +// responseFault returns which of responseFaults err reports, or "" if none. +func responseFault(err error) string { + if err == nil { + return "" + } + msg := err.Error() + for _, fault := range responseFaults { + if strings.Contains(msg, fault) { + return fault + } + } + return "" +} + +// isReadError reports whether err is a failure to read (the connection +// closing, timing out or breaking) rather than a complaint about what was +// read. Such errors carry no text from the peer. +func isReadError(err error) bool { + var ne net.Error + return errors.Is(err, io.EOF) || errors.Is(err, io.ErrUnexpectedEOF) || errors.As(err, &ne) } func (d *Dialer) forward(ctx context.Context, network, addr string) (net.Conn, error) { @@ -200,7 +296,15 @@ func (d *Dialer) connect(ctx context.Context, conn net.Conn, proxyURL *url.URL, br := bufio.NewReader(conn) resp, err := http.ReadResponse(br, &http.Request{Method: http.MethodConnect}) if err != nil { - return nil, fmt.Errorf("read CONNECT response: %w", err) + if isReadError(err) { + return nil, fmt.Errorf("read CONNECT response: %w", err) + } + // Anything else is about the bytes the proxy sent, and quotes them. + what := responseFault(err) + if what == "" { + what = "unparseable response" + } + return nil, &badConnectResponse{what: what} } // The body is never read: on success the rest of the stream belongs to // the tunnel, and on failure the conn is discarded. diff --git a/netproxy/netproxy.go b/netproxy/netproxy.go index 33ceaf6..e0d1087 100644 --- a/netproxy/netproxy.go +++ b/netproxy/netproxy.go @@ -21,12 +21,36 @@ // - Dialer is a DialContext-compatible dialer that tunnels through the // proxy the Resolver picks, or dials directly when there is none. // +// # Rotating credentials +// +// Some egress proxies (Meta Muse's, for one) rotate the credentials embedded +// in HTTPS_PROXY every few minutes. A fresh shell sees the current value; a +// long-running process keeps its launch-time copy and gets 407 Proxy +// Authentication Required on every new CONNECT, while tunnels it already +// opened stay up. A Resolver therefore re-reads its proxy settings: +// +// - on a timer (DefaultRefreshInterval, see WithRefreshInterval): the +// first lookup after the interval has passed starts a refresh in the +// background and, like every lookup, answers from the settings in hand, +// so a slow refresh never holds up a connection; +// - immediately when a proxy rejects its credentials: Dialer and +// RefreshingTransport then refresh once and retry once on a new +// connection, and only when the refresh produced different credentials. +// +// A ModeAuto Resolver re-reads the process environment. WithRefreshCommand +// adds a command whose output is the current proxy URL, for processes whose +// own environment never changes; by convention programs take that command +// from EnvRefreshCommand (PILOT_PROXY_CMD). Concurrent rejections share one +// refresh, and a failed refresh keeps the last good settings. +// // Proxy credentials (URL userinfo) are only ever written to the proxy in a // Proxy-Authorization header. They never appear in errors or in the output // of Redact / Resolver.String, which is what callers must use for logging. +// A refresh command's output is treated the same way and is never logged. package netproxy import ( + "context" "errors" "fmt" "net" @@ -34,6 +58,9 @@ import ( "net/url" "os" "strings" + "sync" + "sync/atomic" + "time" ) // Mode names. Parse accepts ModeAuto and ModeOff (case-insensitive) in @@ -49,11 +76,51 @@ const ( // Resolver decides which proxy, if any, an outbound connection to a given // target should go through. A nil *Resolver is valid and never proxies. -// Resolvers are immutable and safe for concurrent use. +// Resolvers are safe for concurrent use. +// +// Off and Explicit Resolvers without a refresh source never change. A +// ModeAuto Resolver, and any Resolver with a refresh source +// (WithRefreshCommand, WithRefreshFunc), re-reads its settings over time; +// see the package documentation under "Rotating credentials". type Resolver struct { mode string - // fixed is the proxy for every target in ModeExplicit. + // getenv is the environment a ModeAuto Resolver reads (os.Getenv + // outside tests). + getenv func(string) string + // source, when set, supplies the current proxy URL; command reports + // that it runs a shell command (for String). + source func(context.Context) (string, error) + command bool + // interval is the refresh TTL; negative disables timed refreshes. + interval time.Duration + // onError receives the first error of each run of failed refreshes. + onError func(error) + + // state holds the current settings. Refreshes replace it whole, so a + // lookup always sees one consistent set. + state atomic.Pointer[proxyState] + + mu sync.Mutex // guards the fields below + // inflight is the refresh in progress, which every caller joins. + inflight *refreshCall + // attempts counts refreshes started. A caller notes it before picking + // a proxy; if it has moved when the proxy then rejects the + // credentials, a refresh already ran after the pick and running + // another one right away cannot learn anything newer. + attempts uint64 + // lastRun is when the last refresh finished (or the Resolver was built). + lastRun time.Time + // failing is set while refreshes keep failing, so onError hears about + // a run of failures once. + failing bool +} + +// proxyState is one reading of a Resolver's settings. +type proxyState struct { + // fixed is the proxy for every target in ModeExplicit, and in ModeAuto + // the proxy URL from the refresh source, which then replaces the + // environment's proxy URLs (NO_PROXY still applies). fixed *url.URL // ModeAuto: proxy for TLS / opaque TCP targets (HTTPS_PROXY, https_proxy, @@ -69,6 +136,8 @@ type Resolver struct { warnings []error } +var emptyState proxyState + // EnvError reports a proxy environment variable whose value cannot be used. // Its message names the variable and never includes the value's // credentials. @@ -92,7 +161,14 @@ func Off() *Resolver { return &Resolver{mode: ModeOff} } // scheme is taken as http. The port defaults to 80 for http and 443 for // https. Everything up to the last "@" is the userinfo, so credentials may // hold unescaped '/', '?', '#' or '@'; percent-escapes in them are decoded. +// +// The Resolver never changes; NewResolver with WithRefreshCommand builds +// one whose credentials follow a refresh source. func Explicit(proxyURL string) (*Resolver, error) { + return newExplicit(proxyURL, options{}) +} + +func newExplicit(proxyURL string, o options) (*Resolver, error) { if strings.TrimSpace(proxyURL) == "" { return nil, errors.New("netproxy: empty proxy URL") } @@ -100,11 +176,14 @@ func Explicit(proxyURL string) (*Resolver, error) { if err != nil { return nil, err } - return &Resolver{mode: ModeExplicit, fixed: u}, nil + r := &Resolver{mode: ModeExplicit} + r.state.Store(&proxyState{fixed: u}) + r.configure(o) + return r, nil } -// FromEnvironment returns a Resolver built from a snapshot of the -// conventional proxy environment variables, taken now: +// FromEnvironment returns a Resolver built from the conventional proxy +// environment variables: // // - TLS and raw TCP targets (ProxyForAddr, and https:// / wss:// requests) // use the first usable one of HTTPS_PROXY, https_proxy, ALL_PROXY, @@ -127,26 +206,47 @@ func Explicit(proxyURL string) (*Resolver, error) { // https_proxy, which explicitly names the TLS proxy, makes FromEnvironment // fail, with an *EnvError naming it. // +// The variables are read now, and read again at most every +// DefaultRefreshInterval and whenever a proxy rejects the credentials (see +// Refresh), so a process whose environment is updated in place follows it. +// A later reading that fails keeps the previous one. +// // An environment with no usable proxy variables yields a Resolver that -// never proxies (Enabled reports false). Neither errors nor warnings ever +// does not proxy (Enabled reports false). Neither errors nor warnings ever // echo a value's credentials. func FromEnvironment() (*Resolver, error) { - return fromEnv(os.Getenv) + return newAuto(os.Getenv, options{}) } func fromEnv(getenv func(string) string) (*Resolver, error) { - r := &Resolver{mode: ModeAuto} + return newAuto(getenv, options{}) +} + +func newAuto(getenv func(string) string, o options) (*Resolver, error) { + st, err := readEnv(getenv) + if err != nil { + return nil, err + } + r := &Resolver{mode: ModeAuto, getenv: getenv} + r.state.Store(st) + r.configure(o) + return r, nil +} + +// readEnv reads the ModeAuto settings from the environment. +func readEnv(getenv func(string) string) (*proxyState, error) { + st := &proxyState{} for _, name := range []string{"HTTPS_PROXY", "https_proxy", "ALL_PROXY", "all_proxy"} { u, err := envProxy(getenv, name) if err != nil { if name == "HTTPS_PROXY" || name == "https_proxy" { return nil, err } - r.warnings = append(r.warnings, err) + st.warnings = append(st.warnings, err) continue } if u != nil { - r.secure = u + st.secure = u break } } @@ -157,17 +257,17 @@ func fromEnv(getenv func(string) string) (*Resolver, error) { for _, name := range plainVars { u, err := envProxy(getenv, name) if err != nil { - r.warnings = append(r.warnings, err) + st.warnings = append(st.warnings, err) continue } if u != nil { - r.plain = u + st.plain = u break } } - r.noProxyRaw, _ = firstEnv(getenv, "NO_PROXY", "no_proxy") - r.noProxy = parseNoProxy(r.noProxyRaw) - return r, nil + st.noProxyRaw, _ = firstEnv(getenv, "NO_PROXY", "no_proxy") + st.noProxy = parseNoProxy(st.noProxyRaw) + return st, nil } // envProxy parses one proxy variable: nil, nil when it is unset or blank, @@ -184,15 +284,16 @@ func envProxy(getenv func(string) string, name string) (*url.URL, error) { return u, nil } -// Warnings reports the proxy environment variables FromEnvironment skipped -// because their values are unusable, one *EnvError per variable, in -// precedence order. It is empty for other Resolvers. The messages are safe -// to log. +// Warnings reports the proxy environment variables the current reading of +// the environment skipped because their values are unusable, one *EnvError +// per variable, in precedence order. It is empty for Resolvers that do not +// read the environment. The messages are safe to log. func (r *Resolver) Warnings() []error { - if r == nil || len(r.warnings) == 0 { + st := r.snapshot() + if len(st.warnings) == 0 { return nil } - return append([]error(nil), r.warnings...) + return append([]error(nil), st.warnings...) } // Parse builds a Resolver from a -proxy style setting: "auto" (or "") reads @@ -200,15 +301,35 @@ func (r *Resolver) Warnings() []error { // else is an explicit proxy URL (see Explicit). Deciding whether "auto" // applies at all (for example only in a TCP-only transport mode) is the // caller's policy; compare the setting against ModeAuto for that. +// +// Parse(spec) is NewResolver(spec) without options. func Parse(spec string) (*Resolver, error) { + return NewResolver(spec) +} + +// NewResolver is Parse with options, for example a refresh command for +// proxies that rotate their credentials: +// +// r, err := netproxy.NewResolver(spec, +// netproxy.WithRefreshCommand(os.Getenv(netproxy.EnvRefreshCommand)), +// netproxy.WithRefreshErrorHandler(func(err error) { +// slog.Warn("proxy credential refresh failed", "err", err) +// })) +// +// Options do not apply to "off". With a refresh source, NewResolver runs it +// once before returning; if that run fails, the Resolver starts from the +// environment ("auto") or the given URL, and the error goes to the +// WithRefreshErrorHandler function. +func NewResolver(spec string, opts ...Option) (*Resolver, error) { + o := newOptions(opts) s := strings.TrimSpace(spec) switch strings.ToLower(s) { case "", ModeAuto: - return FromEnvironment() + return newAuto(os.Getenv, o) case ModeOff: return Off(), nil } - return Explicit(s) + return newExplicit(s, o) } // Mode reports ModeAuto, ModeOff or ModeExplicit. A nil Resolver is ModeOff. @@ -221,53 +342,76 @@ func (r *Resolver) Mode() string { // Enabled reports whether the Resolver can route any target through a proxy. // It is false for Off, for a nil Resolver and for an environment without -// proxy variables. +// proxy variables. A Resolver with a refresh source reports true: the +// source can supply a proxy at any time. func (r *Resolver) Enabled() bool { if r == nil { return false } switch r.mode { - case ModeExplicit: - return r.fixed != nil - case ModeAuto: - return r.secure != nil || r.plain != nil + case ModeExplicit, ModeAuto: + return r.source != nil || r.snapshot().proxies() } return false } +// proxies reports whether st routes anything through a proxy. +func (st *proxyState) proxies() bool { + return st.fixed != nil || st.secure != nil || st.plain != nil +} + // String describes the Resolver for logs, with credentials redacted, e.g. // "off", "http://***@proxy.internal:3128" or // "auto: http://***@proxy.internal:3128 (NO_PROXY=localhost,.corp)". +// Resolvers with a refresh source add " (credentials refreshed by +// command)" or " (credentials refreshed by callback)". func (r *Resolver) String() string { + st := r.snapshot() switch r.Mode() { case ModeExplicit: - return Redact(r.fixed) + return Redact(st.fixed) + r.refreshSuffix() case ModeAuto: - if !r.Enabled() { - return "auto: no proxy in environment" + r.ignoredSuffix() + if !st.proxies() { + return "auto: no proxy in environment" + st.ignoredSuffix() + r.refreshSuffix() } - s := "auto: " + Redact(r.secure) - if r.secure == nil { - s = "auto: http-only " + Redact(r.plain) - } else if r.plain != nil && Redact(r.plain) != Redact(r.secure) { - s += ", http " + Redact(r.plain) + var s string + switch { + case st.fixed != nil: + s = "auto: " + Redact(st.fixed) + case st.secure == nil: + s = "auto: http-only " + Redact(st.plain) + default: + s = "auto: " + Redact(st.secure) + if st.plain != nil && Redact(st.plain) != Redact(st.secure) { + s += ", http " + Redact(st.plain) + } } - if r.noProxyRaw != "" { - s += " (NO_PROXY=" + r.noProxyRaw + ")" + if st.noProxyRaw != "" { + s += " (NO_PROXY=" + st.noProxyRaw + ")" } - return s + r.ignoredSuffix() + return s + st.ignoredSuffix() + r.refreshSuffix() } return ModeOff } +func (r *Resolver) refreshSuffix() string { + switch { + case r.source == nil: + return "" + case r.command: + return " (credentials refreshed by command)" + } + return " (credentials refreshed by callback)" +} + // ignoredSuffix names the skipped variables for String, e.g. // " (ignored unusable HTTP_PROXY, ALL_PROXY)". -func (r *Resolver) ignoredSuffix() string { - if len(r.warnings) == 0 { +func (st *proxyState) ignoredSuffix() string { + if len(st.warnings) == 0 { return "" } - names := make([]string, 0, len(r.warnings)) - for _, w := range r.warnings { + names := make([]string, 0, len(st.warnings)) + for _, w := range st.warnings { var ee *EnvError if errors.As(w, &ee) { names = append(names, ee.Var) @@ -278,19 +422,30 @@ func (r *Resolver) ignoredSuffix() string { // ProxyForAddr returns the proxy to tunnel a raw TCP (or TLS) connection to // addr ("host:port") through, or nil to dial directly. addr is inspected as -// text only; host names are never resolved. +// text only; host names are never resolved. When the refresh interval has +// passed, the lookup starts a refresh in the background (see +// WithRefreshInterval) and answers from the current settings; it never +// waits for a refresh. func (r *Resolver) ProxyForAddr(addr string) (*url.URL, error) { - if !r.Enabled() { + // current first: a refresh can turn proxying on (a variable set in + // place, a refresh source's first URL). + st := r.current() + if !st.proxies() { return nil, nil } + return r.pickAddr(st, addr) +} + +// pickAddr applies st to a raw TCP target. +func (r *Resolver) pickAddr(st *proxyState, addr string) (*url.URL, error) { if r.mode == ModeExplicit { - return cloneURL(r.fixed), nil + return cloneURL(st.fixed), nil } host, port, err := net.SplitHostPort(addr) if err != nil { return nil, fmt.Errorf("netproxy: invalid target address %q: %w", addr, err) } - return cloneURL(r.pick(host, port, true)), nil + return cloneURL(st.pick(host, port, true)), nil } // ProxyForRequest returns the proxy for req, or nil for a direct connection. @@ -300,13 +455,25 @@ func (r *Resolver) ProxyForAddr(addr string) (*url.URL, error) { // // https:// and wss:// requests follow exactly the same rules as ProxyForAddr // (net/http then tunnels them with CONNECT, sending the host name); http:// -// and ws:// requests prefer HTTP_PROXY in ModeAuto. +// and ws:// requests prefer HTTP_PROXY in ModeAuto. The result is always +// the current proxy URL, and like ProxyForAddr the lookup never waits for a +// refresh; RefreshingTransport also retries a request once when the proxy +// rejects the credentials. func (r *Resolver) ProxyForRequest(req *http.Request) (*url.URL, error) { - if !r.Enabled() || req == nil || req.URL == nil { + if req == nil || req.URL == nil { return nil, nil } + st := r.current() + if !st.proxies() { + return nil, nil + } + return r.pickRequest(st, req), nil +} + +// pickRequest applies st to an HTTP request. +func (r *Resolver) pickRequest(st *proxyState, req *http.Request) *url.URL { if r.mode == ModeExplicit { - return cloneURL(r.fixed), nil + return cloneURL(st.fixed) } secure := true defaultPort := "443" @@ -319,16 +486,19 @@ func (r *Resolver) ProxyForRequest(req *http.Request) (*url.URL, error) { if port == "" { port = defaultPort } - return cloneURL(r.pick(req.URL.Hostname(), port, secure)), nil + return cloneURL(st.pick(req.URL.Hostname(), port, secure)) } // pick applies the ModeAuto rules to one target. -func (r *Resolver) pick(host, port string, secure bool) *url.URL { - proxy := r.secure - if !secure && r.plain != nil { - proxy = r.plain - } - if proxy == nil || !r.noProxy.useProxy(host, port) { +func (st *proxyState) pick(host, port string, secure bool) *url.URL { + proxy := st.secure + switch { + case st.fixed != nil: + proxy = st.fixed + case !secure && st.plain != nil: + proxy = st.plain + } + if proxy == nil || !st.noProxy.useProxy(host, port) { return nil } return proxy diff --git a/netproxy/refresh.go b/netproxy/refresh.go new file mode 100644 index 0000000..d6df288 --- /dev/null +++ b/netproxy/refresh.go @@ -0,0 +1,455 @@ +// SPDX-License-Identifier: AGPL-3.0-or-later + +package netproxy + +import ( + "context" + "errors" + "fmt" + "net/url" + "os/exec" + "strings" + "time" +) + +// EnvRefreshCommand is the environment variable that, by convention, holds +// the refresh command for WithRefreshCommand, e.g. +// +// PILOT_PROXY_CMD='bash -c '\''printf %s "$https_proxy"'\''' +// +// netproxy never reads it on its own: a program opts in by passing +// WithRefreshCommand(os.Getenv(EnvRefreshCommand)). +const EnvRefreshCommand = "PILOT_PROXY_CMD" + +// DefaultRefreshInterval is how long a refreshing Resolver uses the settings +// it last read before a lookup starts reading them again. +const DefaultRefreshInterval = 60 * time.Second + +// refreshTimeout bounds one refresh (one run of the refresh command). A +// variable so tests can shorten it. +var refreshTimeout = 10 * time.Second + +// maxRefreshOutput caps what a refresh command may print. +const maxRefreshOutput = 64 << 10 + +// Option configures a Resolver built by NewResolver. +type Option func(*options) + +type options struct { + source func(context.Context) (string, error) + command bool + interval time.Duration + onError func(error) +} + +func newOptions(opts []Option) options { + var o options + for _, opt := range opts { + if opt != nil { + opt(&o) + } + } + return o +} + +// WithRefreshCommand makes the Resolver take its proxy URL from the output +// of command, run with "sh -c" in the process's environment, stdin and +// stderr discarded, for at most 10 seconds. On Unix the command runs in a +// process group of its own, and a run that times out has the whole group +// killed, so the processes it started do not outlive it. The output, +// surrounding whitespace trimmed, must be one http:// or https:// proxy URL +// with its current credentials. It is never logged or included in errors. +// +// The command runs when NewResolver builds the Resolver, again in the +// background at most every refresh interval, and whenever a proxy rejects +// the credentials. Its URL replaces the explicit URL, or in ModeAuto the +// environment's proxy URLs, while NO_PROXY and the loopback exemption keep +// applying. If a run fails, or prints nothing or something that is not a +// proxy URL, the last good URL stays in use. +// +// A child process inherits this process's environment, so the command must +// read the current value from somewhere that tracks the rotation — in Meta +// Muse a new bash does: bash -c 'printf %s "$https_proxy"'. The command is +// trusted configuration, like the proxy URL itself. An empty command +// leaves the Resolver unchanged, so wiring an unset EnvRefreshCommand is +// harmless. +func WithRefreshCommand(command string) Option { + return func(o *options) { + if strings.TrimSpace(command) == "" { + return + } + o.source = commandSource(command) + o.command = true + } +} + +// WithRefreshFunc is WithRefreshCommand for a Go function: fn returns the +// current proxy URL, under the same rules as a refresh command's output. It +// must honour ctx, which carries the 10 second refresh deadline. Its errors +// are passed on (prefixed "netproxy: proxy refresh: "), so they must not +// contain credentials. A nil fn leaves the Resolver unchanged. +func WithRefreshFunc(fn func(ctx context.Context) (string, error)) Option { + return func(o *options) { + if fn == nil { + return + } + o.source = fn + o.command = false + } +} + +// WithRefreshInterval sets how long the Resolver uses the settings it last +// read before a lookup starts reading them again. That timed refresh runs in +// the background: the lookup that starts it, and every lookup while it runs, +// answers from the settings in hand, so a slow or hung refresh source never +// holds up a dial or a request (a proxy rejecting those settings is what +// waits for a refresh). Zero means DefaultRefreshInterval; a negative +// interval turns timed refreshes off, leaving only the refreshes a rejected +// credential triggers (and explicit Refresh calls). +func WithRefreshInterval(d time.Duration) Option { + return func(o *options) { o.interval = d } +} + +// WithRefreshErrorHandler has fn called with the error when a refresh +// fails, once per run of consecutive failures: after a failure it is not +// called again until a refresh has succeeded. The error never contains +// credentials or the refresh command's output (a WithRefreshFunc +// function's own errors aside). fn is called from the refreshing goroutine +// and must not block. +func WithRefreshErrorHandler(fn func(error)) Option { + return func(o *options) { o.onError = fn } +} + +// configure applies o to a new Resolver and, with a refresh source, runs +// the first refresh. +func (r *Resolver) configure(o options) { + r.source, r.command, r.onError = o.source, o.command, o.onError + r.interval = o.interval + if r.interval == 0 { + r.interval = DefaultRefreshInterval + } + r.lastRun = time.Now() + if r.source != nil { + // Failures are reported through onError and leave the initial + // settings in place. + _ = r.Refresh(context.Background()) + } +} + +// refreshable reports whether the Resolver has settings to re-read. +func (r *Resolver) refreshable() bool { + if r == nil { + return false + } + switch r.mode { + case ModeAuto: + return true + case ModeExplicit: + return r.source != nil + } + return false +} + +// snapshot returns the current settings without refreshing. +func (r *Resolver) snapshot() *proxyState { + if r != nil { + if st := r.state.Load(); st != nil { + return st + } + } + return &emptyState +} + +// current returns the settings for a lookup. When the refresh interval has +// passed it starts a refresh, but never waits for it: the lookup, like every +// other one until the refresh lands, gets the settings in hand. Those are +// the last good ones, which keep working until the proxy rotates them, and +// a rotation is handled by reauth, which does wait. Waiting here instead +// would make a slow refresh source fail dials whose own deadline is shorter +// than the refresh's, although nothing was wrong with their credentials. +func (r *Resolver) current() *proxyState { + if r.refreshable() && r.interval >= 0 { + r.mu.Lock() + if r.inflight == nil && time.Since(r.lastRun) >= r.interval { + r.startLocked() + } + r.mu.Unlock() + } + return r.snapshot() +} + +// Refresh re-reads the Resolver's settings now: the environment in +// ModeAuto, and the refresh source if there is one. A refresh already in +// progress is joined rather than repeated. On error the previous settings +// stay in use. Refresh is a no-op for Resolvers with nothing to re-read +// (nil, Off, and Explicit without a refresh source). The error never +// contains credentials; ctx only bounds the wait. +func (r *Resolver) Refresh(ctx context.Context) error { + if !r.refreshable() { + return nil + } + r.mu.Lock() + call := r.inflight + if call == nil { + call = r.startLocked() + } + r.mu.Unlock() + if !call.wait(ctx) { + return ctx.Err() + } + return call.err +} + +// attemptCount returns the number of refreshes started so far. +func (r *Resolver) attemptCount() uint64 { + if !r.refreshable() { + return 0 + } + r.mu.Lock() + defer r.mu.Unlock() + return r.attempts +} + +// reauth handles a proxy rejecting the credentials in used. start is +// attemptCount from before used was picked; pick applies a reading of the +// settings to the same target. It reports the proxy to retry through (nil +// for a direct connection) and whether a retry is worthwhile, which it is +// only when the current settings differ from used: +// +// - they already differ (another caller refreshed): no new refresh; +// - a refresh is in progress: wait for it; +// - a refresh started after used was picked and has finished: it could +// not do better, so no new refresh and no retry; +// - otherwise: refresh now. +// +// err is the error of the refresh this call waited for, if it failed. +func (r *Resolver) reauth(ctx context.Context, start uint64, used *url.URL, pick func(*proxyState) *url.URL) (next *url.URL, retry bool, err error) { + if !r.refreshable() { + return nil, false, nil + } + r.mu.Lock() + if cur := pick(r.snapshot()); !sameURL(cur, used) { + r.mu.Unlock() + return cur, true, nil + } + call := r.inflight + if call == nil { + if r.attempts != start { + r.mu.Unlock() + return nil, false, nil + } + call = r.startLocked() + } + r.mu.Unlock() + if !call.wait(ctx) { + return nil, false, nil + } + if cur := pick(r.snapshot()); !sameURL(cur, used) { + return cur, true, nil + } + return nil, false, call.err +} + +// refreshCall is one refresh, shared by everyone waiting for it. +type refreshCall struct { + done chan struct{} + err error // set before done is closed +} + +// wait blocks until the refresh finishes (true) or ctx ends (false). It +// gives up after the refresh deadline plus a margin even if ctx never ends, +// in case a WithRefreshFunc function ignores its context. +func (c *refreshCall) wait(ctx context.Context) bool { + t := time.NewTimer(refreshTimeout + 5*time.Second) + defer t.Stop() + select { + case <-c.done: + return true + case <-ctx.Done(): + case <-t.C: + } + return false +} + +// startLocked starts a refresh. r.mu must be held. +func (r *Resolver) startLocked() *refreshCall { + c := &refreshCall{done: make(chan struct{})} + r.inflight = c + r.attempts++ + go r.run(c) + return c +} + +// run performs refresh c. It runs on its own goroutine, detached from the +// callers' contexts, so one caller giving up does not fail the refresh for +// the others. +func (r *Resolver) run(c *refreshCall) { + ctx, cancel := context.WithTimeout(context.Background(), refreshTimeout) + st, err := r.load(ctx) + cancel() + + r.mu.Lock() + report := false + if err == nil { + r.state.Store(st) + r.failing = false + } else { + report = !r.failing + r.failing = true + } + r.lastRun = time.Now() + r.inflight = nil + c.err = err + r.mu.Unlock() + + if report && r.onError != nil { + r.onError(err) + } + close(c.done) +} + +// load reads a fresh copy of the settings. +func (r *Resolver) load(ctx context.Context) (*proxyState, error) { + var st *proxyState + if r.mode == ModeAuto { + var err error + if st, err = readEnv(r.getenv); err != nil { + return nil, err + } + } else { + cp := *r.snapshot() + st = &cp + } + if r.source != nil { + raw, err := r.source(ctx) + if err != nil { + if r.command { + return nil, err + } + return nil, fmt.Errorf("netproxy: proxy refresh: %w", err) + } + u, err := parseRefreshed(raw) + if err != nil { + return nil, err + } + st.fixed = u + } + return st, nil +} + +// errRefreshOutput never echoes the output: it would carry credentials. +var errRefreshOutput = errors.New("netproxy: refreshed proxy URL is unusable (value withheld; want http://[user:pass@]host[:port] or https://...)") + +// parseRefreshed validates a refresh source's output. Unlike an environment +// variable, it must spell out its scheme, so a stray token is never taken +// for a proxy host name (which errors and logs would then show). +func parseRefreshed(raw string) (*url.URL, error) { + s := strings.TrimSpace(raw) + if s == "" { + return nil, errors.New("netproxy: refreshed proxy URL is empty") + } + scheme, _, ok := strings.Cut(s, "://") + if !ok { + return nil, errRefreshOutput + } + switch strings.ToLower(scheme) { + case "http", "https": + default: + return nil, errRefreshOutput + } + u, err := parseProxyURL(s) + if err != nil { + return nil, errRefreshOutput + } + return u, nil +} + +// commandSource runs command with "sh -c" and returns what it prints. +// Errors say how the command failed (exit status, timeout, not startable) +// and never include its output or stderr. +func commandSource(command string) func(context.Context) (string, error) { + return func(ctx context.Context) (string, error) { + cmd := exec.CommandContext(ctx, "sh", "-c", command) + var out cappedBuffer + cmd.Stdout = &out + // Stdin and Stderr stay nil (the null device): stderr could echo + // credentials, e.g. under "set -x". + cmd.WaitDelay = time.Second // a child left holding stdout cannot hang Wait + // On timeout, kill everything the command started, not just sh: + // otherwise every timed-out refresh leaves its hung children behind. + ownProcessGroup(cmd) + err := cmd.Run() + switch { + case ctx.Err() != nil: + return "", fmt.Errorf("netproxy: refresh command timed out after %v", refreshTimeout) + case err != nil: + var ee *exec.ExitError + if errors.As(err, &ee) { + return "", fmt.Errorf("netproxy: refresh command failed: %s", ee.ProcessState) + } + if errors.Is(err, exec.ErrWaitDelay) { + // sh has exited, but a process it started held stdout + // open until now, so the group still has a member: end it. + _ = killProcessGroup(cmd) + return "", errors.New("netproxy: refresh command left a process holding its output open") + } + return "", fmt.Errorf("netproxy: refresh command failed to run: %v", err) + case out.overflow: + return "", fmt.Errorf("netproxy: refresh command printed more than %d bytes", maxRefreshOutput) + } + return string(out.buf), nil + } +} + +// cappedBuffer keeps the first maxRefreshOutput bytes written to it and +// notes whether there were more. +type cappedBuffer struct { + buf []byte + overflow bool +} + +func (b *cappedBuffer) Write(p []byte) (int, error) { + if room := maxRefreshOutput - len(b.buf); len(p) > room { + b.buf = append(b.buf, p[:room]...) + b.overflow = true + return len(p), nil + } + b.buf = append(b.buf, p...) + return len(p), nil +} + +// sameProxies reports whether st and o route through the same proxies with +// the same credentials (NO_PROXY aside). +func (st *proxyState) sameProxies(o *proxyState) bool { + return sameURL(st.fixed, o.fixed) && sameURL(st.secure, o.secure) && sameURL(st.plain, o.plain) +} + +// sameURL reports whether a and b are the same proxy with the same +// credentials (nil meaning a direct connection). +func sameURL(a, b *url.URL) bool { + if a == nil || b == nil { + return a == b + } + return a.String() == b.String() +} + +// refreshFailedError adds the reason a credential refresh failed to the +// error that triggered the refresh. Both stay reachable via errors.As/Is. +type refreshFailedError struct { + err error + refresh error +} + +func (e *refreshFailedError) Error() string { + return e.err.Error() + " (proxy credential refresh failed: " + e.refresh.Error() + ")" +} + +func (e *refreshFailedError) Unwrap() []error { return []error{e.err, e.refresh} } + +// withRefreshError returns err, annotated with refreshErr when there is one. +func withRefreshError(err, refreshErr error) error { + if refreshErr == nil { + return err + } + return &refreshFailedError{err: err, refresh: refreshErr} +} diff --git a/netproxy/refresh_other.go b/netproxy/refresh_other.go new file mode 100644 index 0000000..05b9991 --- /dev/null +++ b/netproxy/refresh_other.go @@ -0,0 +1,14 @@ +// SPDX-License-Identifier: AGPL-3.0-or-later + +//go:build !unix + +package netproxy + +import "os/exec" + +// ownProcessGroup leaves cmd as it is: without Unix process groups, a +// timed-out refresh command is killed on its own (exec.Cmd's default). +func ownProcessGroup(*exec.Cmd) {} + +// killProcessGroup is a no-op without Unix process groups. +func killProcessGroup(*exec.Cmd) error { return nil } diff --git a/netproxy/refresh_unix.go b/netproxy/refresh_unix.go new file mode 100644 index 0000000..5cf5c0b --- /dev/null +++ b/netproxy/refresh_unix.go @@ -0,0 +1,34 @@ +// SPDX-License-Identifier: AGPL-3.0-or-later + +//go:build unix + +package netproxy + +import ( + "errors" + "os" + "os/exec" + "syscall" +) + +// ownProcessGroup makes cmd the leader of a new process group and has its +// context's cancellation SIGKILL that whole group rather than only cmd, so +// that nothing the command started (a pipeline, a nested shell, a +// background job) outlives a timed-out run. +func ownProcessGroup(cmd *exec.Cmd) { + cmd.SysProcAttr = &syscall.SysProcAttr{Setpgid: true} + cmd.Cancel = func() error { return killProcessGroup(cmd) } +} + +// killProcessGroup SIGKILLs the process group cmd leads. It reports +// os.ErrProcessDone when the group has no members left. +func killProcessGroup(cmd *exec.Cmd) error { + if cmd.Process == nil { + return nil + } + err := syscall.Kill(-cmd.Process.Pid, syscall.SIGKILL) + if errors.Is(err, syscall.ESRCH) { + return os.ErrProcessDone + } + return err +} diff --git a/netproxy/transport.go b/netproxy/transport.go new file mode 100644 index 0000000..4f1a718 --- /dev/null +++ b/netproxy/transport.go @@ -0,0 +1,258 @@ +// SPDX-License-Identifier: AGPL-3.0-or-later + +package netproxy + +import ( + "context" + "errors" + "fmt" + "net" + "net/http" + "net/http/httptrace" + "net/url" + "strings" + "sync/atomic" +) + +// RefreshingTransport returns an http.RoundTripper for HTTP clients behind +// a proxy that rotates its credentials. It sends requests through a copy of +// base (http.DefaultTransport when base is nil) whose Proxy is +// r.ProxyForRequest, so every new connection uses the current credentials, +// and, when the proxy rejects them on CONNECT, refreshes r and retries the +// request once — only when the refresh produced different credentials, and +// only for a request that can be sent again: +// +// - GET, HEAD, OPTIONS and TRACE requests, and requests carrying an +// Idempotency-Key or X-Idempotency-Key header, without a body or with +// GetBody set; +// - after a 407, any request with GetBody set: the proxy refused the +// tunnel, so the server never saw the first attempt. +// +// A request that cannot be replayed (a POST whose body cannot be rewound, +// say) gets the error, but the refresh still happens, so the next request +// goes out with fresh credentials. +// +// net/http reports a refused CONNECT only as an error holding the proxy's +// reason phrase — "Proxy Authentication Required", "unknown status code" +// when there is none, or whatever text the proxy chose — without the status +// code. RefreshingTransport therefore reads the status from the CONNECT +// response itself (http.Transport.OnProxyConnectResponse, chained after +// base's own hook) and turns every non-200 answer into a *ConnectError, +// whose message never includes proxy-supplied text. +// +// A CONNECT response net/http cannot parse ("malformed HTTP status code", +// say), which is how some proxies' rejections surface, also triggers a +// refresh, and its error likewise says only what was wrong with the +// response, never its text. That holds only while the tunnel is being set +// up: once net/http has a connection to send the request on, an unparseable +// response came from the server through the tunnel, and its error is +// returned untouched, without a refresh. Having no status code to prove it +// a refusal, such a rejection is retried only for the idempotent methods +// above. +// +// net/http pools connections by proxy URL, credentials included, so once a +// refresh changes the proxy settings, the idle connections opened under the +// old ones can never be picked again. The first request after such a change +// therefore closes the idle connections (in-flight ones finish normally). +// The returned RoundTripper has a CloseIdleConnections method, so +// http.Client.CloseIdleConnections reaches the copy of base. +func RefreshingTransport(base *http.Transport, r *Resolver) http.RoundTripper { + var tr *http.Transport + switch dt, ok := http.DefaultTransport.(*http.Transport); { + case base != nil: + tr = base.Clone() + case ok: + tr = dt.Clone() + default: + tr = &http.Transport{} + } + tr.Proxy = func(req *http.Request) (*url.URL, error) { + u, err := r.ProxyForRequest(req) + if a, ok := req.Context().Value(attemptKey{}).(*attempt); ok { + a.proxy.Store(u) + } + return u, err + } + next := tr.OnProxyConnectResponse + tr.OnProxyConnectResponse = func(ctx context.Context, proxyURL *url.URL, connectReq *http.Request, res *http.Response) error { + if next != nil { + if err := next(ctx, proxyURL, connectReq, res); err != nil { + return err + } + } + // net/http itself accepts exactly 200. + if res.StatusCode == http.StatusOK { + return nil + } + ce := &ConnectError{Target: connectReq.Host, StatusCode: res.StatusCode} + if res.StatusCode == http.StatusProxyAuthRequired { + return &authRejectedError{ConnectError: ce, proxy: cloneURL(proxyURL)} + } + return ce + } + t := &refreshingTransport{tr: tr, r: r} + t.pooled.Store(r.snapshot()) + return t +} + +type refreshingTransport struct { + tr *http.Transport + r *Resolver + + // pooled is the Resolver's reading that the connections in tr's idle + // pool were (last known to be) opened under. + pooled atomic.Pointer[proxyState] +} + +// attempt records, for one try of a request, what net/http did with it. +// It travels in the request's context. +type attempt struct { + // proxy is what the Proxy function returned for the latest connection + // lookup: the proxy URL, with the credentials, that a refused CONNECT + // was sent with. + proxy atomic.Pointer[url.URL] + // connected is set once net/http has a connection to send the request + // on, so any CONNECT for it succeeded. It is reset whenever net/http + // starts looking for a connection, since it retries some failures on a + // new one. + connected atomic.Bool +} + +type attemptKey struct{} + +// authRejectedError is a 407 answer to a CONNECT, with the proxy URL (and +// so the credentials) it was sent with. Its message is the ConnectError's. +type authRejectedError struct { + *ConnectError + proxy *url.URL // never printed +} + +func (e *authRejectedError) Unwrap() error { return e.ConnectError } + +func (t *refreshingTransport) RoundTrip(req *http.Request) (*http.Response, error) { + // Noted before the lookup, so a timed refresh the lookup starts counts as + // having happened after it (see Resolver.reauth). + start := t.r.attemptCount() + t.dropStrandedConns() + resp, rej, err := t.send(req) + if rej.proxy == nil { + return resp, err + } + _, retry, refreshErr := t.r.reauth(req.Context(), start, rej.proxy, func(st *proxyState) *url.URL { + return t.r.pickRequest(st, req) + }) + if !retry || !canReplay(req, rej.refused) { + return nil, withRefreshError(err, refreshErr) + } + again, rewindErr := rewind(req) + if rewindErr != nil { + return nil, err + } + resp, _, err = t.send(again) + return resp, err +} + +// rejection describes a CONNECT the proxy turned down. +type rejection struct { + // proxy is the proxy URL, with the credentials, the CONNECT was sent + // with; nil when the proxy did not reject the credentials. + proxy *url.URL + // refused is set for a 407: the tunnel was never opened, so the server + // cannot have seen the request. + refused bool +} + +// send runs one try of req. When the proxy rejected the credentials it also +// says with which ones, and it replaces net/http's error for an unparseable +// CONNECT response with one that does not quote the response. +func (t *refreshingTransport) send(req *http.Request) (*http.Response, rejection, error) { + a := new(attempt) + ctx := context.WithValue(req.Context(), attemptKey{}, a) + ctx = httptrace.WithClientTrace(ctx, &httptrace.ClientTrace{ + GetConn: func(string) { a.connected.Store(false) }, + GotConn: func(httptrace.GotConnInfo) { a.connected.Store(true) }, + }) + resp, err := t.tr.RoundTrip(req.WithContext(ctx)) + if err == nil { + resp.Request = req // not the copy carrying the trace + return resp, rejection{}, nil + } + var ae *authRejectedError + if errors.As(err, &ae) { + return nil, rejection{proxy: ae.proxy, refused: true}, err + } + // Without a connection, the response net/http could not parse was the + // proxy's answer to CONNECT; with one, it came from the server. + used := a.proxy.Load() + if fault := responseFault(err); fault != "" && used != nil && tunnelled(req) && !a.connected.Load() { + return nil, rejection{proxy: used}, fmt.Errorf("proxy CONNECT %s: %w", connectTarget(req.URL), &badConnectResponse{what: fault}) + } + return nil, rejection{}, err +} + +// dropStrandedConns closes the idle connections once the Resolver's proxy +// settings have changed since they were opened: net/http keys its pool by +// proxy URL, credentials included, and would never pick them again. +func (t *refreshingTransport) dropStrandedConns() { + cur := t.r.snapshot() + if prev := t.pooled.Swap(cur); prev != cur && !prev.sameProxies(cur) { + t.tr.CloseIdleConnections() + } +} + +// connectTarget is the "host:port" net/http sends CONNECT for u. +func connectTarget(u *url.URL) string { + port := u.Port() + if port == "" { + port = "443" + } + return net.JoinHostPort(u.Hostname(), port) +} + +// CloseIdleConnections closes the idle connections of the underlying +// transport. +func (t *refreshingTransport) CloseIdleConnections() { t.tr.CloseIdleConnections() } + +// tunnelled reports whether net/http reaches req's target through a CONNECT +// tunnel when it uses a proxy. +func tunnelled(req *http.Request) bool { + switch strings.ToLower(req.URL.Scheme) { + case "https", "wss": + return true + } + return false +} + +// canReplay reports whether req may be sent again. neverSent says the first +// attempt certainly did not reach the server (the proxy refused the +// tunnel), so any request whose body can be rewound qualifies; otherwise +// only idempotent ones do, as in net/http's own retry rules. +func canReplay(req *http.Request, neverSent bool) bool { + if req.Body != nil && req.Body != http.NoBody && req.GetBody == nil { + return false + } + if neverSent && req.GetBody != nil { + return true + } + switch req.Method { + case "", http.MethodGet, http.MethodHead, http.MethodOptions, http.MethodTrace: + return true + } + _, key := req.Header["Idempotency-Key"] + _, xkey := req.Header["X-Idempotency-Key"] + return key || xkey +} + +// rewind returns a copy of req to send again, with a fresh body from +// GetBody. The RoundTripper contract forbids modifying req itself. +func rewind(req *http.Request) (*http.Request, error) { + again := req.Clone(req.Context()) + if req.Body != nil && req.Body != http.NoBody { + body, err := req.GetBody() + if err != nil { + return nil, err + } + again.Body = body + } + return again, nil +} diff --git a/netproxy/zz_helpers_test.go b/netproxy/zz_helpers_test.go index 62673ca..5336e3b 100644 --- a/netproxy/zz_helpers_test.go +++ b/netproxy/zz_helpers_test.go @@ -34,9 +34,15 @@ type testProxy struct { // wantAuth, when set, is the only Proxy-Authorization value accepted; // anything else gets rejectStatus (407 by default) with rejectReason. + // Guarded by mu: setAuth rotates it while the proxy runs. wantAuth string rejectStatus int rejectReason string + // rejectLine, when set, replaces the whole status line of a rejection + // (to send one net/http cannot parse). + rejectLine string + // rejected counts CONNECTs refused for their credentials. + rejected atomic.Int32 // hang makes the proxy read the CONNECT request and never answer. hang bool @@ -64,6 +70,19 @@ func withReject(status int, reason string) proxyOpt { func withHang() proxyOpt { return func(p *testProxy) { p.hang = true } } +// withRejectLine makes credential rejections use line as the status line. +func withRejectLine(line string) proxyOpt { + return func(p *testProxy) { p.rejectLine = line } +} + +// setAuth rotates the credentials the proxy accepts. Tunnels already open +// are not affected, as with a real rotating proxy. +func (p *testProxy) setAuth(user, pass string) { + p.mu.Lock() + p.wantAuth = basicAuth(user, pass) + p.mu.Unlock() +} + func withTLS(cert tls.Certificate) proxyOpt { return func(p *testProxy) { p.ln = tls.NewListener(p.ln, &tls.Config{Certificates: []tls.Certificate{cert}, MinVersion: tls.VersionTLS12}) @@ -146,6 +165,7 @@ func (p *testProxy) serve(conn net.Conn) { p.targets = append(p.targets, req.RequestURI) p.hostHdrs = append(p.hostHdrs, req.Host) p.auths = append(p.auths, req.Header.Get("Proxy-Authorization")) + wantAuth := p.wantAuth p.mu.Unlock() waitClientClose := func() { @@ -161,8 +181,13 @@ func (p *testProxy) serve(conn net.Conn) { fmt.Fprintf(conn, "HTTP/1.1 405 Method Not Allowed\r\nContent-Length: 0\r\n\r\n") return } - if p.wantAuth != "" && req.Header.Get("Proxy-Authorization") != p.wantAuth { - fmt.Fprintf(conn, "HTTP/1.1 %d %s\r\nProxy-Authenticate: Basic realm=\"test\"\r\nContent-Length: 0\r\n\r\n", p.rejectStatus, p.rejectReason) + if wantAuth != "" && req.Header.Get("Proxy-Authorization") != wantAuth { + p.rejected.Add(1) + line := fmt.Sprintf("HTTP/1.1 %d %s", p.rejectStatus, p.rejectReason) + if p.rejectLine != "" { + line = p.rejectLine + } + fmt.Fprintf(conn, "%s\r\nProxy-Authenticate: Basic realm=\"test\"\r\nContent-Length: 0\r\n\r\n", line) waitClientClose() return } diff --git a/netproxy/zz_refresh_test.go b/netproxy/zz_refresh_test.go new file mode 100644 index 0000000..6c1d6fc --- /dev/null +++ b/netproxy/zz_refresh_test.go @@ -0,0 +1,1000 @@ +// SPDX-License-Identifier: AGPL-3.0-or-later + +package netproxy + +import ( + "context" + "errors" + "fmt" + "io" + "net" + "net/http" + "net/url" + "os" + "path/filepath" + "runtime" + "strings" + "sync" + "sync/atomic" + "testing" + "time" +) + +// credSource simulates a sandbox that rotates its proxy credentials: the +// test proxy only accepts the newest ones, and the refresh command prints +// them, the way a fresh shell sees the current HTTPS_PROXY in Meta Muse. +// The command also logs each run, so tests can count refreshes. +type credSource struct { + t *testing.T + proxy *testProxy + dir string + gen int + // delay is a shell snippet run before printing (e.g. "sleep 0.3;"). + delay string +} + +func newCredSource(t *testing.T, proxy *testProxy) *credSource { + t.Helper() + if runtime.GOOS == "windows" { + t.Skip("refresh commands run with sh -c") + } + c := &credSource{t: t, proxy: proxy, dir: t.TempDir()} + c.rotate() + return c +} + +// rotate issues new credentials. From now on the proxy rejects the old ones +// (tunnels it already opened stay up) and the command prints the new ones. +// The password holds characters that need escaping in a URL. +func (c *credSource) rotate() (user, pass string) { + c.gen++ + user, pass = "muse-agent", fmt.Sprintf("tok/%d?r#s@t", c.gen) + c.proxy.setAuth(user, pass) + c.write(c.urlFor(user, pass)) + return user, pass +} + +func (c *credSource) urlFor(user, pass string) string { + return "http://" + url.UserPassword(user, pass).String() + "@" + c.proxy.addr() +} + +// current is the URL the command prints now. +func (c *credSource) current() string { + b, err := os.ReadFile(filepath.Join(c.dir, "url")) + if err != nil { + c.t.Fatal(err) + } + return strings.TrimSpace(string(b)) +} + +// write replaces what the command prints, atomically. +func (c *credSource) write(proxyURL string) { + c.t.Helper() + tmp := filepath.Join(c.dir, "url.tmp") + if err := os.WriteFile(tmp, []byte(proxyURL+"\n"), 0o600); err != nil { + c.t.Fatal(err) + } + if err := os.Rename(tmp, filepath.Join(c.dir, "url")); err != nil { + c.t.Fatal(err) + } +} + +// fail makes the command exit 3 (true) or work again (false). +func (c *credSource) fail(on bool) { + c.t.Helper() + path := filepath.Join(c.dir, "fail") + if on { + if err := os.WriteFile(path, nil, 0o600); err != nil { + c.t.Fatal(err) + } + return + } + if err := os.Remove(path); err != nil { + c.t.Fatal(err) + } +} + +func (c *credSource) command() string { + return fmt.Sprintf("echo run >> %s; [ -e %s ] && exit 3; %s cat %s", + shellQuote(filepath.Join(c.dir, "runs")), shellQuote(filepath.Join(c.dir, "fail")), + c.delay, shellQuote(filepath.Join(c.dir, "url"))) +} + +// runs counts the command's runs so far. +func (c *credSource) runs() int { + b, err := os.ReadFile(filepath.Join(c.dir, "runs")) + if errors.Is(err, os.ErrNotExist) { + return 0 + } + if err != nil { + c.t.Fatal(err) + } + return strings.Count(string(b), "run\n") +} + +func shellQuote(s string) string { return "'" + strings.ReplaceAll(s, "'", `'\''`) + "'" } + +// errorLog collects WithRefreshErrorHandler calls. +type errorLog struct { + mu sync.Mutex + errs []error +} + +func (l *errorLog) handler() Option { + return WithRefreshErrorHandler(func(err error) { + l.mu.Lock() + l.errs = append(l.errs, err) + l.mu.Unlock() + }) +} + +func (l *errorLog) all() []error { + l.mu.Lock() + defer l.mu.Unlock() + return append([]error(nil), l.errs...) +} + +// syncEnv is a getenv whose values tests change while a Resolver reads it. +type syncEnv struct { + mu sync.Mutex + m map[string]string +} + +func (e *syncEnv) get(k string) string { + e.mu.Lock() + defer e.mu.Unlock() + return e.m[k] +} + +func (e *syncEnv) set(k, v string) { + e.mu.Lock() + e.m[k] = v + e.mu.Unlock() +} + +// settle waits for the refresh in flight, if any, to land. Timed refreshes +// run in the background, so tests that want their result wait for it here. +func settle(t *testing.T, r *Resolver) { + t.Helper() + r.mu.Lock() + c := r.inflight + r.mu.Unlock() + if c == nil { + return + } + select { + case <-c.done: + case <-time.After(refreshTimeout + 5*time.Second): + t.Fatal("the refresh in flight never finished") + } +} + +// lookupAfterTimedRefresh returns the proxy for addr once a timed refresh +// that started after the caller's last change has landed. r's refresh +// interval must be tiny. +func lookupAfterTimedRefresh(t *testing.T, r *Resolver, addr string) string { + t.Helper() + settle(t, r) // one already in flight may have read the old settings + proxyFor(t, r, addr) // the interval has passed: starts a refresh + settle(t, r) + return proxyFor(t, r, addr) +} + +func dialEcho(t *testing.T, d *Dialer, target, msg string) net.Conn { + t.Helper() + c, err := d.DialContext(context.Background(), "tcp", target) + if err != nil { + t.Fatalf("dial %s: %v", target, err) + } + roundTrip(t, c, msg) + return c +} + +// The Muse scenario end to end: a daemon started with HTTPS_PROXY holding +// the launch-time credentials, and PILOT_PROXY_CMD printing the current +// ones. The proxy rotates several times; every new dial still succeeds +// (one 407, one command run, one retry per rotation) and every tunnel +// opened before a rotation keeps working. +func TestRefreshKeepsLongRunningDialerWorkingAcrossRotations(t *testing.T) { + t.Parallel() + echoIP, echoPort := newEchoServer(t) + proxy := newTestProxy(t, map[string]string{echoHost: echoIP}) + creds := newCredSource(t, proxy) + var errs errorLog + env := envMap(map[string]string{"HTTPS_PROXY": creds.current(), "NO_PROXY": "localhost"}) + r, err := newAuto(env, newOptions([]Option{WithRefreshCommand(creds.command()), WithRefreshInterval(time.Hour), errs.handler()})) + if err != nil { + t.Fatalf("newAuto: %v", err) + } + if got := creds.runs(); got != 1 { + t.Fatalf("command ran %d times while building the Resolver, want 1", got) + } + d := NewDialer(r) + target := net.JoinHostPort(echoHost, echoPort) + + var open []net.Conn + defer func() { + for _, c := range open { + c.Close() + } + }() + open = append(open, dialEcho(t, d, target, "before any rotation")) + + const rotations = 4 + for i := 1; i <= rotations; i++ { + user, pass := creds.rotate() + open = append(open, dialEcho(t, d, target, fmt.Sprintf("after rotation %d", i))) + for j, c := range open { + roundTrip(t, c, fmt.Sprintf("tunnel %d still up after rotation %d", j, i)) + } + _, _, auths := proxy.seen() + if last := auths[len(auths)-1]; last != basicAuth(user, pass) { + t.Fatalf("rotation %d: retry sent %q, want the new credentials", i, last) + } + } + if got := proxy.rejected.Load(); got != rotations { + t.Fatalf("proxy rejected %d CONNECTs, want one per rotation (%d)", got, rotations) + } + if got := creds.runs(); got != 1+rotations { + t.Fatalf("command ran %d times, want %d (build + one per rotation)", got, 1+rotations) + } + if e := errs.all(); len(e) != 0 { + t.Fatalf("unexpected refresh errors: %v", e) + } +} + +// When the refreshed credentials are the ones the proxy just rejected, the +// dial fails with the 407 straight away: no retry, no loop. +func TestRefreshUnchangedCredentialsDoNotRetry(t *testing.T) { + t.Parallel() + proxy := newTestProxy(t, nil) + creds := newCredSource(t, proxy) + stale := creds.urlFor("muse-agent", "expired") + creds.write(stale) // the command keeps printing credentials the proxy refuses + r, err := NewResolver(stale, WithRefreshCommand(creds.command()), WithRefreshInterval(time.Hour)) + if err != nil { + t.Fatal(err) + } + d := NewDialer(r) + target := "registry.pilot.invalid:443" + for i := 1; i <= 3; i++ { + _, err := d.DialContext(context.Background(), "tcp", target) + want := "proxy CONNECT " + target + ": 407 Proxy Authentication Required" + if err == nil || err.Error() != want { + t.Fatalf("dial %d: error = %v, want %q", i, err, want) + } + var ce *ConnectError + if !errors.As(err, &ce) || ce.StatusCode != http.StatusProxyAuthRequired { + t.Fatalf("dial %d: %v is not a 407 ConnectError", i, err) + } + if got := proxy.rejected.Load(); got != int32(i) { + t.Fatalf("dial %d: proxy saw %d CONNECTs, want %d (no retries)", i, got, i) + } + // One refresh per failed dial at most: build + i. + if got := creds.runs(); got != 1+i { + t.Fatalf("dial %d: command ran %d times, want %d", i, got, 1+i) + } + } + + // The same holds for ModeAuto without a command: the environment is + // re-read, found unchanged, and the dial fails once. + proxy2 := newTestProxy(t, nil, withAuth("u", "current")) + auto, err := fromEnv(envMap(map[string]string{"HTTPS_PROXY": proxy2.url("u:old")})) + if err != nil { + t.Fatal(err) + } + if _, err := NewDialer(auto).Dial("tcp", target); err == nil { + t.Fatal("dial with stale environment credentials succeeded") + } + if got := proxy2.rejected.Load(); got != 1 { + t.Fatalf("auto: proxy saw %d CONNECTs, want 1", got) + } +} + +// A refresh command that starts failing never costs the working +// credentials: lookups keep the last good URL, the failure is reported +// once per run of failures, and a dial the proxy rejects says why the +// refresh did not help. +func TestRefreshCommandFailureKeepsLastGoodCredentials(t *testing.T) { + t.Parallel() + echoIP, echoPort := newEchoServer(t) + proxy := newTestProxy(t, map[string]string{echoHost: echoIP}) + creds := newCredSource(t, proxy) + _, pass1 := creds.user1() + var errs errorLog + r, err := NewResolver(proxy.url("launch:time"), WithRefreshCommand(creds.command()), WithRefreshInterval(time.Hour), errs.handler()) + if err != nil { + t.Fatal(err) + } + d := NewDialer(r) + target := net.JoinHostPort(echoHost, echoPort) + dialEcho(t, d, target, "working").Close() + + creds.fail(true) + for i := 0; i < 3; i++ { + err := r.Refresh(context.Background()) + if err == nil || err.Error() != "netproxy: refresh command failed: exit status 3" { + t.Fatalf("Refresh = %v, want the exit status", err) + } + } + if e := errs.all(); len(e) != 1 { + t.Fatalf("error handler called %d times for one run of failures, want 1: %v", len(e), e) + } + dialEcho(t, d, target, "last good credentials still used").Close() + if got := proxy.rejected.Load(); got != 0 { + t.Fatalf("proxy rejected %d CONNECTs, want 0", got) + } + + // The proxy rotates while the command is broken: the dial fails, and + // its error carries both the 407 and the refresh failure. + _, pass2 := creds.rotate() + _, err = d.DialContext(context.Background(), "tcp", target) + if err == nil { + t.Fatal("dial succeeded with a broken refresh command after a rotation") + } + var ce *ConnectError + if !errors.As(err, &ce) || ce.StatusCode != http.StatusProxyAuthRequired { + t.Fatalf("%v is not a 407 ConnectError", err) + } + if !strings.Contains(err.Error(), "407 Proxy Authentication Required (proxy credential refresh failed: netproxy: refresh command failed: exit status 3)") { + t.Fatalf("error = %q, want the 407 and the refresh failure", err) + } + if got := proxy.rejected.Load(); got != 1 { + t.Fatalf("proxy saw %d rejected CONNECTs, want 1 (no retry with unchanged credentials)", got) + } + if e := errs.all(); len(e) != 1 { + t.Fatalf("error handler called %d times, want still 1", len(e)) + } + + // The command recovers: the next rejected dial refreshes and succeeds. + creds.fail(false) + dialEcho(t, d, target, "recovered").Close() + if got := proxy.rejected.Load(); got != 2 { + t.Fatalf("proxy saw %d rejected CONNECTs, want 2", got) + } + + // A new run of failures is reported again. + creds.fail(true) + if err := r.Refresh(context.Background()); err == nil { + t.Fatal("Refresh succeeded with a failing command") + } + all := errs.all() + if len(all) != 2 { + t.Fatalf("error handler called %d times, want 2 (one per run of failures)", len(all)) + } + for _, e := range append(all, err) { + for _, secret := range []string{pass1, pass2, "tok/", "tok%2F"} { + if strings.Contains(e.Error(), secret) { + t.Fatalf("error %q leaks %q", e, secret) + } + } + } + if s := r.String(); s != "http://***@"+proxy.addr()+" (credentials refreshed by command)" { + t.Fatalf("String = %q", s) + } +} + +// user1 returns the first generation's credentials. +func (c *credSource) user1() (user, pass string) { return "muse-agent", "tok/1?r#s@t" } + +// Lookups start refreshes on the timer too. With a failing command every +// lookup keeps the last good URL, and the failure is still reported only +// once. +func TestRefreshIntervalKeepsLastGoodOnFailure(t *testing.T) { + t.Parallel() + echoIP, echoPort := newEchoServer(t) + proxy := newTestProxy(t, map[string]string{echoHost: echoIP}) + creds := newCredSource(t, proxy) + var errs errorLog + r, err := NewResolver(proxy.url("launch:time"), WithRefreshCommand(creds.command()), WithRefreshInterval(time.Nanosecond), errs.handler()) + if err != nil { + t.Fatal(err) + } + d := NewDialer(r) + target := net.JoinHostPort(echoHost, echoPort) + creds.fail(true) + for i := 0; i < 3; i++ { + dialEcho(t, d, target, "on the last good credentials").Close() + settle(t, r) // the timed refresh this dial started fails + } + if got := creds.runs(); got != 4 { + t.Fatalf("command ran %d times, want 4 (build + one per lookup)", got) + } + if e := errs.all(); len(e) != 1 { + t.Fatalf("error handler called %d times, want 1: %v", len(e), e) + } + + // Once the command works, a timed refresh picks up rotated credentials + // without waiting for a 407. + creds.fail(false) + creds.rotate() + if got := lookupAfterTimedRefresh(t, r, target); got != creds.current() { + t.Fatalf("after a timed refresh: proxy %q, want the rotated URL", got) + } + dialEcho(t, d, target, "rotated").Close() + if got := proxy.rejected.Load(); got != 0 { + t.Fatalf("proxy rejected %d CONNECTs, want 0: the timed refresh ran first", got) + } +} + +// Fifty dials hit a rotation at once. They share one refresh: the command +// runs exactly once, and every dial succeeds. +func TestRefreshConcurrentRejectionsRunOneCommand(t *testing.T) { + t.Parallel() + echoIP, echoPort := newEchoServer(t) + proxy := newTestProxy(t, map[string]string{echoHost: echoIP}) + creds := newCredSource(t, proxy) + creds.delay = "sleep 0.3;" // keep the refresh in flight while the others arrive + r, err := NewResolver(proxy.url("launch:time"), WithRefreshCommand(creds.command()), WithRefreshInterval(time.Hour)) + if err != nil { + t.Fatal(err) + } + d := NewDialer(r) + target := net.JoinHostPort(echoHost, echoPort) + creds.rotate() + + const n = 50 + var wg sync.WaitGroup + var failed atomic.Int32 + gate := make(chan struct{}) + for i := 0; i < n; i++ { + wg.Add(1) + go func() { + defer wg.Done() + <-gate + c, err := d.DialContext(context.Background(), "tcp", target) + if err != nil { + failed.Add(1) + t.Errorf("dial: %v", err) + return + } + defer c.Close() + msg := "concurrent" + c.SetDeadline(time.Now().Add(5 * time.Second)) + if _, err := c.Write([]byte(msg)); err != nil { + t.Errorf("write: %v", err) + return + } + buf := make([]byte, len(msg)) + if _, err := io.ReadFull(c, buf); err != nil || string(buf) != msg { + t.Errorf("read %q: %v", buf, err) + } + }() + } + close(gate) + wg.Wait() + if failed.Load() != 0 { + t.Fatalf("%d of %d dials failed", failed.Load(), n) + } + if got := creds.runs(); got != 2 { + t.Fatalf("command ran %d times, want 2 (build + exactly one refresh for the rotation)", got) + } + rejected := proxy.rejected.Load() + if rejected < 2 { + t.Fatalf("only %d dials hit the rotation; the test did not exercise concurrent rejections", rejected) + } + t.Logf("%d of %d dials were rejected and shared one refresh", rejected, n) +} + +// Without a command, ModeAuto re-reads the environment, which is enough for +// a process whose environment is updated in place. +func TestRefreshAutoModeRereadsEnvironment(t *testing.T) { + t.Parallel() + echoIP, echoPort := newEchoServer(t) + proxy := newTestProxy(t, map[string]string{echoHost: echoIP}, withAuth("u", "one")) + env := &syncEnv{m: map[string]string{"HTTPS_PROXY": proxy.url("u:one")}} + r, err := newAuto(env.get, options{}) + if err != nil { + t.Fatal(err) + } + d := NewDialer(r) + target := net.JoinHostPort(echoHost, echoPort) + first := dialEcho(t, d, target, "one") + defer first.Close() + + proxy.setAuth("u", "two") + env.set("HTTPS_PROXY", proxy.url("u:two")) + dialEcho(t, d, target, "two").Close() + roundTrip(t, first, "first tunnel untouched") + if got := proxy.rejected.Load(); got != 1 { + t.Fatalf("proxy rejected %d CONNECTs, want 1", got) + } + + // On the timer, a lookup follows the environment without any 407. + env2 := &syncEnv{m: map[string]string{"HTTPS_PROXY": "http://a.invalid:1"}} + r2, err := newAuto(env2.get, newOptions([]Option{WithRefreshInterval(time.Nanosecond)})) + if err != nil { + t.Fatal(err) + } + env2.set("HTTPS_PROXY", "http://b.invalid:2") + env2.set("NO_PROXY", "skip.invalid") + if got := lookupAfterTimedRefresh(t, r2, "x.invalid:443"); got != "http://b.invalid:2" { + t.Fatalf("after the environment changed: proxy %q", got) + } + if got := proxyFor(t, r2, "skip.invalid:443"); got != "" { + t.Fatalf("new NO_PROXY ignored: proxy %q", got) + } + // An environment without a proxy at first starts proxying once one is + // set. + env3 := &syncEnv{m: map[string]string{}} + r3, err := newAuto(env3.get, newOptions([]Option{WithRefreshInterval(time.Nanosecond)})) + if err != nil { + t.Fatal(err) + } + if r3.Enabled() || proxyFor(t, r3, "x.invalid:443") != "" || proxyForURL(t, r3, "https://x.invalid/") != "" { + t.Fatal("empty environment proxies") + } + env3.set("https_proxy", "http://late.invalid:3128") + if got := lookupAfterTimedRefresh(t, r3, "x.invalid:443"); got != "http://late.invalid:3128" { + t.Fatalf("after https_proxy was set: proxy %q", got) + } + if got := proxyForURL(t, r3, "https://x.invalid/"); got != "http://late.invalid:3128" { + t.Fatalf("after https_proxy was set: request proxy %q", got) + } + if !r3.Enabled() { + t.Fatal("Enabled still false after the environment gained a proxy") + } + + // An unusable HTTPS_PROXY is a failed refresh: the last reading stays. + env2.set("HTTPS_PROXY", "socks5://c.invalid:3") + if got := lookupAfterTimedRefresh(t, r2, "x.invalid:443"); got != "http://b.invalid:2" { + t.Fatalf("after an unusable HTTPS_PROXY: proxy %q", got) + } + var ee *EnvError + if err := r2.Refresh(context.Background()); !errors.As(err, &ee) || ee.Var != "HTTPS_PROXY" { + t.Fatalf("Refresh = %v, want an EnvError for HTTPS_PROXY", err) + } +} + +// FromEnvironment follows the real process environment. +func TestFromEnvironmentFollowsInProcessUpdates(t *testing.T) { + echoIP, echoPort := newEchoServer(t) + proxy := newTestProxy(t, map[string]string{echoHost: echoIP}, withAuth("u", "one")) + t.Setenv("HTTPS_PROXY", proxy.url("u:one")) + r, err := FromEnvironment() + if err != nil { + t.Fatal(err) + } + d := NewDialer(r) + target := net.JoinHostPort(echoHost, echoPort) + dialEcho(t, d, target, "one").Close() + proxy.setAuth("u", "two") + t.Setenv("HTTPS_PROXY", proxy.url("u:two")) + dialEcho(t, d, target, "two").Close() +} + +// When refreshed settings stop proxying the target, the retry dials it +// directly. +func TestRefreshRetryFollowsNewRouting(t *testing.T) { + t.Parallel() + echoIP, echoPort := newEchoServer(t) + proxy := newTestProxy(t, map[string]string{echoHost: echoIP}, withAuth("u", "one")) + env := &syncEnv{m: map[string]string{"HTTPS_PROXY": proxy.url("u:one")}} + r, err := newAuto(env.get, options{}) + if err != nil { + t.Fatal(err) + } + fwd := &recordingForward{hosts: map[string]string{echoHost: echoIP}} + d := &Dialer{Resolver: r, Forward: fwd.dial} + proxy.setAuth("u", "two") + env.set("NO_PROXY", echoHost) + target := net.JoinHostPort(echoHost, echoPort) + dialEcho(t, d, target, "direct after refresh").Close() + if got := fwd.seen(); len(got) != 2 || got[0] != proxy.addr() || got[1] != target { + t.Fatalf("dials = %q, want the proxy and then the target directly", got) + } +} + +// Some proxies' rejections arrive as a status line net/http cannot parse +// ("malformed HTTP status code"). That also triggers a refresh. +func TestDialerRefreshesOnMalformedRejection(t *testing.T) { + t.Parallel() + echoIP, echoPort := newEchoServer(t) + proxy := newTestProxy(t, map[string]string{echoHost: echoIP}, withRejectLine("HTTP/1.1 407Proxy Authentication Required")) + creds := newCredSource(t, proxy) + r, err := NewResolver(proxy.url("launch:time"), WithRefreshCommand(creds.command()), WithRefreshInterval(time.Hour)) + if err != nil { + t.Fatal(err) + } + d := NewDialer(r) + target := net.JoinHostPort(echoHost, echoPort) + creds.rotate() + dialEcho(t, d, target, "after a garbled 407").Close() + if got, runs := proxy.rejected.Load(), creds.runs(); got != 1 || runs != 2 { + t.Fatalf("rejected %d, command runs %d; want 1 and 2", got, runs) + } + + // Without anything to refresh, the error is returned, naming what was + // wrong with the response but not quoting it. + _, err = NewDialer(mustExplicit(t, proxy.url("u:wrong"))).Dial("tcp", target) + if err == nil || !strings.HasSuffix(err.Error(), "read CONNECT response: malformed HTTP status code (response text withheld)") || strings.Contains(err.Error(), "407Proxy") { + t.Fatalf("error = %v", err) + } +} + +// A refresh source's output carries credentials: nothing it prints, and +// nothing a failing command writes to stderr, reaches an error or String. +func TestRefreshOutputNeverLeaks(t *testing.T) { + t.Parallel() + for _, out := range []string{ + "socks5://user:SECRET@proxy.invalid:1080", + "SECRET", + "user:SECRET@proxy.invalid:3128", // no scheme: rejected for refresh output + "http://user:SECRET@", + "http://user:SECRET@proxy.invalid:3128\nSECRET-line-two", + " \n", + } { + var errs errorLog + r, err := NewResolver("http://initial.invalid:3128", WithRefreshFunc(func(context.Context) (string, error) { return out, nil }), errs.handler()) + if err != nil { + t.Fatal(err) + } + rerr := r.Refresh(context.Background()) + if rerr == nil { + t.Fatalf("output %q accepted", out) + } + all := errs.all() + if len(all) != 1 { + t.Fatalf("output %q: handler called %d times", out, len(all)) + } + for _, s := range []string{rerr.Error(), all[0].Error(), r.String()} { + if strings.Contains(s, "SECRET") { + t.Fatalf("output %q leaked into %q", out, s) + } + } + if got := proxyFor(t, r, "x.invalid:443"); got != "http://initial.invalid:3128" { + t.Fatalf("output %q: proxy %q, want the initial URL kept", out, got) + } + } + + if runtime.GOOS == "windows" { + return + } + for _, cmd := range []string{ + `printf %s 'socks5://user:SECRET@proxy.invalid:1'`, + `echo 'http://user:SECRET@proxy.invalid:1' >&2; exit 1`, + `echo 'http://user:SECRET@proxy.invalid:1'; kill -9 $$`, + } { + var errs errorLog + r, err := NewResolver("http://initial.invalid:3128", WithRefreshCommand(cmd), errs.handler()) + if err != nil { + t.Fatal(err) + } + all := errs.all() + if len(all) != 1 { + t.Fatalf("command %q: handler called %d times", cmd, len(all)) + } + if strings.Contains(all[0].Error(), "SECRET") || strings.Contains(r.String(), "SECRET") { + t.Fatalf("command %q leaked: %q / %q", cmd, all[0], r) + } + t.Logf("command %q: %v", cmd, all[0]) + } +} + +func TestRefreshCommandTimeout(t *testing.T) { + if runtime.GOOS == "windows" { + t.Skip("refresh commands run with sh -c") + } + defer func(d time.Duration) { refreshTimeout = d }(refreshTimeout) + refreshTimeout = 300 * time.Millisecond + + var errs errorLog + start := time.Now() + r, err := NewResolver("http://initial.invalid:3128", WithRefreshCommand("sleep 5; echo http://late.invalid:1"), errs.handler()) + if err != nil { + t.Fatal(err) + } + if elapsed := time.Since(start); elapsed > 3*time.Second { + t.Fatalf("a hung command held NewResolver for %v", elapsed) + } + all := errs.all() + if len(all) != 1 || all[0].Error() != "netproxy: refresh command timed out after 300ms" { + t.Fatalf("errors = %v", all) + } + if got := proxyFor(t, r, "x.invalid:443"); got != "http://initial.invalid:3128" { + t.Fatalf("proxy %q, want the initial URL", got) + } +} + +func TestNewResolverOptions(t *testing.T) { + t.Parallel() + var calls atomic.Int32 + source := func(u string) Option { + return WithRefreshFunc(func(context.Context) (string, error) { + calls.Add(1) + return u, nil + }) + } + + // "off" ignores options; the source never runs. + off, err := NewResolver("off", source("http://x.invalid:1")) + if err != nil || off.Mode() != ModeOff || off.Enabled() || calls.Load() != 0 { + t.Fatalf("off: %v %v, source ran %d times", off, err, calls.Load()) + } + + // An empty command is no refresh source at all. + plain, err := NewResolver("http://u:p@proxy.invalid:1", WithRefreshCommand(" "), WithRefreshFunc(nil)) + if err != nil { + t.Fatal(err) + } + if plain.refreshable() || plain.String() != "http://***@proxy.invalid:1" { + t.Fatalf("empty command: refreshable=%v String=%q", plain.refreshable(), plain.String()) + } + for _, r := range []*Resolver{nil, Off(), plain, mustExplicit(t, "http://p.invalid:1")} { + if err := r.Refresh(context.Background()); err != nil { + t.Fatalf("Refresh on %v = %v", r, err) + } + } + + // Explicit with a source: the source's URL replaces the given one. + ex, err := NewResolver("http://old:pw@proxy.invalid:1", source("https://new:pw@proxy2.invalid:8443")) + if err != nil { + t.Fatal(err) + } + if ex.Mode() != ModeExplicit || proxyFor(t, ex, "localhost:1") != "https://new:pw@proxy2.invalid:8443" { + t.Fatalf("explicit+source: mode %q proxy %q", ex.Mode(), proxyFor(t, ex, "localhost:1")) + } + if s := ex.String(); s != "https://***@proxy2.invalid:8443 (credentials refreshed by callback)" { + t.Fatalf("String = %q", s) + } + + // Auto with a source: its URL replaces the environment's, NO_PROXY and + // the loopback exemption still apply. + auto, err := newAuto(envMap(map[string]string{"HTTPS_PROXY": "http://env.invalid:1", "HTTP_PROXY": "http://envplain.invalid:1", "NO_PROXY": "skip.invalid"}), + newOptions([]Option{source("http://u:pw@cmd.invalid:3128")})) + if err != nil { + t.Fatal(err) + } + for addr, want := range map[string]string{ + "x.invalid:443": "http://u:pw@cmd.invalid:3128", + "skip.invalid:443": "", + "127.0.0.1:443": "", + } { + if got := proxyFor(t, auto, addr); got != want { + t.Fatalf("auto+source %s: proxy %q, want %q", addr, got, want) + } + } + if got := proxyForURL(t, auto, "http://x.invalid/"); got != "http://u:pw@cmd.invalid:3128" { + t.Fatalf("auto+source plain request: proxy %q", got) + } + if s := auto.String(); s != "auto: http://***@cmd.invalid:3128 (NO_PROXY=skip.invalid) (credentials refreshed by callback)" { + t.Fatalf("String = %q", s) + } + + // Auto with no proxy in the environment and a failing source: enabled + // (the source may supply one later), dialing directly for now. + var errs errorLog + empty, err := newAuto(envMap(nil), newOptions([]Option{ + WithRefreshFunc(func(context.Context) (string, error) { return "", errors.New("not yet") }), + errs.handler(), + })) + if err != nil { + t.Fatal(err) + } + if !empty.Enabled() || proxyFor(t, empty, "x.invalid:443") != "" { + t.Fatalf("auto without a proxy yet: enabled=%v", empty.Enabled()) + } + if s := empty.String(); s != "auto: no proxy in environment (credentials refreshed by callback)" { + t.Fatalf("String = %q", s) + } + if e := errs.all(); len(e) != 1 || e[0].Error() != "netproxy: proxy refresh: not yet" { + t.Fatalf("errors = %v", e) + } + + // Timed refreshes: within the interval lookups reuse the last reading; + // after it, a lookup starts a refresh, whose result later lookups get. + // A negative interval never refreshes. + var n atomic.Int32 + counting := WithRefreshFunc(func(context.Context) (string, error) { + return fmt.Sprintf("http://p%d.invalid:1", n.Add(1)), nil + }) + timed, err := NewResolver("http://p0.invalid:1", counting, WithRefreshInterval(200*time.Millisecond)) + if err != nil { + t.Fatal(err) + } + for i := 0; i < 5; i++ { + if got := proxyFor(t, timed, "x.invalid:443"); got != "http://p1.invalid:1" { + t.Fatalf("within the interval: proxy %q", got) + } + } + time.Sleep(250 * time.Millisecond) + proxyFor(t, timed, "x.invalid:443") // p1, or p2 if the refresh won the race + settle(t, timed) + if got := proxyFor(t, timed, "x.invalid:443"); got != "http://p2.invalid:1" { + t.Fatalf("after the interval: proxy %q", got) + } + n.Store(0) + never, err := NewResolver("http://p0.invalid:1", counting, WithRefreshInterval(-1)) + if err != nil { + t.Fatal(err) + } + time.Sleep(10 * time.Millisecond) + for i := 0; i < 5; i++ { + proxyFor(t, never, "x.invalid:443") + } + if n.Load() != 1 { + t.Fatalf("negative interval: source ran %d times, want 1 (at build)", n.Load()) + } +} + +// A caller that stops waiting does not cancel the refresh: it completes +// and its result is used. +func TestRefreshWaitHonoursContext(t *testing.T) { + t.Parallel() + release := make(chan struct{}) + var n atomic.Int32 + r, err := NewResolver("http://p0.invalid:1", WithRefreshFunc(func(ctx context.Context) (string, error) { + i := n.Add(1) + if i > 1 { + <-release + } + return fmt.Sprintf("http://p%d.invalid:1", i), nil + }), WithRefreshInterval(-1)) + if err != nil { + t.Fatal(err) + } + ctx, cancel := context.WithTimeout(context.Background(), 50*time.Millisecond) + defer cancel() + if err := r.Refresh(ctx); !errors.Is(err, context.DeadlineExceeded) { + t.Fatalf("Refresh = %v, want the context's deadline", err) + } + if got := proxyFor(t, r, "x.invalid:443"); got != "http://p1.invalid:1" { + t.Fatalf("while the refresh is in flight: proxy %q", got) + } + close(release) + waitFor(t, "the abandoned refresh to land", func() bool { + u, _ := r.ProxyForAddr("x.invalid:443") + return u != nil && u.String() == "http://p2.invalid:1" + }) + if n.Load() != 2 { + t.Fatalf("source ran %d times, want 2", n.Load()) + } +} + +// A timed refresh runs in the background. While one hangs, dials with a +// budget far shorter than the refresh timeout (the registry client allows +// 5 s a dial) go out at once on the cached credentials, which still work, +// and so do HTTP proxy lookups. Only a 407 waits for the refresh. +func TestTimedRefreshNeverBlocksALookup(t *testing.T) { + t.Parallel() + echoIP, echoPort := newEchoServer(t) + proxy := newTestProxy(t, map[string]string{echoHost: echoIP}) + creds := newCredSource(t, proxy) + urlFile := filepath.Join(creds.dir, "url") + release := make(chan struct{}) + var runs atomic.Int32 + r, err := NewResolver(proxy.url("launch:time"), WithRefreshFunc(func(ctx context.Context) (string, error) { + if runs.Add(1) > 1 { + select { + case <-release: + case <-ctx.Done(): + return "", ctx.Err() + } + } + b, err := os.ReadFile(urlFile) + return string(b), err + }), WithRefreshInterval(20*time.Millisecond)) + if err != nil { + t.Fatal(err) + } + d := NewDialer(r) + target := net.JoinHostPort(echoHost, echoPort) + time.Sleep(30 * time.Millisecond) // a timed refresh is due + + for i := 0; i < 3; i++ { + ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second) + begin := time.Now() + c, err := d.DialContext(ctx, "tcp", target) + cancel() + if err != nil { + t.Fatalf("dial %d while a timed refresh hangs: %v", i, err) + } + roundTrip(t, c, "on the cached credentials") + c.Close() + if took := time.Since(begin); took > time.Second { + t.Fatalf("dial %d took %v: it waited for the timed refresh", i, took) + } + begin = time.Now() + if got := proxyForURL(t, r, "https://api.pilot.invalid/"); got == "" || time.Since(begin) > time.Second { + t.Fatalf("request lookup %d: proxy %q after %v", i, got, time.Since(begin)) + } + } + if got := runs.Load(); got != 2 { + t.Fatalf("refresh source ran %d times, want 2: build, then one timed refresh that no lookup repeats while it runs", got) + } + if got := proxy.rejected.Load(); got != 0 { + t.Fatalf("proxy rejected %d CONNECTs, want 0", got) + } + + // The proxy rotates while the refresh still hangs: the rejected dial + // waits for that refresh, which reads the new credentials once it is + // released, and retries with them. + creds.rotate() + time.AfterFunc(100*time.Millisecond, func() { close(release) }) + dialEcho(t, d, target, "after the rotation").Close() + if got := proxy.rejected.Load(); got != 1 { + t.Fatalf("proxy rejected %d CONNECTs, want 1", got) + } + if got := runs.Load(); got != 2 { + t.Fatalf("refresh source ran %d times, want still 2: the 407 joined the refresh in flight", got) + } +} + +// A refresh command that times out is killed together with everything it +// started. So is whatever a finished command leaves behind holding its +// output. Before, only sh was killed, and every hung refresh left its +// children running. +func TestRefreshCommandKillsWhatItStarted(t *testing.T) { + if runtime.GOOS == "windows" { + t.Skip("refresh commands run with sh -c and process groups") + } + defer func(d time.Duration) { refreshTimeout = d }(refreshTimeout) + + dir := t.TempDir() + // beat.sh appends to $1 every 50 ms for as long as it lives. + beat := filepath.Join(dir, "beat.sh") + script := "echo $$ >> " + shellQuote(filepath.Join(dir, "pids")) + "\nwhile :; do echo . >> \"$1\"; sleep 0.05; done\n" + if err := os.WriteFile(beat, []byte(script), 0o600); err != nil { + t.Fatal(err) + } + t.Cleanup(func() { + // Only matters if the test fails: do not leave the beaters behind. + b, _ := os.ReadFile(filepath.Join(dir, "pids")) + for _, f := range strings.Fields(string(b)) { + var pid int + if _, err := fmt.Sscan(f, &pid); err == nil && pid > 1 { + if p, err := os.FindProcess(pid); err == nil { + p.Kill() + } + } + } + }) + size := func(name string) int64 { + fi, err := os.Stat(filepath.Join(dir, name)) + if err != nil { + t.Fatalf("%s never started: %v", name, err) + } + return fi.Size() + } + assertStopped := func(names ...string) { + t.Helper() + time.Sleep(100 * time.Millisecond) // a write in progress lands + before := make([]int64, len(names)) + for i, n := range names { + before[i] = size(n) + } + time.Sleep(400 * time.Millisecond) + for i, n := range names { + if after := size(n); after != before[i] { + t.Fatalf("%s is still running after the refresh returned (%d -> %d bytes)", n, before[i], after) + } + } + } + + // Timeout: a background job and a foreground child that never + // finishes. Neither is sh itself. + refreshTimeout = 500 * time.Millisecond + cmd := fmt.Sprintf("sh %[1]s %[2]s & sh %[1]s %[3]s; echo http://late.invalid:1", + shellQuote(beat), shellQuote(filepath.Join(dir, "bg")), shellQuote(filepath.Join(dir, "fg"))) + var errs errorLog + if _, err := NewResolver("http://initial.invalid:3128", WithRefreshCommand(cmd), errs.handler()); err != nil { + t.Fatal(err) + } + if e := errs.all(); len(e) != 1 || e[0].Error() != "netproxy: refresh command timed out after 500ms" { + t.Fatalf("errors = %v", e) + } + assertStopped("bg", "fg") + + // A command that prints its URL and exits, leaving a background job + // holding its stdout: the refresh fails, and the job is ended too. + refreshTimeout = 5 * time.Second + cmd = fmt.Sprintf("sh %s %s & echo http://late.invalid:1", shellQuote(beat), shellQuote(filepath.Join(dir, "held"))) + var errs2 errorLog + if _, err := NewResolver("http://initial.invalid:3128", WithRefreshCommand(cmd), errs2.handler()); err != nil { + t.Fatal(err) + } + if e := errs2.all(); len(e) != 1 || e[0].Error() != "netproxy: refresh command left a process holding its output open" { + t.Fatalf("errors = %v", e) + } + assertStopped("held") +} diff --git a/netproxy/zz_transport_test.go b/netproxy/zz_transport_test.go new file mode 100644 index 0000000..be4b107 --- /dev/null +++ b/netproxy/zz_transport_test.go @@ -0,0 +1,517 @@ +// SPDX-License-Identifier: AGPL-3.0-or-later + +package netproxy + +import ( + "bufio" + "context" + "crypto/tls" + "crypto/x509" + "errors" + "io" + "log" + "net" + "net/http" + "net/http/httptest" + "net/url" + "strings" + "sync" + "sync/atomic" + "testing" + "time" +) + +// apiServer is an HTTPS server reachable only as apiHost through the test +// proxy. It answers " " and records what it received. +type apiServer struct { + ip, port string + pool *x509.CertPool + + mu sync.Mutex + seen []string // " " per request +} + +func newAPIServer(t *testing.T) *apiServer { + t.Helper() + cert, pool := genCert(t, []string{apiHost}, nil) + s := &apiServer{pool: pool} + srv := httptest.NewUnstartedServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + body, _ := io.ReadAll(r.Body) + line := r.Method + " " + string(body) + s.mu.Lock() + s.seen = append(s.seen, line) + s.mu.Unlock() + io.WriteString(w, line) + })) + srv.TLS = &tls.Config{Certificates: []tls.Certificate{cert}} + // A tunnel the client abandons before TLS (a vetoed CONNECT) makes the + // server log a handshake EOF; it is expected noise. + srv.Config.ErrorLog = log.New(io.Discard, "", 0) + srv.StartTLS() + t.Cleanup(srv.Close) + s.ip, s.port, _ = net.SplitHostPort(srv.Listener.Addr().String()) + return s +} + +func (s *apiServer) url(path string) string { + return "https://" + net.JoinHostPort(apiHost, s.port) + path +} + +func (s *apiServer) requests() []string { + s.mu.Lock() + defer s.mu.Unlock() + return append([]string(nil), s.seen...) +} + +func (s *apiServer) base() *http.Transport { + return &http.Transport{TLSClientConfig: &tls.Config{RootCAs: s.pool}} +} + +func do(t *testing.T, c *http.Client, method, rawURL string, body io.Reader) (string, error) { + t.Helper() + req, err := http.NewRequest(method, rawURL, body) + if err != nil { + t.Fatal(err) + } + resp, err := c.Do(req) + if err != nil { + return "", err + } + defer resp.Body.Close() + b, err := io.ReadAll(resp.Body) + if err != nil { + t.Fatalf("read body: %v", err) + } + if resp.StatusCode != http.StatusOK { + t.Fatalf("%s %s: status %d", method, rawURL, resp.StatusCode) + } + return string(b), nil +} + +// Pins what net/http itself reports when a proxy refuses CONNECT: only the +// proxy's reason phrase as a plain error — no status code, "unknown status +// code" without a phrase, the proxy's own words otherwise, and "malformed +// HTTP status code" for a status line it cannot parse. Matching on that +// text cannot tell a 407 apart, which is why RefreshingTransport reads the +// status from OnProxyConnectResponse instead. +func TestNetHTTPReportsRefusedCONNECTWithoutStatusCode(t *testing.T) { + t.Parallel() + for _, tc := range []struct { + line string + want string + }{ + {"HTTP/1.1 407 Proxy Authentication Required", "Proxy Authentication Required"}, + {"HTTP/1.1 407", "unknown status code"}, + {"HTTP/1.1 407 Denied muse-agent:hunter2", "Denied muse-agent:hunter2"}, + {"HTTP/1.1 407Proxy Authentication Required", `malformed HTTP status code "407Proxy"`}, + } { + proxy := newTestProxy(t, nil, withAuth("u", "p"), withRejectLine(tc.line)) + pu, _ := url.Parse(proxy.url("u:stale")) + tr := &http.Transport{Proxy: http.ProxyURL(pu)} + req, _ := http.NewRequest(http.MethodGet, "https://api.pilot.invalid/", nil) + _, err := tr.RoundTrip(req) + if err == nil || err.Error() != tc.want { + t.Fatalf("%q: net/http error = %v, want %q", tc.line, err, tc.want) + } + var ce *ConnectError + if errors.As(err, &ce) { + t.Fatalf("%q: plain net/http error unexpectedly typed", tc.line) + } + tr.CloseIdleConnections() + } +} + +// A GET survives a rotation: a pooled tunnel keeps working without a new +// CONNECT, and the next new connection is refused, refreshed and retried. +func TestRefreshingTransportRetriesGETAfterRotation(t *testing.T) { + t.Parallel() + api := newAPIServer(t) + proxy := newTestProxy(t, map[string]string{apiHost: api.ip}) + creds := newCredSource(t, proxy) + r, err := NewResolver(proxy.url("launch:time"), WithRefreshCommand(creds.command()), WithRefreshInterval(time.Hour)) + if err != nil { + t.Fatal(err) + } + client := &http.Client{Transport: RefreshingTransport(api.base(), r), Timeout: 10 * time.Second} + defer client.CloseIdleConnections() + + if got, err := do(t, client, http.MethodGet, api.url("/1"), nil); err != nil || got != "GET " { + t.Fatalf("GET 1 = %q, %v", got, err) + } + creds.rotate() + // The idle tunnel from GET 1 was opened with the old credentials and is + // still up: no new CONNECT, no refresh. + if got, err := do(t, client, http.MethodGet, api.url("/2"), nil); err != nil || got != "GET " { + t.Fatalf("GET 2 on the pooled tunnel = %q, %v", got, err) + } + if targets, _, _ := proxy.seen(); len(targets) != 1 || creds.runs() != 1 { + t.Fatalf("pooled request opened %d CONNECTs, ran the command %d times", len(targets), creds.runs()) + } + + client.CloseIdleConnections() + if got, err := do(t, client, http.MethodGet, api.url("/3"), nil); err != nil || got != "GET " { + t.Fatalf("GET 3 after a rotation = %q, %v", got, err) + } + if got := proxy.rejected.Load(); got != 1 { + t.Fatalf("proxy rejected %d CONNECTs, want 1", got) + } + if got := creds.runs(); got != 2 { + t.Fatalf("command ran %d times, want 2", got) + } + if got := api.requests(); len(got) != 3 { + t.Fatalf("server saw %d requests, want 3: %q", len(got), got) + } +} + +// A POST whose body cannot be rewound is not retried, but the refresh still +// happens, so the next request goes out with fresh credentials. With +// GetBody, the POST is retried: the proxy refused the tunnel, so the server +// never saw the first attempt, and it sees the body exactly once. +func TestRefreshingTransportPOST(t *testing.T) { + t.Parallel() + api := newAPIServer(t) + proxy := newTestProxy(t, map[string]string{apiHost: api.ip}) + creds := newCredSource(t, proxy) + r, err := NewResolver(proxy.url("launch:time"), WithRefreshCommand(creds.command()), WithRefreshInterval(time.Hour)) + if err != nil { + t.Fatal(err) + } + base := api.base() + base.DisableKeepAlives = true // every request opens a new tunnel + client := &http.Client{Transport: RefreshingTransport(base, r), Timeout: 10 * time.Second} + + creds.rotate() + oneShot := io.NopCloser(strings.NewReader("not rewindable")) // no GetBody + _, err = do(t, client, http.MethodPost, api.url("/post"), oneShot) + var ce *ConnectError + if !errors.As(err, &ce) || ce.StatusCode != http.StatusProxyAuthRequired { + t.Fatalf("POST without GetBody: error = %v, want a 407 ConnectError", err) + } + if got := proxy.rejected.Load(); got != 1 { + t.Fatalf("POST without GetBody was retried: %d rejected CONNECTs", got) + } + if got := api.requests(); len(got) != 0 { + t.Fatalf("server saw %q", got) + } + if got := creds.runs(); got != 2 { + t.Fatalf("command ran %d times, want 2: the refresh happens even without a retry", got) + } + if got, err := do(t, client, http.MethodGet, api.url("/after"), nil); err != nil || got != "GET " { + t.Fatalf("GET after the refused POST = %q, %v", got, err) + } + if got, runs := proxy.rejected.Load(), creds.runs(); got != 1 || runs != 2 { + t.Fatalf("GET after the refresh: rejected %d, runs %d; want 1 and 2", got, runs) + } + + creds.rotate() + got, err := do(t, client, http.MethodPost, api.url("/post"), strings.NewReader("rewindable")) // GetBody set + if err != nil || got != "POST rewindable" { + t.Fatalf("POST with GetBody = %q, %v", got, err) + } + if n := proxy.rejected.Load(); n != 2 { + t.Fatalf("rejected %d CONNECTs, want 2", n) + } + posts := 0 + for _, line := range api.requests() { + if strings.HasPrefix(line, "POST") { + posts++ + } + } + if posts != 1 { + t.Fatalf("server saw %d POSTs, want exactly 1: %q", posts, api.requests()) + } +} + +// Every refused CONNECT becomes a *ConnectError whose message carries none +// of the proxy's text; statuses other than 407 are not retried. +func TestRefreshingTransportConnectErrorsAreTypedAndRedacted(t *testing.T) { + t.Parallel() + for _, status := range []int{http.StatusProxyAuthRequired, http.StatusForbidden} { + proxy := newTestProxy(t, nil, withReject(status, "Denied muse-agent:hunter2")) + r, err := fromEnv(envMap(map[string]string{"HTTPS_PROXY": proxy.url("muse-agent:hunter2")})) + if err != nil { + t.Fatal(err) + } + client := &http.Client{Transport: RefreshingTransport(nil, r), Timeout: 10 * time.Second} + _, err = do(t, client, http.MethodGet, "https://api.pilot.invalid/", nil) + var ce *ConnectError + if !errors.As(err, &ce) || ce.StatusCode != status || ce.Target != "api.pilot.invalid:443" { + t.Fatalf("status %d: error = %v", status, err) + } + if strings.Contains(err.Error(), "hunter2") || strings.Contains(err.Error(), "Denied") { + t.Fatalf("status %d: error leaks proxy text: %q", status, err) + } + // 407: the environment was re-read, unchanged, so no retry. + if targets, _, _ := proxy.seen(); len(targets) != 1 { + t.Fatalf("status %d: %d CONNECTs, want 1", status, len(targets)) + } + client.CloseIdleConnections() + } +} + +// A garbled rejection also refreshes; the GET is retried, a POST (even +// with GetBody) is not, since such an error could have come from the +// server. +func TestRefreshingTransportMalformedRejection(t *testing.T) { + t.Parallel() + api := newAPIServer(t) + proxy := newTestProxy(t, map[string]string{apiHost: api.ip}, withRejectLine("HTTP/1.1 407Proxy Authentication Required")) + creds := newCredSource(t, proxy) + r, err := NewResolver(proxy.url("launch:time"), WithRefreshCommand(creds.command()), WithRefreshInterval(time.Hour)) + if err != nil { + t.Fatal(err) + } + base := api.base() + base.DisableKeepAlives = true + client := &http.Client{Transport: RefreshingTransport(base, r), Timeout: 10 * time.Second} + + creds.rotate() + if got, err := do(t, client, http.MethodGet, api.url("/"), nil); err != nil || got != "GET " { + t.Fatalf("GET after a garbled 407 = %q, %v", got, err) + } + creds.rotate() + _, err = do(t, client, http.MethodPost, api.url("/"), strings.NewReader("x")) + if err == nil || !strings.Contains(err.Error(), "proxy CONNECT "+net.JoinHostPort(apiHost, api.port)+": read CONNECT response: malformed HTTP status code (response text withheld)") { + t.Fatalf("POST after a garbled 407: error = %v", err) + } + if got, runs := proxy.rejected.Load(), creds.runs(); got != 2 || runs != 3 { + t.Fatalf("rejected %d, runs %d; want 2 and 3 (POST refreshed, not retried)", got, runs) + } +} + +// base's own OnProxyConnectResponse runs first and can veto the tunnel. +func TestRefreshingTransportChainsBaseHook(t *testing.T) { + t.Parallel() + api := newAPIServer(t) + proxy := newTestProxy(t, map[string]string{apiHost: api.ip}) + base := api.base() + var hooked []int + var mu sync.Mutex + veto := errors.New("vetoed by base hook") + base.OnProxyConnectResponse = func(_ context.Context, _ *url.URL, _ *http.Request, res *http.Response) error { + mu.Lock() + defer mu.Unlock() + hooked = append(hooked, res.StatusCode) + if len(hooked) > 1 { + return veto + } + return nil + } + client := &http.Client{Transport: RefreshingTransport(base, mustExplicit(t, proxy.url(""))), Timeout: 10 * time.Second} + if _, err := do(t, client, http.MethodGet, api.url("/"), nil); err != nil { + t.Fatalf("GET: %v", err) + } + client.CloseIdleConnections() // reaches the transport's copy of base + if _, err := do(t, client, http.MethodGet, api.url("/"), nil); !errors.Is(err, veto) { + t.Fatalf("GET 2: error = %v, want the base hook's", err) + } + mu.Lock() + defer mu.Unlock() + if len(hooked) != 2 || hooked[0] != http.StatusOK { + t.Fatalf("base hook saw %v", hooked) + } +} + +func TestCanReplay(t *testing.T) { + t.Parallel() + req := func(method string, body io.Reader, hdr ...string) *http.Request { + r, _ := http.NewRequest(method, "https://x.invalid/", body) + for _, h := range hdr { + r.Header.Set(h, "k") + } + return r + } + oneShot := func() io.Reader { return io.NopCloser(strings.NewReader("b")) } + for _, tc := range []struct { + name string + req *http.Request + neverSent bool + want bool + }{ + {"GET", req(http.MethodGet, nil), false, true}, + {"HEAD", req(http.MethodHead, nil), false, true}, + {"OPTIONS", req(http.MethodOptions, nil), false, true}, + {"GET with one-shot body", req(http.MethodGet, oneShot()), true, false}, + {"POST no body", req(http.MethodPost, nil), true, false}, + {"POST one-shot body", req(http.MethodPost, oneShot()), true, false}, + {"POST GetBody, refused tunnel", req(http.MethodPost, strings.NewReader("b")), true, true}, + {"POST GetBody, maybe sent", req(http.MethodPost, strings.NewReader("b")), false, false}, + {"POST Idempotency-Key", req(http.MethodPost, strings.NewReader("b"), "Idempotency-Key"), false, true}, + {"PUT X-Idempotency-Key no body", req(http.MethodPut, nil, "X-Idempotency-Key"), false, true}, + {"DELETE", req(http.MethodDelete, nil), false, false}, + } { + if got := canReplay(tc.req, tc.neverSent); got != tc.want { + t.Errorf("%s: canReplay = %v, want %v", tc.name, got, tc.want) + } + } +} + +// Only the proxy's answer to CONNECT can trigger a refresh. A server that +// answers through the tunnel with a response net/http cannot parse gets its +// error passed on untouched, and never makes the transport run the refresh +// command (which it used to do on every request). +func TestRefreshingTransportIgnoresMalformedServerResponses(t *testing.T) { + t.Parallel() + cert, pool := genCert(t, []string{apiHost}, nil) + ln, err := tls.Listen("tcp", "127.0.0.1:0", &tls.Config{Certificates: []tls.Certificate{cert}, MinVersion: tls.VersionTLS12}) + if err != nil { + t.Fatal(err) + } + var wg sync.WaitGroup + wg.Add(1) + go func() { + defer wg.Done() + for { + c, err := ln.Accept() + if err != nil { + return + } + wg.Add(1) + go func() { + defer wg.Done() + defer c.Close() + c.SetDeadline(time.Now().Add(5 * time.Second)) + if _, err := http.ReadRequest(bufio.NewReader(c)); err != nil { + return + } + io.WriteString(c, "HTTP/1.1 200OK\r\nContent-Length: 0\r\n\r\n") + }() + } + }() + t.Cleanup(func() { ln.Close(); wg.Wait() }) + ip, port, _ := net.SplitHostPort(ln.Addr().String()) + + proxy := newTestProxy(t, map[string]string{apiHost: ip}) + creds := newCredSource(t, proxy) + r, err := NewResolver(creds.current(), WithRefreshCommand(creds.command()), WithRefreshInterval(time.Hour)) + if err != nil { + t.Fatal(err) + } + client := &http.Client{Transport: RefreshingTransport(&http.Transport{TLSClientConfig: &tls.Config{RootCAs: pool}}, r), Timeout: 10 * time.Second} + defer client.CloseIdleConnections() + const n = 5 + for i := 0; i < n; i++ { + _, err := client.Get("https://" + net.JoinHostPort(apiHost, port) + "/") + if err == nil || !strings.Contains(err.Error(), `malformed HTTP status code "200OK"`) || strings.Contains(err.Error(), "withheld") { + t.Fatalf("GET %d: error = %v, want net/http's own", i, err) + } + } + if got := creds.runs(); got != 1 { + t.Fatalf("%d malformed server responses ran the refresh command %d times, want 0", n, got-1) + } + if got := proxy.rejected.Load(); got != 0 { + t.Fatalf("proxy rejected %d CONNECTs, want 0", got) + } +} + +// countingServer is an HTTPS server, reachable as apiHost through the test +// proxy, that counts its open connections. +type countingServer struct { + ip, port string + pool *x509.CertPool + open atomic.Int32 +} + +func newCountingServer(t *testing.T) *countingServer { + t.Helper() + cert, pool := genCert(t, []string{apiHost}, nil) + s := &countingServer{pool: pool} + srv := httptest.NewUnstartedServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { io.WriteString(w, "ok") })) + srv.TLS = &tls.Config{Certificates: []tls.Certificate{cert}} + srv.Config.ErrorLog = log.New(io.Discard, "", 0) + srv.Config.ConnState = func(_ net.Conn, state http.ConnState) { + switch state { + case http.StateNew: + s.open.Add(1) + case http.StateClosed, http.StateHijacked: + s.open.Add(-1) + } + } + srv.StartTLS() + t.Cleanup(srv.Close) + s.ip, s.port, _ = net.SplitHostPort(srv.Listener.Addr().String()) + return s +} + +// net/http pools tunnels by proxy URL, credentials included, so after a +// refresh changes the credentials the idle tunnels opened with the old ones +// can never be picked again. The next request closes them rather than +// leaving one open per rotation (for good, with IdleConnTimeout unset). A +// refresh that reads the same settings again leaves the pool alone. +func TestRefreshingTransportDropsTunnelsStrandedByARotation(t *testing.T) { + t.Parallel() + srv := newCountingServer(t) + proxy := newTestProxy(t, map[string]string{apiHost: srv.ip}) + creds := newCredSource(t, proxy) + r, err := NewResolver(creds.current(), WithRefreshCommand(creds.command()), WithRefreshInterval(time.Hour)) + if err != nil { + t.Fatal(err) + } + base := &http.Transport{TLSClientConfig: &tls.Config{RootCAs: srv.pool}} // no IdleConnTimeout + client := &http.Client{Transport: RefreshingTransport(base, r), Timeout: 10 * time.Second} + defer client.CloseIdleConnections() + get := func() { + t.Helper() + resp, err := client.Get("https://" + net.JoinHostPort(apiHost, srv.port) + "/") + if err != nil { + t.Fatal(err) + } + io.Copy(io.Discard, resp.Body) + resp.Body.Close() + } + + for i := 0; i < 3; i++ { + get() + } + if targets, _, _ := proxy.seen(); len(targets) != 1 { + t.Fatalf("3 GETs opened %d tunnels, want 1 (reused)", len(targets)) + } + + const rotations = 5 + for i := 0; i < rotations; i++ { + creds.rotate() + for j := 0; j < 2; j++ { // the second reads the same credentials again + if err := r.Refresh(context.Background()); err != nil { + t.Fatal(err) + } + } + get() + get() + } + if targets, _, _ := proxy.seen(); len(targets) != 1+rotations { + t.Fatalf("opened %d tunnels, want %d: one per rotation, each then reused", len(targets), 1+rotations) + } + if got := proxy.rejected.Load(); got != 0 { + t.Fatalf("proxy rejected %d CONNECTs, want 0", got) + } + waitFor(t, "the tunnels opened with old credentials to close", func() bool { return srv.open.Load() == 1 }) +} + +// A rejection whose status line or headers net/http cannot parse is +// reported by what was wrong with it, never by its text, which the proxy +// can fill with anything, the credentials it was sent included; a 407's +// reason phrase has always been withheld for the same reason. +func TestUnparseableRejectionNeverQuotesTheProxy(t *testing.T) { + t.Parallel() + for _, tc := range []struct{ line, what string }{ + {"HTTP/1.1 407Denied:muse-agent:hunter2", "malformed HTTP status code"}, + {"HTTP/1.1 407 Proxy Authentication Required\r\nDenied muse-agent hunter2", "malformed MIME header"}, + {"HTTP/1.1 407 Proxy Authentication Required\r\nTransfer-Encoding: muse-agent:hunter2", "unsupported transfer encoding"}, + } { + proxy := newTestProxy(t, nil, withAuth("u", "p"), withRejectLine(tc.line)) + r := mustExplicit(t, proxy.url("muse-agent:hunter2")) + want := "read CONNECT response: " + tc.what + " (response text withheld)" + + _, err := NewDialer(r).DialContext(context.Background(), "tcp", "api.pilot.invalid:443") + if err == nil || !strings.HasSuffix(err.Error(), want) || strings.Contains(err.Error(), "hunter2") { + t.Fatalf("%q: Dialer error = %v, want it to end %q", tc.line, err, want) + } + + client := &http.Client{Transport: RefreshingTransport(nil, r), Timeout: 10 * time.Second} + _, err = client.Get("https://api.pilot.invalid/") + if err == nil || !strings.Contains(err.Error(), "proxy CONNECT api.pilot.invalid:443: "+want) || strings.Contains(err.Error(), "hunter2") { + t.Fatalf("%q: RefreshingTransport error = %v, want %q", tc.line, err, want) + } + client.CloseIdleConnections() + } +}