diff --git a/hscontrol/policy/v2/policy.go b/hscontrol/policy/v2/policy.go index 2847ec2a7..306e00151 100644 --- a/hscontrol/policy/v2/policy.go +++ b/hscontrol/policy/v2/policy.go @@ -534,6 +534,8 @@ func (pm *PolicyManager) SSHCheckParams( } // 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 @@ -551,6 +553,8 @@ func (pm *PolicyManager) SSHCheckParams( continue } + hasOtherDests = true + dstIPs, err := dst.Resolve(pm.pol, pm.users, pm.nodes) if err != nil || dstIPs == nil { continue @@ -560,6 +564,21 @@ func (pm *PolicyManager) SSHCheckParams( return checkPeriodFromRule(rule), true } } + + // Localpart self-access: a source outside dst still gets the rule + // for its own user's nodes, or itself if tagged (compileSSHPolicy). + if hasOtherDests && rule.Users.ContainsLocalpart() && + slices.ContainsFunc(dstNode.IPs(), srcIPs.Contains) { + if dstNode.IsTagged() { + if srcNodeID == dstNodeID { + return checkPeriodFromRule(rule), true + } + } else if !srcNode.IsTagged() && + srcNode.User().Valid() && dstNode.User().Valid() && + srcNode.User().ID() == dstNode.User().ID() { + return checkPeriodFromRule(rule), true + } + } } return 0, false diff --git a/hscontrol/policy/v2/tailscale_ssh_data_compat_test.go b/hscontrol/policy/v2/tailscale_ssh_data_compat_test.go index 60a12d8df..59342560f 100644 --- a/hscontrol/policy/v2/tailscale_ssh_data_compat_test.go +++ b/hscontrol/policy/v2/tailscale_ssh_data_compat_test.go @@ -258,3 +258,69 @@ func TestSSHDataCompat(t *testing.T) { }) } } + +// 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. +func TestSSHCheckParamsMatchesCaptures(t *testing.T) { + t.Parallel() + + files, err := filepath.Glob(filepath.Join("testdata", "ssh*_results", "*.hujson")) + require.NoError(t, err) + require.NotEmpty(t, files) + + users := setupSSHDataCompatUsers() + + for _, file := range files { + tf := loadSSHTestFile(t, file) + if tf.Input.APIResponseCode != 200 { + continue + } + + if _, skip := sshSkipReasons[tf.TestID]; skip { + continue + } + + t.Run(tf.TestID, func(t *testing.T) { + t.Parallel() + + nodes := buildGrantsNodesFromCapture(users, tf) + + pm, err := NewPolicyManager( + []byte(tf.Input.FullPolicy), users, nodes.ViewSlice(), + ) + 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) + } + } + }) + } +}