mirror of
https://github.com/juanfont/headscale.git
synced 2026-10-10 08:40:07 +09:00
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:
committed by
Kristoffer Dalby
parent
dceb584c89
commit
957a332d5d
+15
-8
@@ -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
@@ -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")
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
})
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
@@ -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.
|
||||
|
||||
Reference in New Issue
Block a user