From d6417bcbe5430b6c1e54f9623af3bf0b49aa9ee8 Mon Sep 17 00:00:00 2001 From: Teodor Calin Date: Thu, 24 Sep 2026 02:06:39 +0300 Subject: [PATCH 1/2] feat(netproxy): refresh rotating proxy credentials (407 -> refresh -> retry once) Some egress proxies (Meta Muse) rotate the credentials embedded in HTTPS_PROXY every few minutes. A long-running process keeps its launch-time copy and gets 407 on every new CONNECT while its open tunnels stay up ("node online, all apps broken"). netproxy snapshotted the environment once, so a daemon broke within minutes. Resolver - Settings live in an atomically swapped state; ModeAuto re-reads the environment, and a refresh source (WithRefreshCommand: "sh -c", 10 s timeout, output and stderr never logged; or WithRefreshFunc) supplies the current proxy URL. EnvRefreshCommand documents PILOT_PROXY_CMD as the convention callers wire. - Timed refresh on lookup (DefaultRefreshInterval 60 s, WithRefreshInterval), Refresh(ctx) to force one, singleflight so concurrent callers share one run, last good settings kept on failure, WithRefreshErrorHandler told once per run of failures. - NewResolver(spec, opts...) is Parse with options; existing constructors keep their signatures and behaviour (auto additionally follows in-place environment changes). Dialer - On 407 (or an unparseable CONNECT response, how Muse's rejections surfaced) refresh and retry once on a new connection, only if the refreshed URL differs. A refresh that already ran after the pick is not repeated, so identical credentials never loop. RefreshingTransport(base, r) - Clone of base with Proxy = r.ProxyForRequest and a chained OnProxyConnectResponse that turns non-200 CONNECT answers into *ConnectError. Verified empirically: net/http reports a refused CONNECT only as the proxy's reason phrase ("Proxy Authentication Required", "unknown status code", or arbitrary proxy text) with no status code. - Retries once for GET/HEAD/OPTIONS/TRACE (or Idempotency-Key) and, after a 407, any request with GetBody; a POST without GetBody is not retried but the refresh still happens for the next request. Tests use an in-process CONNECT proxy whose accepted password rotates and a refresh command that counts its runs: long-running dialer across 4 rotations with earlier tunnels still echoing; 407 -> refresh -> success; unchanged credentials -> no retry; failing command -> last good kept and error reported once; 50 concurrent dials hitting a rotation -> exactly one command run; timeout; output never leaked; GET retried, POST without GetBody not retried. Co-Authored-By: Claude Opus 5.5 (1M context) --- netproxy/dialer.go | 67 ++- netproxy/netproxy.go | 283 +++++++++--- netproxy/refresh.go | 441 ++++++++++++++++++ netproxy/transport.go | 174 ++++++++ netproxy/zz_helpers_test.go | 29 +- netproxy/zz_refresh_test.go | 815 ++++++++++++++++++++++++++++++++++ netproxy/zz_transport_test.go | 345 ++++++++++++++ 7 files changed, 2091 insertions(+), 63 deletions(-) create mode 100644 netproxy/refresh.go create mode 100644 netproxy/transport.go create mode 100644 netproxy/zz_refresh_test.go create mode 100644 netproxy/zz_transport_test.go diff --git a/netproxy/dialer.go b/netproxy/dialer.go index 44faf40..01d5aa1 100644 --- a/netproxy/dialer.go +++ b/netproxy/dialer.go @@ -36,14 +36,22 @@ 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. 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,7 +105,10 @@ func (d *Dialer) DialContext(ctx context.Context, network, addr string) (net.Con ctx, cancel := context.WithTimeout(ctx, timeout) defer cancel() - proxyURL, err := resolver.ProxyForAddr(addr) + // Noted before the pick, so a refresh the pick itself runs counts as + // having happened after it. + start := resolver.attemptCount() + proxyURL, err := resolver.proxyForAddr(ctx, addr) if err != nil { return nil, err } @@ -109,7 +120,52 @@ 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 http.ReadResponse rejected as +// malformed (as opposed to an I/O error while reading it). +type badConnectResponse struct{ err error } + +func (e *badConnectResponse) Error() string { return "read CONNECT response: " + e.err.Error() } + +func (e *badConnectResponse) Unwrap() error { return e.err } + +// isMalformedResponse reports whether err is net/http's complaint about an +// unparseable response ("malformed HTTP status code", "malformed HTTP +// response", "malformed HTTP version", "malformed MIME header line"). +func isMalformedResponse(err error) bool { + if err == nil { + return false + } + msg := err.Error() + return strings.Contains(msg, "malformed HTTP ") || strings.Contains(msg, "malformed MIME header") } func (d *Dialer) forward(ctx context.Context, network, addr string) (net.Conn, error) { @@ -200,6 +256,9 @@ 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 { + if isMalformedResponse(err) { + return nil, &badConnectResponse{err: err} + } return nil, fmt.Errorf("read CONNECT response: %w", err) } // The body is never read: on success the rest of the stream belongs to diff --git a/netproxy/netproxy.go b/netproxy/netproxy.go index 33ceaf6..ee58922 100644 --- a/netproxy/netproxy.go +++ b/netproxy/netproxy.go @@ -21,12 +21,34 @@ // - 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 next +// time a proxy is looked up after the interval has passed; +// - 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 +56,9 @@ import ( "net/url" "os" "strings" + "sync" + "sync/atomic" + "time" ) // Mode names. Parse accepts ModeAuto and ModeOff (case-insensitive) in @@ -49,11 +74,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 +134,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 +159,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 +174,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 +204,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 +255,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 +282,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 +299,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 +340,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 +420,32 @@ 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 first refreshes the settings (see Refresh). func (r *Resolver) ProxyForAddr(addr string) (*url.URL, error) { - if !r.Enabled() { + return r.proxyForAddr(context.Background(), addr) +} + +func (r *Resolver) proxyForAddr(ctx context.Context, addr string) (*url.URL, error) { + // current first: a refresh can turn proxying on (a variable set in + // place, a refresh source's first URL). + st := r.current(ctx) + 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,24 @@ 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 (refreshed) proxy URL; 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(req.Context()) + 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 +485,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..0010428 --- /dev/null +++ b/netproxy/refresh.go @@ -0,0 +1,441 @@ +// 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 reads 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. 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 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 reads them again. 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, refreshing them first when the +// refresh interval has passed. If ctx ends before that refresh does, the +// previous settings are returned. +func (r *Resolver) current(ctx context.Context) *proxyState { + if !r.refreshable() || r.interval < 0 { + return r.snapshot() + } + r.mu.Lock() + var call *refreshCall + if time.Since(r.lastRun) >= r.interval { + call = r.inflight + if call == nil { + call = r.startLocked() + } + } + r.mu.Unlock() + if call != nil { + call.wait(ctx) + } + 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 + 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) { + 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 +} + +// 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/transport.go b/netproxy/transport.go new file mode 100644 index 0000000..1a00ebc --- /dev/null +++ b/netproxy/transport.go @@ -0,0 +1,174 @@ +// SPDX-License-Identifier: AGPL-3.0-or-later + +package netproxy + +import ( + "context" + "errors" + "net/http" + "net/url" + "strings" +) + +// 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"), which is how some +// proxies' rejections surface, also triggers a refresh; since such an error +// could in principle come from the server instead, it is retried only for +// the idempotent methods above. +// +// Connections already open, including tunnels in the idle pool, are left +// alone; they keep working until the proxy or the server closes them. 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 = r.ProxyForRequest + 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 + } + return &refreshingTransport{tr: tr, r: r} +} + +type refreshingTransport struct { + tr *http.Transport + r *Resolver +} + +// 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 runs counts as + // having happened after it (see Resolver.reauth). + start := t.r.attemptCount() + used, _ := t.r.ProxyForRequest(req) + resp, err := t.tr.RoundTrip(req) + if err == nil { + return resp, nil + } + + var rejectedWith *url.URL + var replayable bool + var ae *authRejectedError + switch { + case errors.As(err, &ae): + rejectedWith = ae.proxy + replayable = canReplay(req, true) + case used != nil && tunnelled(req) && isMalformedResponse(err): + rejectedWith = used + replayable = canReplay(req, false) + default: + return nil, err + } + _, retry, refreshErr := t.r.reauth(req.Context(), start, rejectedWith, func(st *proxyState) *url.URL { + return t.r.pickRequest(st, req) + }) + if !retry || !replayable { + return nil, withRefreshError(err, refreshErr) + } + again, rewindErr := rewind(req) + if rewindErr != nil { + return nil, err + } + return t.tr.RoundTrip(again) +} + +// 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..b0434a5 --- /dev/null +++ b/netproxy/zz_refresh_test.go @@ -0,0 +1,815 @@ +// 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() +} + +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 refresh 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() + } + 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 lookup picks up rotated credentials without + // waiting for a 407. + creds.fail(false) + creds.rotate() + 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 := proxyFor(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 := proxyFor(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 := proxyFor(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 as before. + _, err = NewDialer(mustExplicit(t, proxy.url("u:wrong"))).Dial("tcp", target) + if err == nil || !strings.Contains(err.Error(), `read CONNECT response: malformed HTTP status code "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, one lookup refreshes. A negative interval never does. + 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) + 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()) + } +} diff --git a/netproxy/zz_transport_test.go b/netproxy/zz_transport_test.go new file mode 100644 index 0000000..9395c51 --- /dev/null +++ b/netproxy/zz_transport_test.go @@ -0,0 +1,345 @@ +// SPDX-License-Identifier: AGPL-3.0-or-later + +package netproxy + +import ( + "context" + "crypto/tls" + "crypto/x509" + "errors" + "io" + "log" + "net" + "net/http" + "net/http/httptest" + "net/url" + "strings" + "sync" + "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(), "malformed HTTP status code") { + 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) + } + } +} From 70607b9a0b5caa5497342e3948fa8ccda23632ea Mon Sep 17 00:00:00 2001 From: Teodor Calin Date: Thu, 24 Sep 2026 02:40:36 +0300 Subject: [PATCH 2/2] fix(netproxy): address review of rotating-credential refresh - F1: a timed refresh no longer blocks the lookup that notices it is due. Resolver.current starts the refresh in the background and answers from the settings in hand (the last good credentials), so a slow or hung refresh command cannot fail dials whose budget (e.g. the registry client's 5 s) is shorter than the 10 s refresh timeout. Only a rejected credential (reauth after a 407) waits for a refresh. - F2: on Unix the refresh command runs in its own process group; a timeout SIGKILLs the whole group (exec.Cmd.Cancel), and so does a command that exits leaving a child holding its stdout, so hung children no longer pile up across refreshes. - F3: RefreshingTransport treats an unparseable response as a proxy rejection only while the tunnel is being set up (no GotConn since the last GetConn, via httptrace). A server sending a malformed response through the tunnel gets net/http's error untouched and never triggers the refresh command. The proxy URL a refused CONNECT used is captured from the Proxy call itself. - F4: net/http keys its pool by proxy URL including credentials, so once a refresh changes the proxy settings, the first request afterwards closes the stranded idle connections. A refresh that reads the same settings again leaves the pool alone. Doc corrected. - F5: an unparseable CONNECT response is reported by what was wrong with it ("malformed HTTP status code (response text withheld)"), never by its text, from both Dialer and RefreshingTransport; the Dialer keeps wrapping only genuine read errors (EOF, net.Error). Each fix has a regression test that fails on the previous commit. Co-Authored-By: Claude Opus 5.5 (1M context) --- netproxy/dialer.go | 79 ++++++++++--- netproxy/netproxy.go | 23 ++-- netproxy/refresh.go | 74 +++++++----- netproxy/refresh_other.go | 14 +++ netproxy/refresh_unix.go | 34 ++++++ netproxy/transport.go | 148 ++++++++++++++++++------ netproxy/zz_refresh_test.go | 205 ++++++++++++++++++++++++++++++++-- netproxy/zz_transport_test.go | 174 ++++++++++++++++++++++++++++- 8 files changed, 650 insertions(+), 101 deletions(-) create mode 100644 netproxy/refresh_other.go create mode 100644 netproxy/refresh_unix.go diff --git a/netproxy/dialer.go b/netproxy/dialer.go index 01d5aa1..a031153 100644 --- a/netproxy/dialer.go +++ b/netproxy/dialer.go @@ -43,6 +43,10 @@ const DefaultTimeout = 30 * time.Second // (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 @@ -105,10 +109,10 @@ 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 refresh the pick itself runs counts as - // having happened after it. + // 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(ctx, addr) + proxyURL, err := resolver.ProxyForAddr(addr) if err != nil { return nil, err } @@ -149,23 +153,59 @@ func credentialsRejected(err error) bool { return errors.As(err, &bad) } -// badConnectResponse is a CONNECT response http.ReadResponse rejected as -// malformed (as opposed to an I/O error while reading it). -type badConnectResponse struct{ err error } +// 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.err.Error() } +func (e *badConnectResponse) Error() string { + return "read CONNECT response: " + e.what + " (response text withheld)" +} -func (e *badConnectResponse) Unwrap() error { return e.err } +// 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", +} -// isMalformedResponse reports whether err is net/http's complaint about an -// unparseable response ("malformed HTTP status code", "malformed HTTP -// response", "malformed HTTP version", "malformed MIME header line"). -func isMalformedResponse(err error) bool { +// responseFault returns which of responseFaults err reports, or "" if none. +func responseFault(err error) string { if err == nil { - return false + return "" } msg := err.Error() - return strings.Contains(msg, "malformed HTTP ") || strings.Contains(msg, "malformed MIME header") + 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) { @@ -256,10 +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 { - if isMalformedResponse(err) { - return nil, &badConnectResponse{err: 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, fmt.Errorf("read CONNECT response: %w", err) + 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 ee58922..e0d1087 100644 --- a/netproxy/netproxy.go +++ b/netproxy/netproxy.go @@ -29,8 +29,10 @@ // 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 next -// time a proxy is looked up after the interval has passed; +// - 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. @@ -421,15 +423,13 @@ func (st *proxyState) 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. When the refresh interval has -// passed, the lookup first refreshes the settings (see Refresh). +// 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) { - return r.proxyForAddr(context.Background(), addr) -} - -func (r *Resolver) proxyForAddr(ctx context.Context, addr string) (*url.URL, error) { // current first: a refresh can turn proxying on (a variable set in // place, a refresh source's first URL). - st := r.current(ctx) + st := r.current() if !st.proxies() { return nil, nil } @@ -456,13 +456,14 @@ func (r *Resolver) pickAddr(st *proxyState, 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. The result is always -// the current (refreshed) proxy URL; RefreshingTransport also retries a -// request once when the proxy rejects the credentials. +// 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 req == nil || req.URL == nil { return nil, nil } - st := r.current(req.Context()) + st := r.current() if !st.proxies() { return nil, nil } diff --git a/netproxy/refresh.go b/netproxy/refresh.go index 0010428..d6df288 100644 --- a/netproxy/refresh.go +++ b/netproxy/refresh.go @@ -22,7 +22,7 @@ import ( const EnvRefreshCommand = "PILOT_PROXY_CMD" // DefaultRefreshInterval is how long a refreshing Resolver uses the settings -// it last read before a lookup reads them again. +// 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 @@ -54,16 +54,18 @@ func newOptions(opts []Option) options { // 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. 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. +// 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 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. +// 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 @@ -97,9 +99,13 @@ func WithRefreshFunc(fn func(ctx context.Context) (string, error)) Option { } // WithRefreshInterval sets how long the Resolver uses the settings it last -// read before a lookup reads them again. Zero means DefaultRefreshInterval; -// a negative interval turns timed refreshes off, leaving only the refreshes -// a rejected credential triggers (and explicit Refresh calls). +// 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 } } @@ -154,24 +160,20 @@ func (r *Resolver) snapshot() *proxyState { return &emptyState } -// current returns the settings for a lookup, refreshing them first when the -// refresh interval has passed. If ctx ends before that refresh does, the -// previous settings are returned. -func (r *Resolver) current(ctx context.Context) *proxyState { - if !r.refreshable() || r.interval < 0 { - return r.snapshot() - } - r.mu.Lock() - var call *refreshCall - if time.Since(r.lastRun) >= r.interval { - call = r.inflight - if call == nil { - call = r.startLocked() +// 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() - if call != nil { - call.wait(ctx) + r.mu.Unlock() } return r.snapshot() } @@ -373,6 +375,9 @@ func commandSource(command string) func(context.Context) (string, error) { // 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: @@ -383,6 +388,9 @@ func commandSource(command string) func(context.Context) (string, error) { 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) @@ -410,6 +418,12 @@ func (b *cappedBuffer) Write(p []byte) (int, error) { 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 { 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 index 1a00ebc..4f1a718 100644 --- a/netproxy/transport.go +++ b/netproxy/transport.go @@ -5,9 +5,13 @@ 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 @@ -34,15 +38,23 @@ import ( // 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"), which is how some -// proxies' rejections surface, also triggers a refresh; since such an error -// could in principle come from the server instead, it is retried only for -// the idempotent methods above. +// whose message never includes proxy-supplied text. // -// Connections already open, including tunnels in the idle pool, are left -// alone; they keep working until the proxy or the server closes them. The -// returned RoundTripper has a CloseIdleConnections method, so +// 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 @@ -54,7 +66,13 @@ func RefreshingTransport(base *http.Transport, r *Resolver) http.RoundTripper { default: tr = &http.Transport{} } - tr.Proxy = r.ProxyForRequest + 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 { @@ -72,14 +90,36 @@ func RefreshingTransport(base *http.Transport, r *Resolver) http.RoundTripper { } return ce } - return &refreshingTransport{tr: tr, r: r} + 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 { @@ -90,39 +130,83 @@ type authRejectedError struct { 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 runs counts as + // Noted before the lookup, so a timed refresh the lookup starts counts as // having happened after it (see Resolver.reauth). start := t.r.attemptCount() - used, _ := t.r.ProxyForRequest(req) - resp, err := t.tr.RoundTrip(req) - if err == nil { - return resp, nil + t.dropStrandedConns() + resp, rej, err := t.send(req) + if rej.proxy == nil { + return resp, err } - - var rejectedWith *url.URL - var replayable bool - var ae *authRejectedError - switch { - case errors.As(err, &ae): - rejectedWith = ae.proxy - replayable = canReplay(req, true) - case used != nil && tunnelled(req) && isMalformedResponse(err): - rejectedWith = used - replayable = canReplay(req, false) - default: - return nil, err - } - _, retry, refreshErr := t.r.reauth(req.Context(), start, rejectedWith, func(st *proxyState) *url.URL { + _, retry, refreshErr := t.r.reauth(req.Context(), start, rej.proxy, func(st *proxyState) *url.URL { return t.r.pickRequest(st, req) }) - if !retry || !replayable { + if !retry || !canReplay(req, rej.refused) { return nil, withRefreshError(err, refreshErr) } again, rewindErr := rewind(req) if rewindErr != nil { return nil, err } - return t.tr.RoundTrip(again) + 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 diff --git a/netproxy/zz_refresh_test.go b/netproxy/zz_refresh_test.go index b0434a5..6c1d6fc 100644 --- a/netproxy/zz_refresh_test.go +++ b/netproxy/zz_refresh_test.go @@ -152,6 +152,34 @@ func (e *syncEnv) set(k, v string) { 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) @@ -349,8 +377,9 @@ func TestRefreshCommandFailureKeepsLastGoodCredentials(t *testing.T) { // user1 returns the first generation's credentials. func (c *credSource) user1() (user, pass string) { return "muse-agent", "tok/1?r#s@t" } -// Lookups refresh on the timer too. With a failing command every lookup -// keeps the last good URL, and the failure is still reported only once. +// 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) @@ -366,6 +395,7 @@ func TestRefreshIntervalKeepsLastGoodOnFailure(t *testing.T) { 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) @@ -374,10 +404,13 @@ func TestRefreshIntervalKeepsLastGoodOnFailure(t *testing.T) { t.Fatalf("error handler called %d times, want 1: %v", len(e), e) } - // Once the command works, a lookup picks up rotated credentials without - // waiting for a 407. + // 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) @@ -475,7 +508,7 @@ func TestRefreshAutoModeRereadsEnvironment(t *testing.T) { } env2.set("HTTPS_PROXY", "http://b.invalid:2") env2.set("NO_PROXY", "skip.invalid") - if got := proxyFor(t, r2, "x.invalid:443"); got != "http://b.invalid:2" { + 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 != "" { @@ -492,7 +525,7 @@ func TestRefreshAutoModeRereadsEnvironment(t *testing.T) { t.Fatal("empty environment proxies") } env3.set("https_proxy", "http://late.invalid:3128") - if got := proxyFor(t, r3, "x.invalid:443"); got != "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" { @@ -504,7 +537,7 @@ func TestRefreshAutoModeRereadsEnvironment(t *testing.T) { // An unusable HTTPS_PROXY is a failed refresh: the last reading stays. env2.set("HTTPS_PROXY", "socks5://c.invalid:3") - if got := proxyFor(t, r2, "x.invalid:443"); got != "http://b.invalid:2" { + 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 @@ -571,9 +604,10 @@ func TestDialerRefreshesOnMalformedRejection(t *testing.T) { t.Fatalf("rejected %d, command runs %d; want 1 and 2", got, runs) } - // Without anything to refresh, the error is returned as before. + // 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.Contains(err.Error(), `read CONNECT response: malformed HTTP status code "407Proxy"`) { + 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) } } @@ -748,7 +782,8 @@ func TestNewResolverOptions(t *testing.T) { } // Timed refreshes: within the interval lookups reuse the last reading; - // after it, one lookup refreshes. A negative interval never does. + // 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 @@ -763,6 +798,8 @@ func TestNewResolverOptions(t *testing.T) { } } 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) } @@ -813,3 +850,151 @@ func TestRefreshWaitHonoursContext(t *testing.T) { 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 index 9395c51..be4b107 100644 --- a/netproxy/zz_transport_test.go +++ b/netproxy/zz_transport_test.go @@ -3,6 +3,7 @@ package netproxy import ( + "bufio" "context" "crypto/tls" "crypto/x509" @@ -15,6 +16,7 @@ import ( "net/url" "strings" "sync" + "sync/atomic" "testing" "time" ) @@ -269,7 +271,7 @@ func TestRefreshingTransportMalformedRejection(t *testing.T) { } creds.rotate() _, err = do(t, client, http.MethodPost, api.url("/"), strings.NewReader("x")) - if err == nil || !strings.Contains(err.Error(), "malformed HTTP status code") { + 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 { @@ -343,3 +345,173 @@ func TestCanReplay(t *testing.T) { } } } + +// 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() + } +}