mirror of
https://github.com/juanfont/headscale.git
synced 2026-10-10 16:50:07 +09:00
0c9186d4a3
It opens a deleted node's stream in place of the initial map.
1234 lines
39 KiB
Go
1234 lines
39 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 _, 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())
|
|
}
|