diff --git a/hscontrol/noise.go b/hscontrol/noise.go index cced5e1c2..4ca7e0f7e 100644 --- a/hscontrol/noise.go +++ b/hscontrol/noise.go @@ -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 } diff --git a/hscontrol/noise_test.go b/hscontrol/noise_test.go index 3ec37ea89..dcaf40969 100644 --- a/hscontrol/noise_test.go +++ b/hscontrol/noise_test.go @@ -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") }) } } diff --git a/hscontrol/policy/pm.go b/hscontrol/policy/pm.go index 78bf62438..ba8a61051 100644 --- a/hscontrol/policy/pm.go +++ b/hscontrol/policy/pm.go @@ -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 diff --git a/hscontrol/policy/v2/filter_test.go b/hscontrol/policy/v2/filter_test.go index 80cbe0ce5..fcdb7e66e 100644 --- a/hscontrol/policy/v2/filter_test.go +++ b/hscontrol/policy/v2/filter_test.go @@ -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 { diff --git a/hscontrol/policy/v2/policy.go b/hscontrol/policy/v2/policy.go index 306e00151..ff4e2258c 100644 --- a/hscontrol/policy/v2/policy.go +++ b/hscontrol/policy/v2/policy.go @@ -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 diff --git a/hscontrol/policy/v2/policy_test.go b/hscontrol/policy/v2/policy_test.go index 57720c7bb..c5a427ebf 100644 --- a/hscontrol/policy/v2/policy_test.go +++ b/hscontrol/policy/v2/policy_test.go @@ -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") } diff --git a/hscontrol/policy/v2/tailscale_ssh_data_compat_test.go b/hscontrol/policy/v2/tailscale_ssh_data_compat_test.go index 59342560f..79fc9b1a2 100644 --- a/hscontrol/policy/v2/tailscale_ssh_data_compat_test.go +++ b/hscontrol/policy/v2/tailscale_ssh_data_compat_test.go @@ -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 + }) }) } } diff --git a/hscontrol/servertest/sshcheck_test.go b/hscontrol/servertest/sshcheck_test.go index 73df0332a..020ded5e5 100644 --- a/hscontrol/servertest/sshcheck_test.go +++ b/hscontrol/servertest/sshcheck_test.go @@ -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) +} diff --git a/hscontrol/state/state.go b/hscontrol/state/state.go index edc49174f..4b5747bd2 100644 --- a/hscontrol/state/state.go +++ b/hscontrol/state/state.go @@ -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.