From 333695d92ac0b76d9b102e17e23a083de29171e5 Mon Sep 17 00:00:00 2001 From: Simon Rozet Date: Mon, 21 Apr 2025 22:19:53 +0200 Subject: [PATCH] redirect https to fqdn --- main.go | 30 +++++++++++++++++++++++- tsproxy_test.go | 61 ++++++++++++++++++++++++++++++++++++++++++++++++- 2 files changed, 89 insertions(+), 2 deletions(-) diff --git a/main.go b/main.go index bd1e54b..0542c64 100644 --- a/main.go +++ b/main.go @@ -16,6 +16,7 @@ import ( "path/filepath" "sort" "strconv" + "strings" "syscall" "github.com/oklog/run" @@ -27,6 +28,7 @@ import ( "github.com/tailscale/hujson" "tailscale.com/client/local" "tailscale.com/client/tailscale/apitype" + "tailscale.com/ipn/ipnstate" "tailscale.com/tsnet" tslogger "tailscale.com/types/logger" ) @@ -356,7 +358,32 @@ func localTailnetHandler(logger *slog.Logger, lc tailscaleLocalClient, next http // localTailnetTLSHandler serves HTTPS on the local tailnet. func localTailnetTLSHandler(logger *slog.Logger, lc tailscaleLocalClient, next http.Handler) http.Handler { - return localTailnetHandler(logger, lc, next) + var ( + handler = localTailnetHandler(logger, lc, next) + dnsName string + ) + return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.TLS == nil { + panic("TLS handler wants TLS") + } + + if dnsName == "" { + st, err := lc.StatusWithoutPeers(r.Context()) + if err != nil { + http.Error(w, http.StatusText(http.StatusInternalServerError), http.StatusInternalServerError) + logger.Error("tailscale status", slog.Any("err", err)) + return + } + dnsName = strings.TrimSuffix(st.Self.DNSName, ".") + } + + if strings.TrimSuffix(r.Host, ".") != dnsName { + http.Redirect(w, r, fmt.Sprintf("https://%s%s", dnsName, r.RequestURI), http.StatusPermanentRedirect) + return + } + + handler.ServeHTTP(w, r) + }) } // insecureFunnelHandler handles HTTPS requests coming from Tailscale Funnel nodes. @@ -368,6 +395,7 @@ func insecureFunnelHandler(logger *slog.Logger, lc tailscaleLocalClient, next ht type tailscaleLocalClient interface { WhoIs(context.Context, string) (*apitype.WhoIsResponse, error) + StatusWithoutPeers(context.Context) (*ipnstate.Status, error) } func tsWhoIs(lc tailscaleLocalClient, r *http.Request) (*apitype.WhoIsResponse, error) { diff --git a/tsproxy_test.go b/tsproxy_test.go index 8d84f53..e575c91 100644 --- a/tsproxy_test.go +++ b/tsproxy_test.go @@ -14,17 +14,23 @@ import ( "github.com/prometheus/client_golang/prometheus" "github.com/prometheus/client_golang/prometheus/testutil" "tailscale.com/client/tailscale/apitype" + "tailscale.com/ipn/ipnstate" "tailscale.com/tailcfg" ) type fakeLocalClient struct { - whois func(context.Context, string) (*apitype.WhoIsResponse, error) + whois func(context.Context, string) (*apitype.WhoIsResponse, error) + status func(context.Context) (*ipnstate.Status, error) } func (c *fakeLocalClient) WhoIs(ctx context.Context, remoteAddr string) (*apitype.WhoIsResponse, error) { return c.whois(ctx, remoteAddr) } +func (c *fakeLocalClient) StatusWithoutPeers(ctx context.Context) (*ipnstate.Status, error) { + return c.status(ctx) +} + func TestLocalTailnetHandler(t *testing.T) { t.Parallel() @@ -111,6 +117,59 @@ func TestLocalTailnetHandler(t *testing.T) { } } +func TestLocalTailnetTLSHandler(t *testing.T) { + t.Parallel() + + lc := &fakeLocalClient{ + whois: func(_ context.Context, _ string) (*apitype.WhoIsResponse, error) { + return &apitype.WhoIsResponse{UserProfile: &tailcfg.UserProfile{LoginName: "tagged-devices"}, Node: &tailcfg.Node{Tags: []string{"foo"}}}, nil + }, + status: func(_ context.Context) (*ipnstate.Status, error) { + return &ipnstate.Status{Self: &ipnstate.PeerStatus{DNSName: "foo.ts.net."}}, nil + }, + } + be := http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + fmt.Fprintln(w, "Hi from the backend.") + }) + px := httptest.NewTLSServer(localTailnetTLSHandler(slog.New(slog.NewTextHandler(io.Discard, &slog.HandlerOptions{})), lc, be)) + defer px.Close() + + cli := px.Client() + cli.CheckRedirect = func(_ *http.Request, _ []*http.Request) error { + return http.ErrUseLastResponse + } + + req, err := http.NewRequestWithContext(t.Context(), http.MethodGet, px.URL+"/bar", nil) + if err != nil { + t.Fatal(err) + } + resp, err := cli.Do(req) + if err != nil { + t.Fatal(err) + } + defer resp.Body.Close() + if want, got := http.StatusPermanentRedirect, resp.StatusCode; want != got { + t.Fatalf("want status %d, got: %d", want, got) + } + if want, got := "https://foo.ts.net/bar", resp.Header.Get("location"); got != want { + t.Fatalf("want Location %s, got: %s", want, got) + } + + req, err = http.NewRequestWithContext(t.Context(), http.MethodGet, px.URL, nil) + if err != nil { + t.Fatal(err) + } + req.Host = "foo.ts.net" + resp, err = px.Client().Do(req) + if err != nil { + t.Fatal(err) + } + defer resp.Body.Close() + if want, got := http.StatusOK, resp.StatusCode; want != got { + t.Fatalf("want status %d, got: %d", want, got) + } +} + func TestServeDiscovery(t *testing.T) { t.Parallel() -- 2.51.2