From 995cdd32cfbcf7b449ee34bd40cd70088764e240 Mon Sep 17 00:00:00 2001 From: Simon Rozet Date: Mon, 21 Apr 2025 15:09:07 +0200 Subject: [PATCH] split server and handler This is in preparation for adding OIDC. --- main.go | 211 ++++++++++++++++++++++++++++++------------------ tsproxy_test.go | 15 +--- 2 files changed, 138 insertions(+), 88 deletions(-) diff --git a/main.go b/main.go index 4f1d131..bd1e54b 100644 --- a/main.go +++ b/main.go @@ -4,6 +4,7 @@ import ( "context" "crypto/tls" "encoding/json" + "errors" "flag" "fmt" "log/slog" @@ -219,116 +220,172 @@ func tsproxy(ctx context.Context) error { if err != nil { return fmt.Errorf("upstream %s: parse backend URL: %w", upstream.Name, err) } + // TODO(sr) Instrument proxy.Transport + proxy := &httputil.ReverseProxy{ + Rewrite: func(req *httputil.ProxyRequest) { + req.SetURL(backendURL) + req.SetXForwarded() + req.Out.Host = req.In.Host + }, + } + proxy.ErrorHandler = func(w http.ResponseWriter, _ *http.Request, err error) { + http.Error(w, http.StatusText(http.StatusBadGateway), http.StatusBadGateway) + logger.Error("upstream error", lerr(err)) + } - srv := &http.Server{ - TLSConfig: &tls.Config{GetCertificate: lc.GetCertificate}, - Handler: promhttp.InstrumentHandlerInFlight(requestsInFlight.With(prometheus.Labels{"upstream": upstream.Name}), - promhttp.InstrumentHandlerDuration(duration.MustCurryWith(prometheus.Labels{"upstream": upstream.Name}), - promhttp.InstrumentHandlerCounter(requests.MustCurryWith(prometheus.Labels{"upstream": upstream.Name}), - newReverseProxy(log, lc, backendURL)))), + instrument := func(h http.Handler) http.Handler { + return promhttp.InstrumentHandlerInFlight( + requestsInFlight.With(prometheus.Labels{"upstream": upstream.Name}), + promhttp.InstrumentHandlerDuration( + duration.MustCurryWith(prometheus.Labels{"upstream": upstream.Name}), + promhttp.InstrumentHandlerCounter( + requests.MustCurryWith(prometheus.Labels{"upstream": upstream.Name}), + h, + ), + ), + ) } - g.Add(func() error { - st, err := ts.Up(ctx) - if err != nil { - return fmt.Errorf("tailscale: wait for node %s to be ready: %w", upstream.Name, err) - } + { + var srv *http.Server + g.Add(func() error { + st, err := ts.Up(ctx) + if err != nil { + return fmt.Errorf("tailscale: wait for tsnet %s to be ready: %w", upstream.Name, err) + } - // register in service discovery when we're ready. - targets[i] = target{name: upstream.Name, prometheus: upstream.Prometheus, magicDNS: st.Self.DNSName} + srv = &http.Server{Handler: instrument(localTailnetHandler(log, lc, proxy))} + ln, err := ts.Listen("tcp", ":80") + if err != nil { + return fmt.Errorf("tailscale: listen for %s on port 80: %w", upstream.Name, err) + } - ln, err := ts.Listen("tcp", ":80") - if err != nil { - return fmt.Errorf("tailscale: listen for %s on port 80: %w", upstream.Name, err) - } - return srv.Serve(ln) - }, func(_ error) { - if err := srv.Close(); err != nil { - log.Error("server shutdown", lerr(err)) - } - cancel() - }) - g.Add(func() error { - _, err := ts.Up(ctx) - if err != nil { - return fmt.Errorf("tailscale: wait for node %s to be ready: %w", upstream.Name, err) - } + // register in service discovery when we're ready. + targets[i] = target{name: upstream.Name, prometheus: upstream.Prometheus, magicDNS: st.Self.DNSName} - if upstream.Funnel { - ln, err := ts.ListenFunnel("tcp", ":443") + return srv.Serve(ln) + }, func(_ error) { + if srv != nil { + if err := srv.Close(); err != nil { + log.Error("server shutdown", lerr(err)) + } + } + cancel() + }) + } + { + var srv *http.Server + g.Add(func() error { + _, err := ts.Up(ctx) if err != nil { - return fmt.Errorf("tailscale: funnel for %s on port 443: %w", upstream.Name, err) + return fmt.Errorf("tailscale: wait for tsnet %s to be ready: %w", upstream.Name, err) + } + srv = &http.Server{ + TLSConfig: &tls.Config{GetCertificate: lc.GetCertificate}, + Handler: instrument(localTailnetTLSHandler(log, lc, proxy)), } - return srv.Serve(ln) - } - ln, err := ts.Listen("tcp", ":443") - if err != nil { - return fmt.Errorf("tailscale: listen for %s on port 443: %w", upstream.Name, err) - } - return srv.ServeTLS(ln, "", "") - }, func(_ error) { - if err := srv.Close(); err != nil { - log.Error("TLS server shutdown", lerr(err)) + ln, err := ts.Listen("tcp", ":443") + if err != nil { + return fmt.Errorf("tailscale: listen for %s on port 443: %w", upstream.Name, err) + } + return srv.ServeTLS(ln, "", "") + }, func(_ error) { + if srv != nil { + if err := srv.Close(); err != nil { + log.Error("TLS server shutdown", lerr(err)) + } + } + cancel() + }) + } + if upstream.Funnel { + { + var srv *http.Server + g.Add(func() error { + _, err := ts.Up(ctx) + if err != nil { + return fmt.Errorf("tailscale: wait for tsnet %s to be ready: %w", upstream.Name, err) + } + srv = &http.Server{ + Handler: instrument(insecureFunnelHandler(log, lc, proxy)), + } + + ln, err := ts.ListenFunnel("tcp", ":443", tsnet.FunnelOnly()) + if err != nil { + return fmt.Errorf("tailscale: funnel for %s on port 443: %w", upstream.Name, err) + } + return srv.Serve(ln) + }, func(_ error) { + if srv != nil { + if err := srv.Close(); err != nil { + log.Error("TLS server shutdown", lerr(err)) + } + } + cancel() + }) } - cancel() - }) + } } return g.Run() } -type tailscaleLocalClient interface { - WhoIs(context.Context, string) (*apitype.WhoIsResponse, error) -} - -func newReverseProxy(logger *slog.Logger, lc tailscaleLocalClient, url *url.URL) http.HandlerFunc { - // TODO(sr) Instrument proxy.Transport - rproxy := &httputil.ReverseProxy{ - Rewrite: func(req *httputil.ProxyRequest) { - req.SetURL(url) - req.SetXForwarded() - req.Out.Host = req.In.Host - }, - } - rproxy.ErrorHandler = func(w http.ResponseWriter, _ *http.Request, err error) { - http.Error(w, http.StatusText(http.StatusBadGateway), http.StatusBadGateway) - logger.Error("upstream error", lerr(err)) - } - +// localTailnetHandler serves plain-HTTP on the local tailnet. +func localTailnetHandler(logger *slog.Logger, lc tailscaleLocalClient, next http.Handler) http.Handler { return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - whois, err := lc.WhoIs(r.Context(), r.RemoteAddr) + whois, err := tsWhoIs(lc, r) if err != nil { http.Error(w, http.StatusText(http.StatusInternalServerError), http.StatusInternalServerError) logger.Error("tailscale whois", lerr(err)) return } - if whois.Node == nil { - http.Error(w, http.StatusText(http.StatusInternalServerError), http.StatusInternalServerError) - logger.Error("tailscale whois", slog.String("err", "node missing")) - return - } - - if whois.UserProfile == nil { - http.Error(w, http.StatusText(http.StatusInternalServerError), http.StatusInternalServerError) - logger.Error("tailscale whois", slog.String("err", "user profile missing")) - return - } - // Proxy requests from tagged nodes as is. if whois.Node.IsTagged() { - rproxy.ServeHTTP(w, r) + next.ServeHTTP(w, r) return } req := r.Clone(r.Context()) req.Header.Set("X-Webauth-User", whois.UserProfile.LoginName) req.Header.Set("X-Webauth-Name", whois.UserProfile.DisplayName) - rproxy.ServeHTTP(w, req) + next.ServeHTTP(w, req) }) } +// localTailnetTLSHandler serves HTTPS on the local tailnet. +func localTailnetTLSHandler(logger *slog.Logger, lc tailscaleLocalClient, next http.Handler) http.Handler { + return localTailnetHandler(logger, lc, next) +} + +// insecureFunnelHandler handles HTTPS requests coming from Tailscale Funnel nodes. +// This is marked insecure because the upstream is exposed to the public Internet. +// The upstream is responsible for implementing authentication. +func insecureFunnelHandler(logger *slog.Logger, lc tailscaleLocalClient, next http.Handler) http.Handler { + return localTailnetHandler(logger, lc, next) +} + +type tailscaleLocalClient interface { + WhoIs(context.Context, string) (*apitype.WhoIsResponse, error) +} + +func tsWhoIs(lc tailscaleLocalClient, r *http.Request) (*apitype.WhoIsResponse, error) { + whois, err := lc.WhoIs(r.Context(), r.RemoteAddr) + if err != nil { + return nil, fmt.Errorf("tailscale whois: %w", err) + } + + if whois.Node == nil { + return nil, errors.New("tailscale whois: node missing") + } + + if whois.UserProfile == nil { + return nil, errors.New("tailscale whois: user profile missing") + } + return whois, nil +} + func serveDiscovery(self string, targets []target) http.Handler { return http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { var tgs []string diff --git a/tsproxy_test.go b/tsproxy_test.go index 7619b0e..8d84f53 100644 --- a/tsproxy_test.go +++ b/tsproxy_test.go @@ -5,11 +5,9 @@ import ( "errors" "fmt" "io" - "log" "log/slog" "net/http" "net/http/httptest" - "net/url" "testing" "github.com/google/go-cmp/cmp" @@ -27,7 +25,7 @@ func (c *fakeLocalClient) WhoIs(ctx context.Context, remoteAddr string) (*apityp return c.whois(ctx, remoteAddr) } -func TestReverseProxy(t *testing.T) { +func TestLocalTailnetHandler(t *testing.T) { t.Parallel() for _, tc := range []struct { @@ -80,18 +78,13 @@ func TestReverseProxy(t *testing.T) { t.Run(tc.name, func(t *testing.T) { t.Parallel() lc := &fakeLocalClient{whois: tc.whois} - be := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + be := http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { for k, v := range r.Header { w.Header().Set(k, v[0]) } fmt.Fprintln(w, "Hi from the backend.") - })) - defer be.Close() - beURL, err := url.Parse(be.URL) - if err != nil { - log.Fatal(err) - } - px := httptest.NewServer(newReverseProxy(slog.New(slog.NewTextHandler(io.Discard, &slog.HandlerOptions{})), lc, beURL)) + }) + px := httptest.NewServer(localTailnetHandler(slog.New(slog.NewTextHandler(io.Discard, &slog.HandlerOptions{})), lc, be)) defer px.Close() resp, err := http.Get(px.URL) -- 2.51.2