Files
headscale/hscontrol/noise_test.go
T
Kristoffer Dalby 957a332d5d policy/v2: match the login user in SSHCheckParams
The client picks the check rule by login user; the server took the
first rule for the node pair, so a root login could get a 12h
localpart period instead of "always", or be approved after its rule
was removed while another user's rule remained. Hold URLs now carry
the concrete user: tailssh never expanded the encoded $LOCAL_USER.

Updates #3508
2026-10-08 10:29:23 +02:00

1346 lines
42 KiB
Go

package hscontrol
import (
"bytes"
"context"
"encoding/binary"
"encoding/json"
"errors"
"fmt"
"io"
"net/http"
"net/http/httptest"
"net/url"
"slices"
"strconv"
"sync"
"strings"
"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/rs/zerolog"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"tailscale.com/tailcfg"
"tailscale.com/types/key"
"tailscale.com/util/zstdframe"
)
// newNoiseRouterWithBodyLimit builds a chi router with the same body-limit
// middleware used in the real Noise router but wired to a test handler that
// captures the [io.ReadAll] result. This lets us verify the limit without
// needing a full [Headscale] instance.
func newNoiseRouterWithBodyLimit(readBody *[]byte, readErr *error) http.Handler {
r := chi.NewRouter()
r.Use(func(next http.Handler) http.Handler {
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
r.Body = http.MaxBytesReader(w, r.Body, noiseBodyLimit)
next.ServeHTTP(w, r)
})
})
handler := func(w http.ResponseWriter, r *http.Request) {
*readBody, *readErr = io.ReadAll(r.Body)
if *readErr != nil {
http.Error(w, "body too large", http.StatusRequestEntityTooLarge)
return
}
w.WriteHeader(http.StatusOK)
}
r.Post("/machine/map", handler)
r.Post("/machine/register", handler)
return r
}
func TestNoiseBodyLimit_MapEndpoint(t *testing.T) {
t.Parallel()
t.Run("normal_map_request", func(t *testing.T) {
t.Parallel()
var body []byte
var readErr error
router := newNoiseRouterWithBodyLimit(&body, &readErr)
mapReq := tailcfg.MapRequest{Version: 100, Stream: true}
payload, err := json.Marshal(mapReq)
require.NoError(t, err)
req := httptest.NewRequestWithContext(context.Background(), http.MethodPost, "/machine/map", bytes.NewReader(payload))
rec := httptest.NewRecorder()
router.ServeHTTP(rec, req)
require.NoError(t, readErr)
assert.Equal(t, http.StatusOK, rec.Code)
assert.Len(t, body, len(payload))
})
t.Run("oversized_body_rejected", func(t *testing.T) {
t.Parallel()
var body []byte
var readErr error
router := newNoiseRouterWithBodyLimit(&body, &readErr)
oversized := bytes.Repeat([]byte("x"), int(noiseBodyLimit)+1)
req := httptest.NewRequestWithContext(context.Background(), http.MethodPost, "/machine/map", bytes.NewReader(oversized))
rec := httptest.NewRecorder()
router.ServeHTTP(rec, req)
require.Error(t, readErr)
assert.Equal(t, http.StatusRequestEntityTooLarge, rec.Code)
assert.LessOrEqual(t, len(body), int(noiseBodyLimit))
})
}
func TestNoiseBodyLimit_RegisterEndpoint(t *testing.T) {
t.Parallel()
t.Run("normal_register_request", func(t *testing.T) {
t.Parallel()
var body []byte
var readErr error
router := newNoiseRouterWithBodyLimit(&body, &readErr)
regReq := tailcfg.RegisterRequest{Version: 100}
payload, err := json.Marshal(regReq)
require.NoError(t, err)
req := httptest.NewRequestWithContext(context.Background(), http.MethodPost, "/machine/register", bytes.NewReader(payload))
rec := httptest.NewRecorder()
router.ServeHTTP(rec, req)
require.NoError(t, readErr)
assert.Equal(t, http.StatusOK, rec.Code)
assert.Len(t, body, len(payload))
})
t.Run("oversized_body_rejected", func(t *testing.T) {
t.Parallel()
var body []byte
var readErr error
router := newNoiseRouterWithBodyLimit(&body, &readErr)
oversized := bytes.Repeat([]byte("x"), int(noiseBodyLimit)+1)
req := httptest.NewRequestWithContext(context.Background(), http.MethodPost, "/machine/register", bytes.NewReader(oversized))
rec := httptest.NewRecorder()
router.ServeHTTP(rec, req)
require.Error(t, readErr)
assert.Equal(t, http.StatusRequestEntityTooLarge, rec.Code)
assert.LessOrEqual(t, len(body), int(noiseBodyLimit))
})
}
func TestNoiseBodyLimit_AtExactLimit(t *testing.T) {
t.Parallel()
var body []byte
var readErr error
router := newNoiseRouterWithBodyLimit(&body, &readErr)
payload := bytes.Repeat([]byte("a"), int(noiseBodyLimit))
req := httptest.NewRequestWithContext(context.Background(), http.MethodPost, "/machine/map", bytes.NewReader(payload))
rec := httptest.NewRecorder()
router.ServeHTTP(rec, req)
require.NoError(t, readErr)
assert.Equal(t, http.StatusOK, rec.Code)
assert.Len(t, body, int(noiseBodyLimit))
}
// TestPollNetMapHandler_OversizedBody calls the real handler with a
// [http.MaxBytesReader]-wrapped body to verify it fails gracefully (json decode
// error on truncated data) rather than consuming unbounded memory.
func TestPollNetMapHandler_OversizedBody(t *testing.T) {
t.Parallel()
ns := &noiseServer{}
oversized := bytes.Repeat([]byte("x"), int(noiseBodyLimit)+1)
req := httptest.NewRequestWithContext(context.Background(), http.MethodPost, "/machine/map", bytes.NewReader(oversized))
rec := httptest.NewRecorder()
req.Body = http.MaxBytesReader(rec, req.Body, noiseBodyLimit)
ns.PollNetMapHandler(rec, req)
// Body is truncated → [json.Decoder.Decode] fails → [httpError] returns 500.
assert.Equal(t, http.StatusInternalServerError, rec.Code)
}
// TestRegistrationHandler_OversizedBody calls the real handler with a
// [http.MaxBytesReader]-wrapped body to verify it returns an error response
// rather than consuming unbounded memory.
func TestRegistrationHandler_OversizedBody(t *testing.T) {
t.Parallel()
ns := &noiseServer{}
oversized := bytes.Repeat([]byte("x"), int(noiseBodyLimit)+1)
req := httptest.NewRequestWithContext(context.Background(), http.MethodPost, "/machine/register", bytes.NewReader(oversized))
rec := httptest.NewRecorder()
req.Body = http.MaxBytesReader(rec, req.Body, noiseBodyLimit)
ns.RegistrationHandler(rec, req)
// [json.Decoder.Decode] returns [http.MaxBytesError] before any field is
// decoded, so [rejectUnsupported] sees version 0 and answers 400 before
// the decode error is reported.
assert.Equal(t, http.StatusBadRequest, rec.Code)
}
// serveRegister guards against panics so a handler that reaches a nil
// dependency fails its own row instead of the whole test binary. body is
// any so a [json.RawMessage] can carry a request that fails to decode.
func serveRegister(t *testing.T, ns *noiseServer, body any) *httptest.ResponseRecorder {
t.Helper()
payload, err := json.Marshal(body)
require.NoError(t, err)
req := httptest.NewRequestWithContext(context.Background(), http.MethodPost, "/machine/register", bytes.NewReader(payload))
rec := httptest.NewRecorder()
require.NotPanics(t, func() {
ns.RegistrationHandler(rec, 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))
})
}
})
}
})
// A request that passes the floor but fails to decode must be answered
// with RegisterResponse.Error before anything reaches the nil headscale.
// NodeKey 1 is a type error, which still leaves Version decoded.
t.Run("decode_error_at_floor", func(t *testing.T) {
t.Parallel()
body := json.RawMessage(fmt.Sprintf(`{"Version":%d,"NodeKey":1}`, capver.MinSupportedCapabilityVersion))
ns := &noiseServer{machineKey: key.NewMachine().Public()}
rec := serveRegister(t, ns, body)
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.NotEmpty(t, resp.Error)
})
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
// endpoint is /machine/ssh/action/{src}/to/{dst}?local_user=.... If
// someone re-adds the old route shape, this fails.
func TestSSHActionRoute_OldPathReturns404(t *testing.T) {
t.Parallel()
r := chi.NewRouter()
r.Route("/machine", func(r chi.Router) {
r.Get("/ssh/action/{src_node_id}/to/{dst_node_id}", func(w http.ResponseWriter, _ *http.Request) {
w.WriteHeader(http.StatusOK)
})
})
cases := []struct {
name string
path string
want int
}{
{"new", "/machine/ssh/action/1/to/2", http.StatusOK},
{"old-with-from", "/machine/ssh/action/from/1/to/2", http.StatusNotFound},
}
for _, tc := range cases {
t.Run(tc.name, func(t *testing.T) {
t.Parallel()
req := httptest.NewRequestWithContext(context.Background(), http.MethodGet, tc.path, nil)
rec := httptest.NewRecorder()
r.ServeHTTP(rec, req)
assert.Equal(t, tc.want, rec.Code)
})
}
}
// sshTestLocalUser is the non-root local user the SSH action tests log in as.
const sshTestLocalUser = "alice"
// newSSHActionRequest builds an httptest request with the chi URL params
// [noiseServer.SSHActionHandler] reads (src_node_id and dst_node_id), so the handler
// can be exercised directly without going through the chi router.
func newSSHActionRequest(t *testing.T, src, dst types.NodeID) *http.Request {
t.Helper()
url := fmt.Sprintf("/machine/ssh/action/%d/to/%d?local_user=%s",
src.Uint64(), dst.Uint64(), sshTestLocalUser)
req := httptest.NewRequestWithContext(context.Background(), http.MethodGet, url, nil)
rctx := chi.NewRouteContext()
rctx.URLParams.Add("src_node_id", strconv.FormatUint(src.Uint64(), 10))
rctx.URLParams.Add("dst_node_id", strconv.FormatUint(dst.Uint64(), 10))
req = req.WithContext(context.WithValue(req.Context(), chi.RouteCtxKey, rctx))
return req
}
// putTestNodeInStore creates a node via the database test helper and
// also stages it into the in-memory [state.NodeStore] so handlers that read
// [state.NodeStore]-backed APIs (e.g. [state.State.GetNodeByID]) can see it.
func putTestNodeInStore(t *testing.T, app *Headscale, user *types.User, hostname string) *types.Node {
t.Helper()
node := app.state.CreateNodeForTest(user, hostname)
app.state.PutNodeInStoreForTest(*node)
return node
}
// sshCheckPolicy subjects every same-user pair of userName's nodes to an
// SSH check.
func sshCheckPolicy(userName string) string {
return fmt.Sprintf(`{"ssh": [{
"action": "check",
"src": [%q],
"dst": ["autogroup:self"],
"users": ["autogroup:nonroot"]
}]}`, userName+"@")
}
// putSSHCheckNodes stages one node per hostname for a new user and sets
// [sshCheckPolicy], so the SSH action handler holds and delegates for any
// pair of them.
func putSSHCheckNodes(t *testing.T, app *Headscale, userName string, hostnames ...string) []*types.Node {
t.Helper()
user := app.state.CreateUserForTest(userName)
require.NoError(t, app.state.UpdatePolicyManagerUsersForTest())
nodes := make([]*types.Node, 0, len(hostnames))
for _, hostname := range hostnames {
node := app.state.CreateRegisteredNodeForTest(user, hostname)
node.User = user
app.state.PutNodeInStoreForTest(*node)
nodes = append(nodes, node)
}
require.NoError(t, app.state.UpdatePolicyManagerNodesForTest())
_, err := app.state.SetPolicy([]byte(sshCheckPolicy(userName)))
require.NoError(t, err)
_, checkFound := app.state.SSHCheckParams(nodes[0].ID, nodes[len(nodes)-1].ID, sshTestLocalUser)
require.True(t, checkFound, "test setup: nodes must be subject to an SSH check")
return nodes
}
// TestSSHActionHandler_RejectsRogueMachineKey verifies that the SSH
// check action endpoint rejects a Noise session whose machine key does
// not match the dst node.
func TestSSHActionHandler_RejectsRogueMachineKey(t *testing.T) {
t.Parallel()
app := createTestApp(t)
user := app.state.CreateUserForTest("ssh-handler-user")
src := putTestNodeInStore(t, app, user, "src-node")
dst := putTestNodeInStore(t, app, user, "dst-node")
// [noiseServer] carries the wrong machine key — a fresh throwaway key,
// not dst.MachineKey.
rogue := key.NewMachine().Public()
require.NotEqual(t, dst.MachineKey, rogue, "test sanity: rogue key must differ from dst")
ns := &noiseServer{
headscale: app,
machineKey: rogue,
}
rec := httptest.NewRecorder()
ns.SSHActionHandler(rec, newSSHActionRequest(t, src.ID, dst.ID))
assert.Equal(t, http.StatusUnauthorized, rec.Code,
"rogue machine key must be rejected with 401")
// And the auth cache must not have been mutated by the rejected request.
if last, ok := app.state.GetLastSSHAuth(src.ID, dst.ID); ok {
t.Fatalf("rejected SSH action must not record lastSSHAuth, got %v", last)
}
}
// TestSSHActionHandler_RejectsUnknownDst verifies that the handler
// rejects a request for a dst_node_id that does not exist with 404.
func TestSSHActionHandler_RejectsUnknownDst(t *testing.T) {
t.Parallel()
app := createTestApp(t)
user := app.state.CreateUserForTest("ssh-handler-unknown-user")
src := putTestNodeInStore(t, app, user, "src-node")
ns := &noiseServer{
headscale: app,
machineKey: key.NewMachine().Public(),
}
rec := httptest.NewRecorder()
ns.SSHActionHandler(rec, newSSHActionRequest(t, src.ID, 9999))
assert.Equal(t, http.StatusNotFound, rec.Code,
"unknown dst node id must be rejected with 404")
}
// TestSSHActionFollowUp_RejectsBindingMismatch verifies that the
// follow-up handler refuses to honour an auth_id whose cached binding
// does not match the (src, dst) pair on the request URL. Without this
// check an attacker holding any auth_id could route its verdict to a
// different node pair.
func TestSSHActionFollowUp_RejectsBindingMismatch(t *testing.T) {
t.Parallel()
app := createTestApp(t)
nodes := putSSHCheckNodes(t, app, "ssh-binding-user",
"src-cached", "dst-cached", "src-other", "dst-other")
srcCached, dstCached, srcOther, dstOther := nodes[0], nodes[1], nodes[2], nodes[3]
// Mint an SSH-check auth request bound to (srcCached, dstCached).
authID := types.MustAuthID()
app.state.SetAuthCacheEntry(
authID,
types.NewSSHCheckAuthRequest(srcCached.ID, dstCached.ID),
)
// Build a follow-up that claims to be for (srcOther, dstOther) but
// reuses the bound auth_id. The Noise machineKey matches dstOther so
// the outer machine-key check passes — only the binding check
// should reject it.
ns := &noiseServer{
headscale: app,
machineKey: dstOther.MachineKey,
}
url := fmt.Sprintf(
"/machine/ssh/action/%d/to/%d?local_user=%s&auth_id=%s",
srcOther.ID.Uint64(), dstOther.ID.Uint64(), sshTestLocalUser, authID.String(),
)
req := httptest.NewRequestWithContext(context.Background(), http.MethodGet, url, nil)
rctx := chi.NewRouteContext()
rctx.URLParams.Add("src_node_id", strconv.FormatUint(srcOther.ID.Uint64(), 10))
rctx.URLParams.Add("dst_node_id", strconv.FormatUint(dstOther.ID.Uint64(), 10))
req = req.WithContext(context.WithValue(req.Context(), chi.RouteCtxKey, rctx))
rec := httptest.NewRecorder()
ns.SSHActionHandler(rec, req)
assert.Equal(t, http.StatusUnauthorized, rec.Code,
"binding mismatch must be rejected with 401")
}
// TestOverrideRemoteAddr asserts the middleware used inside the Noise
// tunnel pins r.RemoteAddr to the value captured from the outer
// (pre-hijack) request, so /machine/* requests log the trusted-proxy
// resolved client IP instead of the hijacked TCP socket's loopback peer.
func TestOverrideRemoteAddr(t *testing.T) {
t.Parallel()
const clientAddr = "192.168.91.240"
r := chi.NewRouter()
r.Use(overrideRemoteAddr(clientAddr))
var observed string
r.Get("/x", func(w http.ResponseWriter, r *http.Request) {
observed = r.RemoteAddr
w.WriteHeader(http.StatusOK)
})
req := httptest.NewRequestWithContext(t.Context(), http.MethodGet, "/x", nil)
req.RemoteAddr = "127.0.0.1:44388"
r.ServeHTTP(httptest.NewRecorder(), req)
assert.Equal(t, clientAddr, observed)
}
// TestSSHActionHoldAndDelegate_PersistsAuthSession guards the happy path: the
// initial SSH-check poll returns a HoldAndDelegate URL carrying an auth_id, and
// that auth session must remain in the cache for the follow-up poll to find.
func TestSSHActionHoldAndDelegate_PersistsAuthSession(t *testing.T) {
t.Parallel()
app := createTestApp(t)
nodes := putSSHCheckNodes(t, app, "ssh-persist-user", "src-node", "dst-node")
src, dst := nodes[0], nodes[1]
ns := &noiseServer{headscale: app, machineKey: dst.MachineKey}
rec := httptest.NewRecorder()
ns.SSHActionHandler(rec, newSSHActionRequest(t, src.ID, dst.ID))
require.Equal(t, http.StatusOK, rec.Code, "initial poll body=%s", rec.Body.String())
var action tailcfg.SSHAction
require.NoError(t, json.Unmarshal(rec.Body.Bytes(), &action))
require.NotEmpty(t, action.HoldAndDelegate, "expected HoldAndDelegate, got %+v", action)
u, err := url.Parse(action.HoldAndDelegate)
require.NoError(t, err)
authIDStr := u.Query().Get("auth_id")
require.NotEmpty(t, authIDStr, "HoldAndDelegate URL missing auth_id: %s", action.HoldAndDelegate)
authID, err := types.AuthIDFromString(authIDStr)
require.NoError(t, err)
_, ok := app.state.GetAuthCacheEntry(authID)
require.True(t, ok, "auth session %s must persist after HoldAndDelegate", authID)
}
// TestSSHActionHandler_RejectsWithoutCheck verifies that a pair no check rule
// covers is denied with a 200 Reject on both the initial and the follow-up
// poll. Such a call comes from a client holding a stale check rule; an HTTP
// error would make tailssh retry for up to 30 minutes. Re-delegation for a
// missing session (issue #3305) applies only while a check is required.
// https://github.com/juanfont/headscale/issues/3508
func TestSSHActionHandler_RejectsWithoutCheck(t *testing.T) {
t.Parallel()
app := createTestApp(t)
user := app.state.CreateUserForTest("ssh-nocheck-user")
src := putTestNodeInStore(t, app, user, "src-node")
dst := putTestNodeInStore(t, app, user, "dst-node")
_, checkFound := app.state.SSHCheckParams(src.ID, dst.ID, sshTestLocalUser)
require.False(t, checkFound, "test setup: pair must not be subject to a check")
ns := &noiseServer{headscale: app, machineKey: dst.MachineKey}
for name, req := range map[string]*http.Request{
"initial": newSSHActionRequest(t, src.ID, dst.ID),
"follow-up": newSSHActionFollowUpRequest(t, src.ID, dst.ID, types.MustAuthID()),
} {
rec := httptest.NewRecorder()
ns.SSHActionHandler(rec, req)
require.Equal(t, http.StatusOK, rec.Code, "%s: body=%s", name, rec.Body.String())
var action tailcfg.SSHAction
require.NoError(t, json.Unmarshal(rec.Body.Bytes(), &action))
assert.True(t, action.Reject, "%s: must reject, got %+v", name, action)
assert.Empty(t, action.HoldAndDelegate, "%s: must not delegate", name)
}
}
// TestSSHActionFollowUp_RejectsVerdictAfterCheckRemoved verifies the check is
// re-evaluated once the user authenticates: a rule removed while the login
// prompt was open must not grant access.
// https://github.com/juanfont/headscale/issues/3508
func TestSSHActionFollowUp_RejectsVerdictAfterCheckRemoved(t *testing.T) {
t.Parallel()
const userName = "ssh-verdict-user"
nonrootOnly := sshCheckPolicy(userName)
rootAndNonroot := strings.Replace(nonrootOnly,
`["autogroup:nonroot"]`, `["root", "autogroup:nonroot"]`, 1)
for name, tc := range map[string]struct {
after string // policy once the user has authenticated; "" keeps it
localUser string
}{
"rule kept": {"", sshTestLocalUser},
"rule removed": {`{}`, sshTestLocalUser},
// Another check rule still covers the pair, but not for root.
"root rule removed": {nonrootOnly, "root"},
} {
t.Run(name, func(t *testing.T) {
t.Parallel()
app := createTestApp(t)
nodes := putSSHCheckNodes(t, app, userName, "src-node", "dst-node")
src, dst := nodes[0], nodes[1]
_, err := app.state.SetPolicy([]byte(rootAndNonroot))
require.NoError(t, err)
authID := types.MustAuthID()
app.state.SetAuthCacheEntry(authID, types.NewSSHCheckAuthRequest(src.ID, dst.ID))
auth, ok := app.state.GetAuthCacheEntry(authID)
require.True(t, ok)
auth.FinishAuth(types.AuthVerdict{})
if tc.after != "" {
_, err := app.state.SetPolicy([]byte(tc.after))
require.NoError(t, err)
}
ns := &noiseServer{headscale: app, machineKey: dst.MachineKey}
// Call the follow-up directly: the verdict is already in, as if
// the policy changed while the user was authenticating.
action, err := ns.sshActionFollowUp(
t.Context(), zerolog.Nop(), &tailcfg.SSHAction{},
authID.String(), src.ID, dst.ID, tc.localUser,
)
require.NoError(t, err)
kept := tc.after == ""
_, recorded := app.state.GetLastSSHAuth(src.ID, dst.ID)
assert.Equal(t, kept, action.Accept, "accept, got %+v", action)
assert.Equal(t, !kept, action.Reject, "reject, got %+v", action)
assert.Equal(t, kept, recorded, "auth recorded for auto-approval")
})
}
}
// TestTS2021Route_AcceptsGETAndPOST reproduces a regression where the
// browser/WASM control client could not connect. Tailscale's JS/WASM control
// client opens /ts2021 as a WebSocket, which is an HTTP GET upgrade; the native
// Go client uses an HTTP POST upgrade. The gorilla->chi router migration
// registered /ts2021 for POST only, so the GET WebSocket handshake was rejected
// with 405 Method Not Allowed by the router before it could reach
// NoiseUpgradeHandler. Both methods must route to the handler.
//
// NoiseUpgradeHandler dispatches on the Upgrade header, not the HTTP method, so
// once the route is reachable the handler handles both upgrade styles. The
// httptest recorder is not an http.Hijacker, so the upgrade itself fails past
// the router (501 for the WebSocket path, 400 for the native path) — the point
// is only that neither is 405, i.e. the router no longer rejects GET early.
func TestTS2021Route_AcceptsGETAndPOST(t *testing.T) {
t.Parallel()
handler := createTestApp(t).HTTPHandler()
tests := []struct {
name string
method string
headers map[string]string
}{
{
name: "websocket_get_from_wasm_client",
method: http.MethodGet,
headers: map[string]string{
"Connection": "Upgrade",
"Upgrade": "websocket",
"Sec-WebSocket-Version": "13",
"Sec-WebSocket-Key": "dGhlIHNhbXBsZSBub25jZQ==",
"Sec-WebSocket-Protocol": "tailscale-control-protocol",
},
},
{
name: "native_post_upgrade",
method: http.MethodPost,
headers: map[string]string{
"Connection": "upgrade",
"Upgrade": "tailscale-control-protocol",
"X-Tailscale-Handshake": "AAAA",
},
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
t.Parallel()
req := httptest.NewRequestWithContext(context.Background(), tt.method,
"/ts2021?X-Tailscale-Handshake=AAAA", nil)
for k, v := range tt.headers {
req.Header.Set(k, v)
}
rec := httptest.NewRecorder()
handler.ServeHTTP(rec, req)
assert.NotEqual(t, http.StatusMethodNotAllowed, rec.Code,
"%s /ts2021 must reach NoiseUpgradeHandler, not be rejected by the router with 405",
tt.method)
})
}
}
// newSSHActionFollowUpRequest is like newSSHActionRequest but carries the
// auth_id query parameter that marks a follow-up poll.
func newSSHActionFollowUpRequest(t *testing.T, src, dst types.NodeID, authID types.AuthID) *http.Request {
t.Helper()
req := newSSHActionRequest(t, src, dst)
q := req.URL.Query()
q.Set("auth_id", authID.String())
req.URL.RawQuery = q.Encode()
return req
}
var errSSHCheckRejectedForTest = errors.New("ssh check rejected for test")
// sshVerdictCases are the two verdicts a check session can resolve to.
var sshVerdictCases = []struct {
name string
verdict types.AuthVerdict
accept bool
}{
{name: "accept", verdict: types.AuthVerdict{}, accept: true},
{name: "reject", verdict: types.AuthVerdict{Err: errSSHCheckRejectedForTest}, accept: false},
}
// sshVerdictFixture is a same-user (src, dst) pair under an SSH check with
// checkPeriod "always". Period 0 never auto-approves from the ledger, so any
// Accept not backed by a verdict is a replay.
type sshVerdictFixture struct {
ns *noiseServer
src, dst types.NodeID
}
func newSSHVerdictFixture(t *testing.T) *sshVerdictFixture {
t.Helper()
app := createTestApp(t)
user := app.state.CreateUserForTest("ssh-verdict-user")
require.NoError(t, app.state.UpdatePolicyManagerUsersForTest())
var ids [2]types.NodeID
for i, name := range []string{"src-node", "dst-node"} {
node := app.state.CreateRegisteredNodeForTest(user, name)
// autogroup:self compares hydrated users.
node.User = user
// SaveNode refreshes the policy manager's nodes, which
// SSHCheckParams resolves against.
_, _, err := app.state.SaveNode(node.View())
require.NoError(t, err)
ids[i] = node.ID
}
_, err := app.state.SetPolicy(fmt.Appendf(nil, `{
"ssh": [{
"action": "check",
"checkPeriod": "always",
"src": [%q],
"dst": ["autogroup:self"],
"users": ["autogroup:nonroot"]
}]
}`, user.Name+"@"))
require.NoError(t, err)
period, checkFound := app.state.SSHCheckParams(ids[0], ids[1])
require.True(t, checkFound, "test setup: pair must be subject to a check")
require.Zero(t, period, "test setup: checkPeriod must be always")
dst, ok := app.state.GetNodeByID(ids[1])
require.True(t, ok)
return &sshVerdictFixture{
ns: &noiseServer{headscale: app, machineKey: dst.MachineKey()},
src: ids[0],
dst: ids[1],
}
}
// mint runs the initial poll and returns the check session it created.
func (f *sshVerdictFixture) mint(t *testing.T) (types.AuthID, *types.AuthRequest) {
t.Helper()
rec := httptest.NewRecorder()
f.ns.SSHActionHandler(rec, newSSHActionRequest(t, f.src, f.dst))
authID := requireSSHHold(t, sshActionFromRecorder(t, rec))
auth, ok := f.ns.headscale.state.GetAuthCacheEntry(authID)
require.True(t, ok, "minted session must be cached")
return authID, auth
}
// cancellableFollowUp returns a follow-up request and the cancel for its
// context, derived from the request's own so chi's route values survive.
func (f *sshVerdictFixture) cancellableFollowUp(
t *testing.T,
authID types.AuthID,
) (*http.Request, context.CancelFunc) {
t.Helper()
req := newSSHActionFollowUpRequest(t, f.src, f.dst, authID)
ctx, cancel := context.WithCancel(req.Context())
return req.WithContext(ctx), cancel
}
func (f *sshVerdictFixture) serve(req *http.Request) *httptest.ResponseRecorder {
rec := httptest.NewRecorder()
f.ns.SSHActionHandler(rec, req)
return rec
}
func (f *sshVerdictFixture) followUp(t *testing.T, authID types.AuthID) *httptest.ResponseRecorder {
t.Helper()
return f.serve(newSSHActionFollowUpRequest(t, f.src, f.dst, authID))
}
func sshActionFromRecorder(t *testing.T, rec *httptest.ResponseRecorder) tailcfg.SSHAction {
t.Helper()
require.Equal(t, http.StatusOK, rec.Code, "body=%s", rec.Body.String())
var action tailcfg.SSHAction
require.NoError(t, json.Unmarshal(rec.Body.Bytes(), &action))
return action
}
// carriesSSHVerdict reports whether action is the answer to the verdict.
func carriesSSHVerdict(action tailcfg.SSHAction, accept bool) bool {
if action.HoldAndDelegate != "" || action.Accept == action.Reject {
return false
}
return action.Accept == accept
}
// requireSSHHold asserts action re-delegates and returns its fresh auth_id.
func requireSSHHold(t *testing.T, action tailcfg.SSHAction) types.AuthID {
t.Helper()
require.False(t, action.Accept, "expected HoldAndDelegate, got Accept: %+v", action)
require.False(t, action.Reject, "expected HoldAndDelegate, got Reject: %+v", action)
require.NotEmpty(t, action.HoldAndDelegate, "expected HoldAndDelegate: %+v", action)
u, err := url.Parse(action.HoldAndDelegate)
require.NoError(t, err)
authID, err := types.AuthIDFromString(u.Query().Get("auth_id"))
require.NoError(t, err)
return authID
}
// TestSSHActionFollowUp_ConsumedVerdictNotReplayed guards the one-shot
// verdict channel: after a follow-up consumed the verdict, a second follow-up
// on the same auth_id must re-decide instead of reading the closed channel's
// zero value, which Accept() reports as success.
func TestSSHActionFollowUp_ConsumedVerdictNotReplayed(t *testing.T) {
t.Parallel()
for _, tc := range sshVerdictCases {
t.Run(tc.name, func(t *testing.T) {
t.Parallel()
f := newSSHVerdictFixture(t)
authID, auth := f.mint(t)
auth.FinishAuth(tc.verdict)
first := sshActionFromRecorder(t, f.followUp(t, authID))
require.True(t, carriesSSHVerdict(first, tc.accept),
"first follow-up must carry the verdict, got %+v", first)
second := sshActionFromRecorder(t, f.followUp(t, authID))
replayID := requireSSHHold(t, second)
assert.NotEqual(t, authID, replayID, "re-delegation must mint a new session")
// A replay that recorded auth would feed auto-approval.
_, recorded := f.ns.headscale.state.GetLastSSHAuth(f.src, f.dst)
assert.Equal(t, tc.accept, recorded, "only an accepted verdict may record auth")
})
}
// With the check gone there is nothing to re-delegate, but the replay
// must still not be answered from the consumed verdict.
for _, tc := range sshVerdictCases {
t.Run(tc.name+"-check-removed", func(t *testing.T) {
t.Parallel()
f := newSSHVerdictFixture(t)
authID, auth := f.mint(t)
auth.FinishAuth(tc.verdict)
first := sshActionFromRecorder(t, f.followUp(t, authID))
require.True(t, carriesSSHVerdict(first, tc.accept),
"first follow-up must carry the verdict, got %+v", first)
_, err := f.ns.headscale.state.SetPolicy([]byte(`{}`))
require.NoError(t, err)
_, checkFound := f.ns.headscale.state.SSHCheckParams(f.src, f.dst)
require.False(t, checkFound, "test setup: pair must no longer be subject to a check")
rec := f.followUp(t, authID)
assert.Equal(t, http.StatusBadRequest, rec.Code,
"replay without a check must be refused, body=%s", rec.Body.String())
})
}
}
// TestSSHActionFollowUp_ConcurrentWaiters parks two follow-ups on one
// session: exactly one may consume the verdict, the other re-decides.
// FinishAuth can land before either waiter parks and nothing signals the
// park, so the body repeats to exercise the both-parked order.
func TestSSHActionFollowUp_ConcurrentWaiters(t *testing.T) {
t.Parallel()
const iterations = 200
for _, tc := range sshVerdictCases {
t.Run(tc.name, func(t *testing.T) {
t.Parallel()
f := newSSHVerdictFixture(t)
for range iterations {
authID, auth := f.mint(t)
var (
wg sync.WaitGroup
recs [2]*httptest.ResponseRecorder
)
for i := range recs {
wg.Go(func() {
recs[i] = f.followUp(t, authID)
})
}
auth.FinishAuth(tc.verdict)
wg.Wait()
var carried, held int
for _, rec := range recs {
action := sshActionFromRecorder(t, rec)
if carriesSSHVerdict(action, tc.accept) {
carried++
continue
}
requireSSHHold(t, action)
held++
}
require.Equal(t, 1, carried, "exactly one waiter must carry the verdict")
require.Equal(t, 1, held, "the other waiter must re-delegate")
}
})
}
}
// TestSSHActionFollowUp_CancelledWaiterLeavesVerdict: a follow-up that
// returns before FinishAuth must not consume the verdict, so the retry gets
// it exactly once.
func TestSSHActionFollowUp_CancelledWaiterLeavesVerdict(t *testing.T) {
t.Parallel()
for _, tc := range sshVerdictCases {
t.Run(tc.name, func(t *testing.T) {
t.Parallel()
f := newSSHVerdictFixture(t)
authID, auth := f.mint(t)
req1, cancel := f.cancellableFollowUp(t, authID)
done1 := make(chan *httptest.ResponseRecorder)
go func() {
done1 <- f.serve(req1)
}()
cancel()
rec1 := <-done1
require.Equal(t, http.StatusUnauthorized, rec1.Code,
"cancelled follow-up must return 401, body=%s", rec1.Body.String())
// The handler has returned, so its select could only take ctx.Done.
auth.FinishAuth(tc.verdict)
second := sshActionFromRecorder(t, f.followUp(t, authID))
require.True(t, carriesSSHVerdict(second, tc.accept),
"retry must carry the verdict, got %+v", second)
third := sshActionFromRecorder(t, f.followUp(t, authID))
requireSSHHold(t, third)
})
}
}
// TestSSHActionFollowUp_CancelRacingVerdictAtMostOnce makes both select cases
// ready before the handler parks. Either outcome is allowed; the verdict
// must reach at most one response.
func TestSSHActionFollowUp_CancelRacingVerdictAtMostOnce(t *testing.T) {
t.Parallel()
const iterations = 200
for _, tc := range sshVerdictCases {
t.Run(tc.name, func(t *testing.T) {
t.Parallel()
f := newSSHVerdictFixture(t)
for range iterations {
authID, auth := f.mint(t)
auth.FinishAuth(tc.verdict)
req1, cancel := f.cancellableFollowUp(t, authID)
cancel()
rec1 := f.serve(req1)
retry := sshActionFromRecorder(t, f.followUp(t, authID))
if rec1.Code == http.StatusUnauthorized {
require.True(t, carriesSSHVerdict(retry, tc.accept),
"cancelled follow-up left the verdict; retry must carry it, got %+v", retry)
continue
}
first := sshActionFromRecorder(t, rec1)
require.True(t, carriesSSHVerdict(first, tc.accept),
"cancelled follow-up consumed the verdict, got %+v", first)
requireSSHHold(t, retry)
}
})
}
}
// TestSSHActionFollowUp_LostResponseRedecides: the client never saw the
// Accept (dropped response) and retries. The retry must re-decide through a
// fresh session, which then completes normally.
func TestSSHActionFollowUp_LostResponseRedecides(t *testing.T) {
t.Parallel()
f := newSSHVerdictFixture(t)
authID, auth := f.mint(t)
auth.FinishAuth(types.AuthVerdict{})
// The Accept response is lost on the way to the client.
_ = f.followUp(t, authID)
retry := sshActionFromRecorder(t, f.followUp(t, authID))
freshID := requireSSHHold(t, retry)
require.NotEqual(t, authID, freshID)
fresh, ok := f.ns.headscale.state.GetAuthCacheEntry(freshID)
require.True(t, ok, "re-delegated session must be cached")
fresh.FinishAuth(types.AuthVerdict{})
final := sshActionFromRecorder(t, f.followUp(t, freshID))
assert.True(t, carriesSSHVerdict(final, true),
"fresh session must complete with its verdict, got %+v", final)
}
// newMapRequest builds a streaming [tailcfg.MapRequest] POST for
// /machine/map. Version is mandatory: [rejectUnsupported] runs before the
// handler looks the node up, and a zero version is rejected with 400.
func newMapRequest(t *testing.T, req tailcfg.MapRequest) *http.Request {
t.Helper()
body, err := json.Marshal(req)
require.NoError(t, err)
return httptest.NewRequestWithContext(context.Background(), http.MethodPost, "/machine/map", bytes.NewReader(body))
}
// decodeMapResponse reads a map response frame the way a Tailscale client
// does: a little-endian length prefix followed by a body that is zstd-framed
// whenever the request asked for compression.
func decodeMapResponse(t *testing.T, compress string, body []byte) tailcfg.MapResponse {
t.Helper()
require.GreaterOrEqual(t, len(body), reservedResponseHeaderSize, "response too short to carry a length prefix")
size := binary.LittleEndian.Uint32(body[:reservedResponseHeaderSize])
payload := body[reservedResponseHeaderSize:]
require.Len(t, payload, int(size), "length prefix must match the body it precedes")
if compress == util.ZstdCompression {
decoded, err := zstdframe.AppendDecode(nil, payload)
require.NoError(t, err, "client decodes every frame as zstd when it asked for zstd")
payload = decoded
}
var resp tailcfg.MapResponse
require.NoError(t, json.Unmarshal(payload, &resp))
return resp
}
// TestPollNetMapHandler_DeletedNodeGetsExpiredSelf verifies that a streaming
// map request for a node that no longer exists is answered with an expired
// self node instead of a bare 404. A Tailscale client treats every non-200 on
// the map path identically and retries forever with loggedIn still set; only a
// self node whose KeyExpiry is in the past moves it to NeedsLogin.
//
// See: https://github.com/juanfont/headscale/issues/3410
func TestPollNetMapHandler_DeletedNodeGetsExpiredSelf(t *testing.T) {
t.Parallel()
for _, tc := range []struct {
compress string
logTail bool
wantDebug *tailcfg.Debug
}{
{"", false, &tailcfg.Debug{DisableLogTail: true}},
{util.ZstdCompression, false, &tailcfg.Debug{DisableLogTail: true}},
{"", true, nil},
{util.ZstdCompression, true, nil},
} {
compress := tc.compress
t.Run(fmt.Sprintf("compress=%s/logtail=%t", compress, tc.logTail), func(t *testing.T) {
t.Parallel()
app := createTestApp(t)
app.cfg.LogTail.Enabled = tc.logTail
user := app.state.CreateUserForTest("deleted-node-user")
node := putTestNodeInStore(t, app, user, "deleted-node")
nodeView, ok := app.state.GetNodeByID(node.ID)
require.True(t, ok)
_, err := app.state.DeleteNode(nodeView)
require.NoError(t, err)
ns := &noiseServer{headscale: app, machineKey: node.MachineKey}
rec := httptest.NewRecorder()
ns.PollNetMapHandler(rec, newMapRequest(t, tailcfg.MapRequest{
Version: tailcfg.CurrentCapabilityVersion,
NodeKey: node.NodeKey,
Stream: true,
Compress: compress,
}))
require.Equal(t, http.StatusOK, rec.Code, "body=%q", rec.Body.String())
resp := decodeMapResponse(t, compress, rec.Body.Bytes())
require.NotNil(t, resp.Node, "clients reject an initial map response without a node")
assert.Equal(t, node.NodeKey, resp.Node.Key)
assert.Equal(t, time.Unix(1, 0).UTC(), resp.Node.KeyExpiry,
"a fixed ancient KeyExpiry must remain expired despite client clock skew")
assert.True(t, resp.Node.Expired)
// This frame starts the stream, so it must carry the logtail
// instruction the initial map would have.
assert.Equal(t, tc.wantDebug, resp.Debug)
})
}
}
// TestPollNetMapHandler_ForeignMachineKeyStillRejected pins that the
// expired-self response is limited to a genuinely unknown node. A known
// NodeKey presented by the wrong machine key is an impostor, and answering it
// with "your key expired" would wipe the real client's persisted node ID.
func TestPollNetMapHandler_ForeignMachineKeyStillRejected(t *testing.T) {
t.Parallel()
app := createTestApp(t)
user := app.state.CreateUserForTest("impostor-user")
victim := putTestNodeInStore(t, app, user, "victim-node")
impostor := putTestNodeInStore(t, app, user, "impostor-node")
ns := &noiseServer{headscale: app, machineKey: impostor.MachineKey}
rec := httptest.NewRecorder()
ns.PollNetMapHandler(rec, newMapRequest(t, tailcfg.MapRequest{
Version: tailcfg.CurrentCapabilityVersion,
NodeKey: victim.NodeKey,
Stream: true,
}))
assert.Equal(t, http.StatusNotFound, rec.Code, "body=%q", rec.Body.String())
}