policy/v2: find SSH check params for localpart self-access

compileSSHPolicy sends these check rules; SSHCheckParams missed them,
so they never auto-approved within checkPeriod.

Updates #3508
This commit is contained in:
Kristoffer Dalby
2026-10-02 09:03:27 +00:00
committed by Kristoffer Dalby
parent d7b333d2aa
commit 03e723c611
2 changed files with 85 additions and 0 deletions
+19
View File
@@ -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
@@ -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)
}
}
})
}
}