diff --git a/prober/http.go b/prober/http.go index 53996e655..5ae527ac3 100644 --- a/prober/http.go +++ b/prober/http.go @@ -6,30 +6,73 @@ package prober import ( "bytes" "context" + "expvar" "fmt" "io" + "net" "net/http" + "net/netip" + "time" + "github.com/prometheus/client_golang/prometheus" "tailscale.com/net/netutil" ) const maxHTTPBody = 4 << 20 // MiB +// NewProbeTransport returns a fresh *http.Transport for a single probe run (so +// we never reuse a past connection). +// +// If dialAddr is valid (its zero value means "no override"), every connection is +// dialed to dialAddr instead of resolving the request URL's host, while SNI, the +// Host header, and TLS certificate validation continue to derive from the URL +// host. This is the HTTP analog of TLSWithIP: it lets a probe target a specific +// backend that serves a given hostname (e.g. one particular Funnel ingress node). +// +// Custom probe classes that dial a specific backend should use this rather than +// reconstructing the dial override, so the SNI/Host/cert semantics stay +// identical across probes. +func NewProbeTransport(dialAddr netip.AddrPort) *http.Transport { + tr := netutil.NewDefaultTransport() + if dialAddr.IsValid() { + dst := dialAddr.String() + // Reuse the transport's own dialer (preserving its Timeout/KeepAlive and + // any future tuning); only substitute the dial target so connections go + // to dialAddr instead of the resolved URL host. + dial := tr.DialContext + tr.DialContext = func(ctx context.Context, network, _ string) (net.Conn, error) { + return dial(ctx, network, dst) + } + } + return tr +} + // HTTP returns a ProbeClass that healthchecks an HTTP URL. // // The probe function sends a GET request for url, expects an HTTP 200 // response, and verifies that want is present in the response // body. func HTTP(url, wantText string) ProbeClass { + return httpProbe(url, netip.AddrPort{}, wantText) +} + +// HTTPWithDialAddr is like HTTP, but dials dialAddr (an ip:port) instead of the +// URL's host. SNI, the Host header, and TLS certificate validation still use the +// URL host, so this probes a specific backend serving the URL's hostname. +func HTTPWithDialAddr(url string, dialAddr netip.AddrPort, wantText string) ProbeClass { + return httpProbe(url, dialAddr, wantText) +} + +func httpProbe(url string, dialAddr netip.AddrPort, wantText string) ProbeClass { return ProbeClass{ Probe: func(ctx context.Context) error { - return probeHTTP(ctx, url, []byte(wantText)) + return probeHTTP(ctx, url, []byte(wantText), dialAddr) }, Class: "http", } } -func probeHTTP(ctx context.Context, url string, want []byte) error { +func probeHTTP(ctx context.Context, url string, want []byte, dialAddr netip.AddrPort) error { req, err := http.NewRequestWithContext(ctx, "GET", url, nil) if err != nil { return fmt.Errorf("constructing request: %w", err) @@ -37,7 +80,7 @@ func probeHTTP(ctx context.Context, url string, want []byte) error { // Get a completely new transport each time, so we don't reuse a // past connection. - tr := netutil.NewDefaultTransport() + tr := NewProbeTransport(dialAddr) defer tr.CloseIdleConnections() c := &http.Client{ Transport: tr, @@ -68,3 +111,88 @@ func probeHTTP(ctx context.Context, url string, want []byte) error { return nil } + +// HTTPBandwidth returns a ProbeClass that downloads size bytes from url and +// records how long the transfer took, for bandwidth measurement. It issues a +// GET, expects an HTTP 200 response, and reads size bytes from the body. +// +// Because the transfer is measured at the receiver (this prober reads and times +// the body it pulls), the recorded byte count and duration are exact even on a +// truncated response. This probe does not carry a direction label; callers that +// run it alongside an upload probe can attach one at registration time (e.g. +// Labels{"direction": "down"}). +// +// size must be positive. A non-positive size reads nothing from the body, so +// the probe records a zero-byte transfer and trivially succeeds. +func HTTPBandwidth(url string, size int64) ProbeClass { + return httpBandwidthProbe(url, size, netip.AddrPort{}) +} + +// HTTPBandwidthWithDialAddr is like HTTPBandwidth, but dials dialAddr (an +// ip:port) instead of the URL's host, while SNI/Host/cert validation still use +// the URL host. It measures download bandwidth from a specific backend serving +// the URL's hostname. +func HTTPBandwidthWithDialAddr(url string, size int64, dialAddr netip.AddrPort) ProbeClass { + return httpBandwidthProbe(url, size, dialAddr) +} + +func httpBandwidthProbe(url string, size int64, dialAddr netip.AddrPort) ProbeClass { + var transferTimeSeconds expvar.Float + var totalBytesTransferred expvar.Float + return ProbeClass{ + Probe: func(ctx context.Context) error { + return probeHTTPBandwidth(ctx, url, size, dialAddr, &transferTimeSeconds, &totalBytesTransferred) + }, + Class: "http_bw", + Metrics: HTTPBandwidthMetrics(size, &transferTimeSeconds, &totalBytesTransferred), + } +} + +// HTTPBandwidthMetrics returns the Metrics function for an "http_bw" bandwidth +// probe, exposing the configured payload size and the running transfer +// time/bytes accumulators. It is shared so probes that measure bandwidth +// differently (e.g. a receiver-reported upload probe) still emit an identical +// metric set and can be compared under a common direction label. +func HTTPBandwidthMetrics(size int64, transferTimeSeconds, totalBytesTransferred *expvar.Float) func(prometheus.Labels) []prometheus.Metric { + return func(lb prometheus.Labels) []prometheus.Metric { + return []prometheus.Metric{ + prometheus.MustNewConstMetric(prometheus.NewDesc("http_bw_probe_size_bytes", "Payload size of the bandwidth prober", nil, lb), prometheus.GaugeValue, float64(size)), + prometheus.MustNewConstMetric(prometheus.NewDesc("http_bw_transfer_time_seconds_total", "Time it took to transfer data", nil, lb), prometheus.CounterValue, transferTimeSeconds.Value()), + prometheus.MustNewConstMetric(prometheus.NewDesc("http_bw_bytes_total", "Amount of data transferred", nil, lb), prometheus.CounterValue, totalBytesTransferred.Value()), + } + } +} + +func probeHTTPBandwidth(ctx context.Context, url string, size int64, dialAddr netip.AddrPort, transferTimeSeconds, totalBytesTransferred *expvar.Float) error { + req, err := http.NewRequestWithContext(ctx, "GET", url, nil) + if err != nil { + return fmt.Errorf("constructing request: %w", err) + } + + // Get a completely new transport each time, so we don't reuse a + // past connection. + tr := NewProbeTransport(dialAddr) + defer tr.CloseIdleConnections() + c := &http.Client{ + Transport: tr, + } + + resp, err := c.Do(req) + if err != nil { + return fmt.Errorf("fetching %q: %w", url, err) + } + defer resp.Body.Close() + if resp.StatusCode != 200 { + return fmt.Errorf("fetching %q: status code %d, want 200", url, resp.StatusCode) + } + start := time.Now() + n, err := io.CopyN(io.Discard, resp.Body, size) + // Measure transfer time and bytes transferred irrespective of whether + // it succeeded or failed. + transferTimeSeconds.Add(time.Since(start).Seconds()) + totalBytesTransferred.Add(float64(n)) + if err != nil { + return fmt.Errorf("reading body of %q: %w", url, err) + } + return nil +} diff --git a/prober/http_test.go b/prober/http_test.go new file mode 100644 index 000000000..127256017 --- /dev/null +++ b/prober/http_test.go @@ -0,0 +1,171 @@ +// Copyright (c) Tailscale Inc & contributors +// SPDX-License-Identifier: BSD-3-Clause + +package prober + +import ( + "context" + "io" + "net" + "net/http" + "net/http/httptest" + "strings" + "testing" + "time" + + "github.com/prometheus/client_golang/prometheus" + dto "github.com/prometheus/client_model/go" +) + +// zeroReader is an io.Reader that yields an unlimited stream of zero bytes, used +// to generate fixed-size test payloads via io.CopyN. +type zeroReader struct{} + +func (zeroReader) Read(p []byte) (int, error) { + clear(p) + return len(p), nil +} + +// metricValue extracts the numeric value of the (gauge or counter) metric whose +// descriptor contains name from a slice returned by a ProbeClass.Metrics call. +func metricValue(t *testing.T, metrics []prometheus.Metric, name string) float64 { + t.Helper() + for _, m := range metrics { + if !strings.Contains(m.Desc().String(), name) { + continue + } + var dm dto.Metric + if err := m.Write(&dm); err != nil { + t.Fatalf("writing metric %q: %v", name, err) + } + switch { + case dm.Counter != nil: + return dm.Counter.GetValue() + case dm.Gauge != nil: + return dm.Gauge.GetValue() + default: + t.Fatalf("metric %q is neither counter nor gauge", name) + } + } + t.Fatalf("metric %q not found", name) + return 0 +} + +func TestHTTPBandwidth(t *testing.T) { + const size = 1 << 16 // 64 KiB + + mux := http.NewServeMux() + // /download writes exactly `size` zero bytes. + mux.HandleFunc("/download", func(w http.ResponseWriter, r *http.Request) { + w.WriteHeader(http.StatusOK) + io.CopyN(w, zeroReader{}, size) + }) + // /bad returns a non-200 status. + mux.HandleFunc("/bad", func(w http.ResponseWriter, r *http.Request) { + w.WriteHeader(http.StatusInternalServerError) + }) + // /short writes fewer than `size` bytes for a download. + mux.HandleFunc("/short", func(w http.ResponseWriter, r *http.Request) { + w.WriteHeader(http.StatusOK) + io.CopyN(w, zeroReader{}, size/2) + }) + + srv := httptest.NewServer(mux) + defer srv.Close() + + for _, tc := range []struct { + name string + path string + size int64 + wantErr bool + }{ + {name: "download_ok", path: "/download", size: size}, + {name: "download_non200", path: "/bad", size: size, wantErr: true}, + {name: "download_truncated", path: "/short", size: size, wantErr: true}, + } { + t.Run(tc.name, func(t *testing.T) { + pc := HTTPBandwidth(srv.URL+tc.path, tc.size) + + if got, want := pc.Class, "http_bw"; got != want { + t.Errorf("Class = %q, want %q", got, want) + } + + ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second) + defer cancel() + err := pc.Probe(ctx) + if tc.wantErr { + if err == nil { + t.Fatalf("Probe() = nil, want error") + } + return + } + if err != nil { + t.Fatalf("Probe() = %v, want nil", err) + } + + // On success, the Metrics callback should return the expected + // descriptors. + if pc.Metrics == nil { + t.Fatal("Metrics callback is nil") + } + metrics := pc.Metrics(prometheus.Labels{}) + wantDescs := map[string]bool{ + "http_bw_probe_size_bytes": false, + "http_bw_transfer_time_seconds_total": false, + "http_bw_bytes_total": false, + } + for _, m := range metrics { + if m == nil { + t.Fatal("got nil metric") + } + desc := m.Desc().String() + for name := range wantDescs { + if strings.Contains(desc, name) { + wantDescs[name] = true + } + } + } + for name, seen := range wantDescs { + if !seen { + t.Errorf("metric %q not emitted", name) + } + } + + // On a successful transfer the recorded byte count should equal the + // full payload size, and the transfer should take a positive, + // finite amount of time. + if got := metricValue(t, metrics, "http_bw_bytes_total"); got != float64(tc.size) { + t.Errorf("http_bw_bytes_total = %v, want %v", got, tc.size) + } + if got := metricValue(t, metrics, "http_bw_transfer_time_seconds_total"); got <= 0 { + t.Errorf("http_bw_transfer_time_seconds_total = %v, want > 0", got) + } + }) + } +} + +// TestHTTPWithDialAddr verifies that the dial-address override sends the +// connection to dialAddr while the URL host still drives the Host header (and, +// for HTTPS, SNI/cert validation). The URL host here is an unresolvable name, so +// the probe can only succeed if the dial override is honored. +func TestHTTPWithDialAddr(t *testing.T) { + const wantHost = "funnel-host.invalid" + var gotHost string + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + gotHost = r.Host + io.WriteString(w, "ok") + })) + defer srv.Close() + + dialAddr := srv.Listener.Addr().(*net.TCPAddr).AddrPort() + pc := HTTPWithDialAddr("http://"+wantHost+"/probe", dialAddr, "ok") + + ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second) + defer cancel() + if err := pc.Probe(ctx); err != nil { + t.Fatalf("Probe() = %v, want nil", err) + } + if gotHost != wantHost { + t.Errorf("server saw Host %q, want %q (URL host should drive the Host header)", gotHost, wantHost) + } +}