mirror of
https://github.com/juanfont/headscale.git
synced 2026-10-10 08:40:07 +09:00
97c4bf264e
Nothing signals that both waiters parked before FinishAuth, so one pass mostly checks them one after the other.
1220 lines
38 KiB
Go
1220 lines
38 KiB
Go
package hscontrol
|
|
|
|
import (
|
|
"bytes"
|
|
"context"
|
|
"encoding/binary"
|
|
"encoding/json"
|
|
"errors"
|
|
"fmt"
|
|
"io"
|
|
"net/http"
|
|
"net/http/httptest"
|
|
"net/url"
|
|
"slices"
|
|
"strconv"
|
|
"sync"
|
|
"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"
|
|
"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)
|
|
})
|
|
}
|
|
}
|
|
|
|
// 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", src.Uint64(), dst.Uint64())
|
|
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
|
|
}
|
|
|
|
// 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)
|
|
user := app.state.CreateUserForTest("ssh-binding-user")
|
|
|
|
srcCached := putTestNodeInStore(t, app, user, "src-cached")
|
|
dstCached := putTestNodeInStore(t, app, user, "dst-cached")
|
|
srcOther := putTestNodeInStore(t, app, user, "src-other")
|
|
dstOther := putTestNodeInStore(t, app, user, "dst-other")
|
|
|
|
// 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?auth_id=%s",
|
|
srcOther.ID.Uint64(), dstOther.ID.Uint64(), 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)
|
|
user := app.state.CreateUserForTest("ssh-persist-user")
|
|
src := putTestNodeInStore(t, app, user, "src-node")
|
|
dst := putTestNodeInStore(t, app, user, "dst-node")
|
|
|
|
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_RejectsMissingSessionWithoutCheck verifies that without
|
|
// an SSH check covering the pair, a follow-up poll for an unknown auth_id is a
|
|
// genuinely bogus request and is rejected. The re-delegation behaviour for a
|
|
// missing session (issue #3305, exercised end to end with a real client in the
|
|
// servertest package) applies only when the pair is still subject to a check.
|
|
func TestSSHActionHandler_RejectsMissingSessionWithoutCheck(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")
|
|
|
|
// No SSH-check policy is set, so the pair is not subject to a check.
|
|
_, checkFound := app.state.SSHCheckParams(src.ID, dst.ID)
|
|
require.False(t, checkFound, "test setup: pair must not be subject to a check")
|
|
|
|
ns := &noiseServer{headscale: app, machineKey: dst.MachineKey}
|
|
|
|
missing := types.MustAuthID()
|
|
|
|
rec := httptest.NewRecorder()
|
|
ns.SSHActionHandler(rec, newSSHActionFollowUpRequest(t, src.ID, dst.ID, missing))
|
|
require.Equal(t, http.StatusBadRequest, rec.Code,
|
|
"a bogus auth_id with no active check must be rejected, body=%s", rec.Body.String())
|
|
}
|
|
|
|
// 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 _, compress := range []string{"", util.ZstdCompression} {
|
|
t.Run("compress="+compress, func(t *testing.T) {
|
|
t.Parallel()
|
|
|
|
app := createTestApp(t)
|
|
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)
|
|
})
|
|
}
|
|
}
|
|
|
|
// 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())
|
|
}
|