From ed8d5466758cbafa6d187e43ed068b3a19eaaeef Mon Sep 17 00:00:00 2001 From: Kristoffer Dalby Date: Wed, 30 Sep 2026 15:24:49 +0000 Subject: [PATCH] noise: check capability floor before handleRegister Below-floor register got 400 only after logout, key use or auth-cache write had run. --- CHANGELOG.md | 1 + hscontrol/noise.go | 49 +++++----- hscontrol/noise_test.go | 198 +++++++++++++++++++++++++++++++++++++++- 3 files changed, 220 insertions(+), 28 deletions(-) diff --git a/CHANGELOG.md b/CHANGELOG.md index 58086a0fe..3e5ef1261 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -127,6 +127,7 @@ clients, and how to run the same setup without Nix. - `headscale nodes backfillips` now sends the new IPs to connected clients [#3501](https://github.com/juanfont/headscale/pull/3501) - Fix packet filters under `autogroup:self` not updating after a user is added or renamed [#3501](https://github.com/juanfont/headscale/pull/3501) - Fix a rejected policy leaving its packet filter active [#3501](https://github.com/juanfont/headscale/pull/3501) +- Fix a registration request from a client below the minimum supported version still being acted on before it was rejected: it could log the node out, delete an ephemeral node, consume a pre-auth key or start a login. It is now rejected before anything changes ## 0.29.4 (2026-09-23) diff --git a/hscontrol/noise.go b/hscontrol/noise.go index 7ce89e8c4..d863d721c 100644 --- a/hscontrol/noise.go +++ b/hscontrol/noise.go @@ -765,37 +765,36 @@ func (ns *noiseServer) RegistrationHandler( return } - registerRequest, registerResponse := func() (*tailcfg.RegisterRequest, *tailcfg.RegisterResponse) { //nolint:contextcheck - var resp *tailcfg.RegisterResponse + var registerRequest tailcfg.RegisterRequest - var regReq tailcfg.RegisterRequest + decodeErr := json.NewDecoder(req.Body).Decode(®isterRequest) - err := json.NewDecoder(req.Body).Decode(®Req) - if err != nil { - return ®Req, regErr(err) - } - - resp, err = ns.headscale.handleRegister(req.Context(), regReq, ns.conn.Peer()) - if err != nil { - if httpErr, ok := errors.AsType[HTTPError](err); ok { - resp = &tailcfg.RegisterResponse{ - Error: httpErr.Msg, - } - - return ®Req, resp - } - - return ®Req, regErr(err) - } - - return ®Req, resp - }() - - // Reject unsupported versions + // The floor must be enforced before handleRegister: a logout, pre-auth + // key use or auth-cache write it performs is not undone by a later 400. + // A failed decode still gets checked against whatever Version it read. if rejectUnsupported(writer, registerRequest.Version, ns.machineKey, registerRequest.NodeKey) { return } + registerResponse := func() *tailcfg.RegisterResponse { //nolint:contextcheck + if decodeErr != nil { + return regErr(decodeErr) + } + + resp, err := ns.headscale.handleRegister(req.Context(), registerRequest, ns.machineKey) + if err != nil { + if httpErr, ok := errors.AsType[HTTPError](err); ok { + return &tailcfg.RegisterResponse{ + Error: httpErr.Msg, + } + } + + return regErr(err) + } + + return resp + }() + writer.Header().Set("Content-Type", "application/json; charset=utf-8") writer.WriteHeader(http.StatusOK) diff --git a/hscontrol/noise_test.go b/hscontrol/noise_test.go index 942b1bd74..2ae254e13 100644 --- a/hscontrol/noise_test.go +++ b/hscontrol/noise_test.go @@ -10,11 +10,13 @@ import ( "net/http" "net/http/httptest" "net/url" + "slices" "strconv" "testing" "time" "github.com/go-chi/chi/v5" + "github.com/juanfont/headscale/hscontrol/capver" "github.com/juanfont/headscale/hscontrol/types" "github.com/juanfont/headscale/hscontrol/util" "github.com/stretchr/testify/assert" @@ -197,12 +199,202 @@ func TestRegistrationHandler_OversizedBody(t *testing.T) { ns.RegistrationHandler(rec, req) - // [json.Decoder.Decode] returns [http.MaxBytesError] → [regErr] wraps it → handler writes - // a [tailcfg.RegisterResponse] with the error and then [rejectUnsupported] kicks in - // for version 0 → returns 400. + // [json.Decoder.Decode] returns [http.MaxBytesError] before any field is + // decoded, so [rejectUnsupported] sees version 0 and answers 400 before + // the decode error would be. assert.Equal(t, http.StatusBadRequest, rec.Code) } +func newRegisterRequest(t *testing.T, req tailcfg.RegisterRequest) *http.Request { + t.Helper() + + body, err := json.Marshal(req) + require.NoError(t, err) + + return httptest.NewRequestWithContext(context.Background(), http.MethodPost, "/machine/register", bytes.NewReader(body)) +} + +// serveRegister guards against panics so a handler that reaches a nil +// dependency fails its own row instead of the whole test binary. +func serveRegister(t *testing.T, ns *noiseServer, req tailcfg.RegisterRequest) *httptest.ResponseRecorder { + t.Helper() + + rec := httptest.NewRecorder() + + require.NotPanics(t, func() { + ns.RegistrationHandler(rec, newRegisterRequest(t, req)) + }) + + return rec +} + +func requireBelowFloorRejected(t *testing.T, rec *httptest.ResponseRecorder) { + t.Helper() + + require.Equal(t, http.StatusBadRequest, rec.Code, "body=%q", rec.Body.String()) + assert.Contains(t, rec.Body.String(), ErrUnsupportedClientVersion.Error()) +} + +// TestRegistrationHandler_BelowFloorLeavesNoStateChange pins that the +// capability floor is checked before [Headscale.handleRegister]. A logout, +// pre-auth key use or auth-cache write that ran first would not be undone by +// the 400 the client then receives. +func TestRegistrationHandler_BelowFloorLeavesNoStateChange(t *testing.T) { + t.Parallel() + + type versionCase struct { + name string + version tailcfg.CapabilityVersion + } + + belowFloor := []versionCase{ + {"v0", 0}, + {"floor-1", capver.MinSupportedCapabilityVersion - 1}, + } + + // A nil headscale proves the rejected request never reaches it. + t.Run("nil_server", func(t *testing.T) { + t.Parallel() + + authID := types.MustAuthID() + + requests := []struct { + name string + req tailcfg.RegisterRequest + }{ + {"interactive", tailcfg.RegisterRequest{ + NodeKey: key.NewNode().Public(), + Hostinfo: &tailcfg.Hostinfo{Hostname: "floor-interactive"}, + }}, + {"authkey", tailcfg.RegisterRequest{ + NodeKey: key.NewNode().Public(), + Auth: &tailcfg.RegisterResponseAuth{AuthKey: "floor-authkey"}, + }}, + {"followup", tailcfg.RegisterRequest{ + NodeKey: key.NewNode().Public(), + Followup: "http://localhost:8080/register/" + authID.String(), + }}, + {"logout", tailcfg.RegisterRequest{ + NodeKey: key.NewNode().Public(), + Expiry: time.Unix(123, 0), + }}, + } + + for _, rc := range requests { + t.Run(rc.name, func(t *testing.T) { + t.Parallel() + + for _, vc := range belowFloor { + t.Run(vc.name, func(t *testing.T) { + t.Parallel() + + req := rc.req + req.Version = vc.version + + ns := &noiseServer{machineKey: key.NewMachine().Public()} + requireBelowFloorRejected(t, serveRegister(t, ns, req)) + }) + } + }) + } + }) + + t.Run("authkey", func(t *testing.T) { + t.Parallel() + + // Positive control: exactly the floor registers, pinning the >= boundary. + cases := append(slices.Clone(belowFloor), versionCase{"floor", capver.MinSupportedCapabilityVersion}) + + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + t.Parallel() + + app := createTestApp(t) + user := app.state.CreateUserForTest("floor-authkey-user") + + pak, err := app.state.CreatePreAuthKey(user.TypedID(), false, false, nil, nil) + require.NoError(t, err) + + ns := &noiseServer{headscale: app, machineKey: key.NewMachine().Public()} + rec := serveRegister(t, ns, tailcfg.RegisterRequest{ + Version: tc.version, + NodeKey: key.NewNode().Public(), + Auth: &tailcfg.RegisterResponseAuth{AuthKey: pak.Key}, + Hostinfo: &tailcfg.Hostinfo{Hostname: "floor-authkey-node"}, + }) + + stored, err := app.state.GetPreAuthKey(pak.Key) + require.NoError(t, err) + + if tc.version < capver.MinSupportedCapabilityVersion { + requireBelowFloorRejected(t, rec) + assert.Equal(t, 0, app.state.ListNodes().Len(), "rejected request must not register a node") + assert.False(t, stored.Used, "rejected request must not consume the key") + + return + } + + require.Equal(t, http.StatusOK, rec.Code, "body=%q", rec.Body.String()) + + var resp tailcfg.RegisterResponse + require.NoError(t, json.Unmarshal(rec.Body.Bytes(), &resp)) + assert.True(t, resp.MachineAuthorized, "resp=%+v", resp) + assert.Equal(t, 1, app.state.ListNodes().Len()) + assert.True(t, stored.Used) + }) + } + }) + + t.Run("logout", func(t *testing.T) { + t.Parallel() + + for _, tc := range []struct { + name string + ephemeral bool + }{ + {"regular", false}, + {"ephemeral", true}, + } { + t.Run(tc.name, func(t *testing.T) { + t.Parallel() + + app := createTestApp(t) + user := app.state.CreateUserForTest("floor-logout-user") + + pak, err := app.state.CreatePreAuthKey(user.TypedID(), false, tc.ephemeral, nil, nil) + require.NoError(t, err) + + machineKey := key.NewMachine().Public() + nodeKey := key.NewNode().Public() + + _, err = app.handleRegisterWithAuthKey(tailcfg.RegisterRequest{ + Auth: &tailcfg.RegisterResponseAuth{AuthKey: pak.Key}, + NodeKey: nodeKey, + Hostinfo: &tailcfg.Hostinfo{Hostname: "floor-logout-node"}, + }, machineKey) + require.NoError(t, err) + + before, ok := app.state.GetNodeByNodeKey(nodeKey) + require.True(t, ok) + require.Equal(t, tc.ephemeral, before.IsEphemeral()) + require.False(t, before.IsExpired()) + + ns := &noiseServer{headscale: app, machineKey: machineKey} + rec := serveRegister(t, ns, tailcfg.RegisterRequest{ + Version: capver.MinSupportedCapabilityVersion - 1, + NodeKey: nodeKey, + Expiry: time.Unix(123, 0), + }) + requireBelowFloorRejected(t, rec) + + after, ok := app.state.GetNodeByNodeKey(nodeKey) + require.True(t, ok, "rejected logout must not delete the node") + assert.False(t, after.IsExpired(), "rejected logout must not expire the node") + }) + } + }) +} + // TestSSHActionRoute_OldPathReturns404 pins the wire-format shape of the // SSH check-action endpoint. Pre-alignment headscale served // /machine/ssh/action/from/{src}/to/{dst}?ssh_user=...; the current