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
This commit is contained in:
Kristoffer Dalby
2026-10-02 10:51:40 +00:00
committed by Kristoffer Dalby
parent dceb584c89
commit 957a332d5d
9 changed files with 416 additions and 82 deletions
+15 -8
View File
@@ -430,6 +430,7 @@ func (ns *noiseServer) SSHActionHandler(
req.Context(),
reqLog,
srcNodeID, dstNodeID,
req.URL.Query().Get("local_user"),
req.URL.Query().Get("auth_id"),
)
if err != nil {
@@ -470,7 +471,7 @@ func (ns *noiseServer) sshAction(
ctx context.Context,
reqLog zerolog.Logger,
srcNodeID, dstNodeID types.NodeID,
authIDStr string,
localUser, authIDStr string,
) (*tailcfg.SSHAction, error) {
action := tailcfg.SSHAction{
AllowAgentForwarding: true,
@@ -480,8 +481,10 @@ func (ns *noiseServer) sshAction(
// Look up check params from the server's own policy rather than
// trusting URL parameters, which the client could tamper with.
// local_user only narrows which rule applies, and it comes from dst,
// the node enforcing the login.
checkPeriod, checkFound := ns.headscale.state.SSHCheckParams(
srcNodeID, dstNodeID,
srcNodeID, dstNodeID, localUser,
)
// Clients only call back for check rules they were sent. Without one in
@@ -494,7 +497,7 @@ func (ns *noiseServer) sshAction(
if authIDStr != "" {
return ns.sshActionFollowUp(
ctx, reqLog, &action, authIDStr,
srcNodeID, dstNodeID,
srcNodeID, dstNodeID, localUser,
)
}
@@ -515,7 +518,7 @@ func (ns *noiseServer) sshAction(
}
// No auto-approval — create an auth session and hold.
return ns.sshActionHoldAndDelegate(reqLog, &action, srcNodeID, dstNodeID)
return ns.sshActionHoldAndDelegate(reqLog, &action, srcNodeID, dstNodeID, localUser)
}
// sshActionDeny rejects a check the current policy does not require. It is a
@@ -536,11 +539,11 @@ func (ns *noiseServer) sshActionHoldAndDelegate(
reqLog zerolog.Logger,
action *tailcfg.SSHAction,
srcNodeID, dstNodeID types.NodeID,
localUser string,
) (*tailcfg.SSHAction, error) {
holdURL, err := url.Parse(
ns.headscale.cfg.ServerURL +
"/machine/ssh/action/$SRC_NODE_ID/to/$DST_NODE_ID" +
"?local_user=$LOCAL_USER",
"/machine/ssh/action/$SRC_NODE_ID/to/$DST_NODE_ID",
)
if err != nil {
return nil, NewHTTPError(
@@ -566,7 +569,10 @@ func (ns *noiseServer) sshActionHoldAndDelegate(
authURL := ns.headscale.authProvider.AuthURL(authID)
// The concrete user, not $LOCAL_USER: Encode escapes the placeholder
// and tailssh only expands it literally.
q := holdURL.Query()
q.Set("local_user", localUser)
q.Set("auth_id", authID.String())
holdURL.RawQuery = q.Encode()
@@ -597,6 +603,7 @@ func (ns *noiseServer) sshActionFollowUp(
action *tailcfg.SSHAction,
authIDStr string,
srcNodeID, dstNodeID types.NodeID,
localUser string,
) (*tailcfg.SSHAction, error) {
authID, err := types.AuthIDFromString(authIDStr)
if err != nil {
@@ -616,7 +623,7 @@ func (ns *noiseServer) sshActionFollowUp(
reqLog.Info().Caller().Msg(logMsg)
return ns.sshActionHoldAndDelegate(
reqLog, action, srcNodeID, dstNodeID,
reqLog, action, srcNodeID, dstNodeID, localUser,
)
}
@@ -691,7 +698,7 @@ func (ns *noiseServer) sshActionFollowUp(
// The policy may have changed while the user authenticated, and the
// client won't drop a connection still waiting on its check.
if _, ok := ns.headscale.state.SSHCheckParams(srcNodeID, dstNodeID); !ok {
if _, ok := ns.headscale.state.SSHCheckParams(srcNodeID, dstNodeID, localUser); !ok {
return sshActionDeny(reqLog, action), nil
}
+35 -15
View File
@@ -14,6 +14,7 @@ import (
"slices"
"strconv"
"sync"
"strings"
"testing"
"time"
@@ -448,13 +449,17 @@ func TestSSHActionRoute_OldPathReturns404(t *testing.T) {
}
}
// 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", src.Uint64(), dst.Uint64())
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()
@@ -510,7 +515,7 @@ func putSSHCheckNodes(t *testing.T, app *Headscale, userName string, hostnames .
_, err := app.state.SetPolicy([]byte(sshCheckPolicy(userName)))
require.NoError(t, err)
_, checkFound := app.state.SSHCheckParams(nodes[0].ID, nodes[len(nodes)-1].ID)
_, 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
@@ -601,8 +606,8 @@ func TestSSHActionFollowUp_RejectsBindingMismatch(t *testing.T) {
}
url := fmt.Sprintf(
"/machine/ssh/action/%d/to/%d?auth_id=%s",
srcOther.ID.Uint64(), dstOther.ID.Uint64(), authID.String(),
"/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)
@@ -693,7 +698,7 @@ func TestSSHActionHandler_RejectsWithoutCheck(t *testing.T) {
src := putTestNodeInStore(t, app, user, "src-node")
dst := putTestNodeInStore(t, app, user, "dst-node")
_, checkFound := app.state.SSHCheckParams(src.ID, dst.ID)
_, 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}
@@ -720,17 +725,31 @@ func TestSSHActionHandler_RejectsWithoutCheck(t *testing.T) {
func TestSSHActionFollowUp_RejectsVerdictAfterCheckRemoved(t *testing.T) {
t.Parallel()
for name, removeRule := range map[string]bool{
"rule kept": false,
"rule removed": true,
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, "ssh-verdict-user", "src-node", "dst-node")
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))
@@ -738,8 +757,8 @@ func TestSSHActionFollowUp_RejectsVerdictAfterCheckRemoved(t *testing.T) {
require.True(t, ok)
auth.FinishAuth(types.AuthVerdict{})
if removeRule {
_, err := app.state.SetPolicy([]byte(`{}`))
if tc.after != "" {
_, err := app.state.SetPolicy([]byte(tc.after))
require.NoError(t, err)
}
@@ -749,15 +768,16 @@ func TestSSHActionFollowUp_RejectsVerdictAfterCheckRemoved(t *testing.T) {
// the policy changed while the user was authenticating.
action, err := ns.sshActionFollowUp(
t.Context(), zerolog.Nop(), &tailcfg.SSHAction{},
authID.String(), src.ID, dst.ID,
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, !removeRule, action.Accept, "accept, got %+v", action)
assert.Equal(t, removeRule, action.Reject, "reject, got %+v", action)
assert.Equal(t, !removeRule, recorded, "auth recorded for auto-approval")
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")
})
}
}
+4 -3
View File
@@ -21,9 +21,10 @@ type PolicyManager interface {
// BuildPeerMap constructs peer relationship maps for the given nodes
BuildPeerMap(nodes views.Slice[types.NodeView]) map[types.NodeID][]types.NodeID
SSHPolicy(baseURL string, node types.NodeView) (*tailcfg.SSHPolicy, error)
// SSHCheckParams resolves the SSH check period for a (src, dst) pair
// from the current policy, avoiding trust of client-provided URL params.
SSHCheckParams(srcNodeID, dstNodeID types.NodeID) (time.Duration, bool)
// SSHCheckParams resolves the SSH check period for src logging in to
// dst as localUser from the current policy, avoiding trust of
// client-provided URL params.
SSHCheckParams(srcNodeID, dstNodeID types.NodeID, localUser string) (time.Duration, bool)
SetPolicy(pol []byte) (bool, error)
// SetUsers replaces the user list. policyChanged reports whether clients
// need a policy refresh; peerMapChanged reports whether user-derived peer
+91 -1
View File
@@ -2753,11 +2753,35 @@ func TestIPSetToPrincipals(t *testing.T) {
}
}
// overlappingSSHChecks has a localpart check rule whose self-access reaches
// user1's own nodes ahead of an always-check root rule for the same pair.
var overlappingSSHChecks = []byte(`{
"tagOwners": {"tag:server": ["user1@"]},
"ssh": [
{
"action": "check",
"checkPeriod": "12h",
"src": ["user1@"],
"dst": ["tag:server"],
"users": ["localpart:*@example.com"]
},
{
"action": "check",
"checkPeriod": "always",
"src": ["user1@"],
"dst": ["autogroup:self"],
"users": ["root"]
}
]
}`)
func TestSSHCheckParams(t *testing.T) {
users := types.Users{
{Name: "user1", ID: 1},
{Name: "user2", ID: 2},
}
users[0].Email = "user1@example.com"
users[1].Email = "user2@example.com"
nodeUser1 := types.Node{
ID: 1,
@@ -2789,6 +2813,7 @@ func TestSSHCheckParams(t *testing.T) {
policy []byte
srcID types.NodeID
dstID types.NodeID
localUser string // defaults to a non-root user
wantPeriod time.Duration
wantOK bool
}{
@@ -2936,6 +2961,66 @@ func TestSSHCheckParams(t *testing.T) {
dstID: types.NodeID(2),
wantOK: false,
},
{
name: "root rejected by a nonroot rule",
policy: []byte(`{
"tagOwners": {"tag:server": ["user1@"]},
"ssh": [{
"action": "check",
"src": ["user2@"],
"dst": ["tag:server"],
"users": ["autogroup:nonroot"]
}]
}`),
srcID: types.NodeID(2),
dstID: types.NodeID(3),
localUser: "root",
wantOK: false,
},
{
name: "literal user must match",
policy: []byte(`{
"tagOwners": {"tag:server": ["user1@"]},
"ssh": [{
"action": "check",
"src": ["user2@"],
"dst": ["tag:server"],
"users": ["deploy"]
}]
}`),
srcID: types.NodeID(2),
dstID: types.NodeID(3),
localUser: "other",
wantOK: false,
},
{
// The client skips the localpart rule for root and enforces
// the later always-check root rule; so must the server.
name: "overlapping rules: root skips localpart self-access",
policy: overlappingSSHChecks,
srcID: types.NodeID(1),
dstID: types.NodeID(1),
localUser: "root",
wantPeriod: 0,
wantOK: true,
},
{
name: "overlapping rules: localpart user gets its period",
policy: overlappingSSHChecks,
srcID: types.NodeID(1),
dstID: types.NodeID(1),
localUser: "user1",
wantPeriod: 12 * time.Hour,
wantOK: true,
},
{
name: "overlapping rules: another user's localpart is rejected",
policy: overlappingSSHChecks,
srcID: types.NodeID(1),
dstID: types.NodeID(1),
localUser: "user2",
wantOK: false,
},
}
for _, tt := range tests {
@@ -2943,7 +3028,12 @@ func TestSSHCheckParams(t *testing.T) {
pm, err := NewPolicyManager(tt.policy, users, nodes.ViewSlice())
require.NoError(t, err)
period, ok := pm.SSHCheckParams(tt.srcID, tt.dstID)
localUser := tt.localUser
if localUser == "" {
localUser = "alice"
}
period, ok := pm.SSHCheckParams(tt.srcID, tt.dstID, localUser)
assert.Equal(t, tt.wantOK, ok, "ok mismatch")
if tt.wantOK {
+43 -15
View File
@@ -477,16 +477,18 @@ func (pm *PolicyManager) SSHPolicy(baseURL string, node types.NodeView) (*tailcf
return sshPol, nil
}
// SSHCheckParams resolves the SSH check period for a source-destination
// node pair by looking up the current policy. This avoids trusting URL
// SSHCheckParams resolves the SSH check period for src logging in to dst
// as localUser by looking up the current policy. This avoids trusting URL
// parameters that a client could tamper with. First-match wins across
// the policy's SSH rules.
// the policy's SSH rules; a rule only matches when it lets src log in as
// localUser, as the client does when it picks the rule.
//
// Returns (duration, true) when a matching rule is found and
// (0, false) when none is. A (0, true) return means the matched rule
// uses a zero check period (re-check every session).
func (pm *PolicyManager) SSHCheckParams(
srcNodeID, dstNodeID types.NodeID,
localUser string,
) (time.Duration, bool) {
pm.mu.RLock()
defer pm.mu.RUnlock()
@@ -533,20 +535,16 @@ func (pm *PolicyManager) SSHCheckParams(
continue
}
if !pm.sshRuleAllowsUser(rule, srcNode, localUser) {
continue
}
// Check if dst node matches any destination.
hasOtherDests := false
for _, dst := range rule.Destinations {
if ag, isAG := dst.(*AutoGroup); isAG && ag.Is(AutoGroupSelf) {
// User().Valid() guards the User().ID() dereference: the
// NodeStore can hold a non-tagged node with UserID set but
// the User association unhydrated (nil), and IsTagged()
// alone does not cover that. Mirrors filter.go's
// autogroup:self guard. Without it, a tailnet client on the
// Noise SSH-check path crashes the server (nil deref).
if !srcNode.IsTagged() && !dstNode.IsTagged() &&
srcNode.User().Valid() && dstNode.User().Valid() &&
srcNode.User().ID() == dstNode.User().ID() {
if sshNodesShareUser(srcNode, dstNode) {
return checkPeriodFromRule(rule), true
}
@@ -573,9 +571,7 @@ func (pm *PolicyManager) SSHCheckParams(
if srcNodeID == dstNodeID {
return checkPeriodFromRule(rule), true
}
} else if !srcNode.IsTagged() &&
srcNode.User().Valid() && dstNode.User().Valid() &&
srcNode.User().ID() == dstNode.User().ID() {
} else if sshNodesShareUser(srcNode, dstNode) {
return checkPeriodFromRule(rule), true
}
}
@@ -584,6 +580,38 @@ func (pm *PolicyManager) SSHCheckParams(
return 0, false
}
// sshNodesShareUser matches user-owned nodes with hydrated user associations.
// IsTagged alone does not guard User().ID(): the NodeStore can hold a
// non-tagged node whose UserID is set but whose User association is nil.
func sshNodesShareUser(srcNode, dstNode types.NodeView) bool {
return !srcNode.IsTagged() && !dstNode.IsTagged() &&
srcNode.User().Valid() && dstNode.User().Valid() &&
srcNode.User().ID() == dstNode.User().ID()
}
// sshRuleAllowsUser reports whether rule lets src log in as localUser,
// mirroring the SSHUsers maps compileSSHPolicy emits: root when listed,
// other users via autogroup:nonroot or a literal name, and anyone, root
// included, whose name is the source user's localpart.
func (pm *PolicyManager) sshRuleAllowsUser(rule SSH, srcNode types.NodeView, localUser string) bool {
isRoot := localUser == "root"
switch {
case localUser == "":
return false
case isRoot && rule.Users.ContainsRoot(),
!isRoot && rule.Users.ContainsNonRoot(),
!isRoot && slices.Contains(rule.Users.NormalUsers(), SSHUser(localUser)):
return true
case srcNode.IsTagged() || !srcNode.User().Valid() || !rule.Users.ContainsLocalpart():
return false
}
lp, ok := resolveLocalparts(rule.Users.LocalpartEntries(), pm.users)[srcNode.User().ID()]
return ok && lp == localUser
}
func (pm *PolicyManager) SetPolicy(polB []byte) (bool, error) {
if len(polB) == 0 {
return false, nil
+1 -1
View File
@@ -392,7 +392,7 @@ func TestSSHCheckParamsUnhydratedUserNoPanic(t *testing.T) {
require.NoError(t, err)
require.NotPanics(t, func() {
pm.SSHCheckParams(types.NodeID(1), types.NodeID(2))
pm.SSHCheckParams(types.NodeID(1), types.NodeID(2), "alice")
}, "SSHCheckParams must not panic when a non-tagged node has an unhydrated User")
}
@@ -8,7 +8,10 @@
package v2
import (
"encoding/json"
"fmt"
"path/filepath"
"slices"
"strings"
"testing"
@@ -259,9 +262,89 @@ func TestSSHDataCompat(t *testing.T) {
}
}
// sshLocalUser maps sshUser to a local user the way tailssh does for one
// rule: an exact entry wins, then "*"; "" means the rule does not apply.
func sshLocalUser(users map[string]string, sshUser string) string {
local, ok := users[sshUser]
if !ok {
local = users["*"]
}
if local == "=" {
return sshUser
}
return local
}
// assertSSHCheckParamsMatchRules checks SSHCheckParams against check rules as
// tailssh reads them: (src, dst, local user) must be found exactly when a
// holdAndDelegate rule for dst lists src as a principal and maps some SSH
// user to that local user.
func assertSSHCheckParamsMatchRules(
t *testing.T,
pm *PolicyManager,
nodes types.Nodes,
rulesFor func(dst *types.Node) []*tailcfg.SSHRule,
) {
t.Helper()
byIP := make(map[string]*types.Node)
for _, n := range nodes {
for _, ip := range n.IPs() {
byIP[ip.String()] = n
}
}
candidates := []string{"root", "nonroot-probe"}
for _, dst := range nodes {
for _, rule := range rulesFor(dst) {
for user := range rule.SSHUsers {
if user != "*" && !slices.Contains(candidates, user) {
candidates = append(candidates, user)
}
}
}
}
for _, dst := range nodes {
want := make(map[types.NodeID]map[string]bool)
for _, rule := range rulesFor(dst) {
if rule.Action == nil || rule.Action.HoldAndDelegate == "" {
continue
}
for _, p := range rule.Principals {
src, ok := byIP[p.NodeIP]
require.Truef(t, ok, "principal %q is not a node", p.NodeIP)
if want[src.ID] == nil {
want[src.ID] = make(map[string]bool)
}
for _, user := range candidates {
if local := sshLocalUser(rule.SSHUsers, user); local != "" {
want[src.ID][local] = true
}
}
}
}
for _, src := range nodes {
for _, user := range candidates {
_, got := pm.SSHCheckParams(src.ID, dst.ID, user)
assert.Equalf(t, want[src.ID][user], got,
"SSHCheckParams(%s -> %s as %s)", src.Hostname, dst.Hostname, user)
}
}
}
}
// TestSSHCheckParamsMatchesCaptures pins SSHCheckParams, which decides the
// SSH check callback, to the check rules Tailscale sent: exactly the
// captured holdAndDelegate principals must be found for each node.
// SSH check callback, to the check rules Tailscale sent.
func TestSSHCheckParamsMatchesCaptures(t *testing.T) {
t.Parallel()
@@ -291,36 +374,69 @@ func TestSSHCheckParamsMatchesCaptures(t *testing.T) {
)
require.NoError(t, err)
byIP := make(map[string]*types.Node)
for _, n := range nodes {
for _, ip := range n.IPs() {
byIP[ip.String()] = n
}
}
for _, dst := range nodes {
want := make(map[types.NodeID]bool)
for _, rule := range tf.Captures[dst.GivenName].SSHRules {
if rule.Action == nil || rule.Action.HoldAndDelegate == "" {
continue
}
for _, p := range rule.Principals {
src, ok := byIP[p.NodeIP]
require.Truef(t, ok, "principal %q is not a node", p.NodeIP)
want[src.ID] = true
}
}
for _, src := range nodes {
_, got := pm.SSHCheckParams(src.ID, dst.ID)
assert.Equalf(t, want[src.ID], got,
"SSHCheckParams(%s -> %s)", src.GivenName, dst.GivenName)
}
}
assertSSHCheckParamsMatchRules(t, pm, nodes,
func(dst *types.Node) []*tailcfg.SSHRule {
return tf.Captures[dst.GivenName].SSHRules
})
})
}
}
// TestSSHCheckParamsMatchesCompiledRules pins SSHCheckParams to headscale's
// own compiled rules for shapes the captures lack, such as a user whose
// email localpart is root.
func TestSSHCheckParamsMatchesCompiledRules(t *testing.T) {
t.Parallel()
users := types.Users{
{Name: "root", Email: "root@example.com"},
{Name: "alice", Email: "alice@example.com"},
}
users[0].ID = 1
users[1].ID = 2
nodes := types.Nodes{
node("root-1", "100.64.0.1", "fd7a:115c:a1e0::1", users[0]),
node("root-2", "100.64.0.2", "fd7a:115c:a1e0::2", users[0]),
node("alice-1", "100.64.0.3", "fd7a:115c:a1e0::3", users[1]),
node("server", "100.64.0.4", "fd7a:115c:a1e0::4", users[1]),
}
for i, n := range nodes {
n.ID = types.NodeID(i + 1) //nolint:gosec
}
nodes[3].Tags = []string{"tag:server"}
check := func(dst string, users ...string) string {
usersJSON, err := json.Marshal(users)
require.NoError(t, err)
return fmt.Sprintf(`{"action": "check", "src": ["autogroup:member"], "dst": [%q], "users": %s}`,
dst, usersJSON)
}
for name, rules := range map[string][]string{
"localpart on tag": {check("tag:server", "localpart:*@example.com")},
"localpart on self": {check("autogroup:self", "localpart:*@example.com")},
"localpart then root": {check("tag:server", "localpart:*@example.com"), check("autogroup:self", "root")},
"nonroot and literal on self": {check("autogroup:self", "autogroup:nonroot", "deploy")},
} {
t.Run(name, func(t *testing.T) {
t.Parallel()
pol := fmt.Sprintf(`{"tagOwners": {"tag:server": ["alice@"]}, "ssh": [%s]}`,
strings.Join(rules, ","))
pm, err := NewPolicyManager([]byte(pol), users, nodes.ViewSlice())
require.NoError(t, err)
assertSSHCheckParamsMatchRules(t, pm, nodes,
func(dst *types.Node) []*tailcfg.SSHRule {
sshPol, err := pm.SSHPolicy("", dst.View())
require.NoError(t, err)
return sshPol.Rules
})
})
}
}
+75 -4
View File
@@ -6,6 +6,7 @@ import (
"fmt"
"net/http"
"net/url"
"strconv"
"strings"
"testing"
"time"
@@ -15,6 +16,7 @@ import (
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"tailscale.com/tailcfg"
"tailscale.com/types/netmap"
)
// TestSSHCheckReDelegatesWhenSessionMissing exercises the fix for
@@ -45,7 +47,7 @@ func TestSSHCheckReDelegatesWhenSessionMissing(t *testing.T) {
// Sanity: the policy must actually subject this pair to a check, otherwise
// the test would pass for the wrong reason.
_, checkFound := h.Server.State().SSHCheckParams(srcID, dstID)
_, checkFound := h.Server.State().SSHCheckParams(srcID, dstID, sshCheckLocalUser)
require.True(t, checkFound, "test setup: (src, dst) must be subject to an SSH check")
// The dst node's first poll yields a real HoldAndDelegate carrying a real,
@@ -72,6 +74,9 @@ func TestSSHCheckReDelegatesWhenSessionMissing(t *testing.T) {
"re-delegation must mint a fresh auth_id")
}
// sshCheckLocalUser is the non-root local user the SSH-check polls log in as.
const sshCheckLocalUser = "alice"
// pollSSHAction issues an /machine/ssh/action poll from the given node over its
// real Noise connection, as tailscaled does. An empty authID is the initial
// poll; a non-empty one is a follow-up.
@@ -84,11 +89,20 @@ func pollSSHAction(
) tailcfg.SSHAction {
t.Helper()
actionURL := fmt.Sprintf("%s/machine/ssh/action/%d/to/%d", serverURL, srcID, dstID)
actionURL := fmt.Sprintf("%s/machine/ssh/action/%d/to/%d?local_user=%s",
serverURL, srcID, dstID, sshCheckLocalUser)
if authID != "" {
actionURL += "?auth_id=" + authID
actionURL += "&auth_id=" + authID
}
return pollSSHActionURL(t, node, actionURL)
}
// pollSSHActionURL polls actionURL from the given node over its real Noise
// connection.
func pollSSHActionURL(t *testing.T, node *servertest.TestClient, actionURL string) tailcfg.SSHAction {
t.Helper()
// Noise requests are addressed with the https scheme; the control client
// routes them over the established Noise connection (mirroring how
// controlclient issues its own register/map calls).
@@ -163,7 +177,7 @@ func TestSSHCheckRejectedAfterRuleRemoved(t *testing.T) {
h.ChangePolicy(t, []byte(after))
_, checkFound := h.Server.State().SSHCheckParams(srcID, dstID)
_, checkFound := h.Server.State().SSHCheckParams(srcID, dstID, sshCheckLocalUser)
require.False(t, checkFound, "test setup: check must be gone")
for poll, id := range map[string]string{"initial": "", "follow-up": authID.String()} {
@@ -174,3 +188,60 @@ func TestSSHCheckRejectedAfterRuleRemoved(t *testing.T) {
})
}
}
// expandHoldURL substitutes a HoldAndDelegate URL's variables literally, as
// tailssh does before polling it.
func expandHoldURL(holdURL string, srcID, dstID types.NodeID, localUser string) string {
return strings.NewReplacer(
"$SRC_NODE_ID", strconv.FormatUint(srcID.Uint64(), 10),
"$DST_NODE_ID", strconv.FormatUint(dstID.Uint64(), 10),
"$SSH_USER", url.QueryEscape(localUser),
"$LOCAL_USER", url.QueryEscape(localUser),
).Replace(holdURL)
}
// TestSSHCheckFollowsReturnedHoldURL walks a check the way tailssh does:
// expand the HoldAndDelegate URL from the netmap, poll it, then poll the
// URL the server returns. The follow-up must still carry the login user,
// or a root-only rule denies a user who already authenticated.
// https://github.com/juanfont/headscale/issues/3508
func TestSSHCheckFollowsReturnedHoldURL(t *testing.T) {
t.Parallel()
h := servertest.NewHarness(t, 2)
srcID := types.NodeID(h.Client(0).Netmap().SelfNode.ID()) //nolint:gosec
dstID := types.NodeID(h.Client(1).Netmap().SelfNode.ID()) //nolint:gosec
h.ChangePolicy(t, []byte(`{"ssh": [{
"action": "check",
"src": ["harness-default@"],
"dst": ["autogroup:self"],
"users": ["root"]
}]}`))
var ruleURL string
h.Client(1).WaitForCondition(t, "check rule in netmap", 10*time.Second,
func(nm *netmap.NetworkMap) bool {
if nm.SSHPolicy == nil || len(nm.SSHPolicy.Rules) == 0 ||
nm.SSHPolicy.Rules[0].Action == nil {
return false
}
ruleURL = nm.SSHPolicy.Rules[0].Action.HoldAndDelegate
return ruleURL != ""
})
initial := pollSSHActionURL(t, h.Client(1), expandHoldURL(ruleURL, srcID, dstID, "root"))
require.NotEmpty(t, initial.HoldAndDelegate, "check must hold, got %+v", initial)
auth, ok := h.Server.State().GetAuthCacheEntry(authIDFromHoldURL(t, initial.HoldAndDelegate))
require.True(t, ok)
auth.FinishAuth(types.AuthVerdict{})
followUp := pollSSHActionURL(t, h.Client(1),
expandHoldURL(initial.HoldAndDelegate, srcID, dstID, "root"))
assert.True(t, followUp.Accept, "authenticated root login must be accepted, got %+v", followUp)
}
+4 -3
View File
@@ -1260,12 +1260,13 @@ func (s *State) SSHPolicy(node types.NodeView) (*tailcfg.SSHPolicy, error) {
return s.polMan.SSHPolicy(s.cfg.ServerURL, node)
}
// SSHCheckParams resolves the SSH check period for a source-destination
// node pair from the current policy.
// SSHCheckParams resolves the SSH check period for src logging in to dst
// as localUser from the current policy.
func (s *State) SSHCheckParams(
srcNodeID, dstNodeID types.NodeID,
localUser string,
) (time.Duration, bool) {
return s.polMan.SSHCheckParams(srcNodeID, dstNodeID)
return s.polMan.SSHCheckParams(srcNodeID, dstNodeID, localUser)
}
// Filter returns the current network filter rules and matches.