Fix review findings for embedded VNC server

This commit is contained in:
Viktor Liu
2026-06-10 10:57:50 +02:00
parent 699ac9c203
commit 4adaa73253
39 changed files with 694 additions and 289 deletions
@@ -2,6 +2,7 @@ package peers
import (
"bytes"
"context"
"encoding/json"
"net/http"
"net/http/httptest"
@@ -34,7 +35,7 @@ func TestCreateTemporaryAccess_RejectsCallerWithoutPeersCreate(t *testing.T) {
// nil so the test fails loudly if the handler tries to call it.
permMgr.EXPECT().
ValidateUserPermissions(gomock.Any(), gomock.Any(), gomock.Any(), gomock.Eq(modules.Peers), gomock.Eq(operations.Create)).
Return(false, nil).
Return(false, context.Background(), nil).
Times(1)
h := &Handler{
@@ -74,11 +75,11 @@ func TestCreateTemporaryAccess_RejectsCallerWithoutPoliciesCreate(t *testing.T)
permMgr.EXPECT().
ValidateUserPermissions(gomock.Any(), gomock.Any(), gomock.Any(), gomock.Eq(modules.Peers), gomock.Eq(operations.Create)).
Return(true, nil).
Return(true, context.Background(), nil).
Times(1)
permMgr.EXPECT().
ValidateUserPermissions(gomock.Any(), gomock.Any(), gomock.Any(), gomock.Eq(modules.Policies), gomock.Eq(operations.Create)).
Return(false, nil).
Return(false, context.Background(), nil).
Times(1)
h := &Handler{
-1
View File
@@ -138,7 +138,6 @@ type Flags struct {
DisableIPv6 bool
LazyConnectionEnabled bool
}
// PeerSystemMeta is a metadata of a Peer machine system
+9 -3
View File
@@ -2536,7 +2536,7 @@ func (s *SqlStore) getPolicyRules(ctx context.Context, policyIDs []string) ([]*t
if len(policyIDs) == 0 {
return nil, nil
}
const query = `SELECT id, policy_id, name, description, enabled, action, destinations, destination_resource, sources, source_resource, bidirectional, protocol, ports, port_ranges, authorized_groups, authorized_user FROM policy_rules WHERE policy_id = ANY($1)`
const query = `SELECT id, policy_id, name, description, enabled, action, destinations, destination_resource, sources, source_resource, bidirectional, protocol, ports, port_ranges, authorized_groups, authorized_user, session_pub_key, session_display_name FROM policy_rules WHERE policy_id = ANY($1)`
rows, err := s.pool.Query(ctx, query, policyIDs)
if err != nil {
return nil, err
@@ -2545,8 +2545,8 @@ func (s *SqlStore) getPolicyRules(ctx context.Context, policyIDs []string) ([]*t
var r types.PolicyRule
var dest, destRes, sources, sourceRes, ports, portRanges, authorizedGroups []byte
var enabled, bidirectional sql.NullBool
var authorizedUser sql.NullString
err := row.Scan(&r.ID, &r.PolicyID, &r.Name, &r.Description, &enabled, &r.Action, &dest, &destRes, &sources, &sourceRes, &bidirectional, &r.Protocol, &ports, &portRanges, &authorizedGroups, &authorizedUser)
var authorizedUser, sessionPubKey, sessionDisplayName sql.NullString
err := row.Scan(&r.ID, &r.PolicyID, &r.Name, &r.Description, &enabled, &r.Action, &dest, &destRes, &sources, &sourceRes, &bidirectional, &r.Protocol, &ports, &portRanges, &authorizedGroups, &authorizedUser, &sessionPubKey, &sessionDisplayName)
if err == nil {
if enabled.Valid {
r.Enabled = enabled.Bool
@@ -2578,6 +2578,12 @@ func (s *SqlStore) getPolicyRules(ctx context.Context, policyIDs []string) ([]*t
if authorizedUser.Valid {
r.AuthorizedUser = authorizedUser.String
}
if sessionPubKey.Valid {
r.SessionPubKey = sessionPubKey.String
}
if sessionDisplayName.Valid {
r.SessionDisplayName = sessionDisplayName.String
}
}
return &r, err
})
+1 -2
View File
@@ -14,7 +14,6 @@ import (
"github.com/rs/xid"
log "github.com/sirupsen/logrus"
auth "github.com/netbirdio/netbird/shared/sessionauth"
nbdns "github.com/netbirdio/netbird/dns"
proxydomain "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/domain"
"github.com/netbirdio/netbird/management/internals/modules/reverseproxy/service"
@@ -29,6 +28,7 @@ import (
"github.com/netbirdio/netbird/route"
"github.com/netbirdio/netbird/shared/management/domain"
"github.com/netbirdio/netbird/shared/management/status"
auth "github.com/netbirdio/netbird/shared/sessionauth"
"github.com/netbirdio/netbird/version"
)
@@ -170,7 +170,6 @@ func (a *Account) GetGroup(groupID string) *Group {
return a.Groups[groupID]
}
func (a *Account) addNetworksRoutingPeers(
networkResourcesRoutes []*route.Route,
peer *nbpeer.Peer,
@@ -9,13 +9,13 @@ import (
"strings"
"time"
auth "github.com/netbirdio/netbird/shared/sessionauth"
nbdns "github.com/netbirdio/netbird/dns"
resourceTypes "github.com/netbirdio/netbird/management/server/networks/resources/types"
routerTypes "github.com/netbirdio/netbird/management/server/networks/routers/types"
nbpeer "github.com/netbirdio/netbird/management/server/peer"
"github.com/netbirdio/netbird/route"
"github.com/netbirdio/netbird/shared/management/domain"
auth "github.com/netbirdio/netbird/shared/sessionauth"
)
type NetworkMapComponents struct {
@@ -109,7 +109,7 @@ func (c *NetworkMapComponents) Calculate(ctx context.Context) *NetworkMap {
peerGroups := c.GetPeerGroups(targetPeerID)
connRes := c.getPeerConnectionResources(targetPeerID)
connRes := c.getPeerConnectionResources(ctx, targetPeerID)
aclPeers := connRes.peers
peersToConnect, expiredPeers := c.filterPeersByLoginExpiration(aclPeers)
@@ -182,7 +182,7 @@ type peerConnectionResult struct {
sshEnabled bool
}
func (c *NetworkMapComponents) getPeerConnectionResources(targetPeerID string) peerConnectionResult {
func (c *NetworkMapComponents) getPeerConnectionResources(ctx context.Context, targetPeerID string) peerConnectionResult {
targetPeer := c.GetPeerInfo(targetPeerID)
if targetPeer == nil {
return peerConnectionResult{}
@@ -202,7 +202,7 @@ func (c *NetworkMapComponents) getPeerConnectionResources(targetPeerID string) p
if !rule.Enabled {
continue
}
c.applyPolicyRule(rule, policy.SourcePostureChecks, targetPeer, targetPeerID, generateResources, state)
c.applyPolicyRule(ctx, rule, policy.SourcePostureChecks, targetPeer, targetPeerID, generateResources, state)
}
}
@@ -218,6 +218,7 @@ func (c *NetworkMapComponents) getPeerConnectionResources(targetPeerID string) p
}
func (c *NetworkMapComponents) applyPolicyRule(
ctx context.Context,
rule *PolicyRule,
sourcePostureChecks []string,
targetPeer *nbpeer.Peer,
@@ -229,8 +230,12 @@ func (c *NetworkMapComponents) applyPolicyRule(
destinationPeers, peerInDestinations := c.resolveRuleEndpoint(rule.DestinationResource, rule.Destinations, targetPeerID, nil)
cb := ruleAuthCallbacks{
collectSSHUsers: c.collectAuthorizedUsers,
collectVNCUsers: c.collectAuthorizedUsers,
collectSSHUsers: func(r *PolicyRule, t map[string]map[string]struct{}) {
c.collectAuthorizedUsers(ctx, r, t)
},
collectVNCUsers: func(r *PolicyRule, t map[string]map[string]struct{}) {
c.collectAuthorizedUsers(ctx, r, t)
},
getAllowedUserIDs: c.getAllowedUserIDs,
}
applyResolvedRuleToState(rule, sourcePeers, destinationPeers, peerInSources, peerInDestinations, targetPeer.SSHEnabled, generateResources, cb, state)
@@ -249,10 +254,10 @@ func (c *NetworkMapComponents) resolveRuleEndpoint(
}
// collectAuthorizedUsers populates the target map with authorized user mappings from the rule.
func (c *NetworkMapComponents) collectAuthorizedUsers(rule *PolicyRule, target map[string]map[string]struct{}) {
func (c *NetworkMapComponents) collectAuthorizedUsers(ctx context.Context, rule *PolicyRule, target map[string]map[string]struct{}) {
switch {
case len(rule.AuthorizedGroups) > 0:
mergeAuthorizedGroupUsers(context.Background(), rule.AuthorizedGroups, c.GroupIDToUserIDs, target)
mergeAuthorizedGroupUsers(ctx, rule.AuthorizedGroups, c.GroupIDToUserIDs, target)
case rule.AuthorizedUser != "":
ensureWildcardUser(target, rule.AuthorizedUser)
default:
@@ -6,8 +6,8 @@ import (
log "github.com/sirupsen/logrus"
auth "github.com/netbirdio/netbird/shared/sessionauth"
nbpeer "github.com/netbirdio/netbird/management/server/peer"
auth "github.com/netbirdio/netbird/shared/sessionauth"
)
// peerConnResolveState carries the in-progress maps mutated by per-rule
@@ -1,6 +1,10 @@
package types
import "testing"
import (
"testing"
nbpeer "github.com/netbirdio/netbird/management/server/peer"
)
// TestHandleVNCRule_BidirectionalDistributesPubkeyToSourcePeer covers the
// latent bug where a bidirectional VNC rule used to drop the
@@ -83,3 +87,68 @@ func TestHandleVNCRule_DestinationAlwaysGetsPubkey(t *testing.T) {
t.Fatalf("expected 1 session pubkey for destination peer, got %d", len(state.vncSessionPubKeys))
}
}
// TestApplyResolvedRule_BidirectionalSSHEnablesSourcePeer locks the
// bidirectional widening for netbird-ssh rules: a peer that appears only
// in the rule's sources of a bidirectional SSH rule must get SSH enabled
// and its authorized users collected, because the rule grants access in
// both directions. A unidirectional rule must not do this for a
// source-only peer.
func TestApplyResolvedRule_BidirectionalSSHEnablesSourcePeer(t *testing.T) {
collected := false
cb := ruleAuthCallbacks{
collectSSHUsers: func(_ *PolicyRule, target map[string]map[string]struct{}) {
collected = true
target["local"] = map[string]struct{}{"user1": {}}
},
}
rule := &PolicyRule{
Protocol: PolicyRuleProtocolNetbirdSSH,
Bidirectional: true,
}
state := &peerConnResolveState{
authorizedUsers: make(map[string]map[string]struct{}),
vncAuthorizedUsers: make(map[string]map[string]struct{}),
}
applyResolvedRuleToState(rule, nil, nil, true /*peerInSources*/, false /*peerInDestinations*/, false, func(*PolicyRule, []*nbpeer.Peer, int) {}, cb, state)
if !state.sshEnabled {
t.Fatal("expected SSH enabled on source-side peer of bidirectional SSH rule")
}
if !collected {
t.Fatal("expected authorized users collected on source-side peer of bidirectional SSH rule")
}
if _, ok := state.authorizedUsers["local"]; !ok {
t.Fatal("expected authorized users map populated for source-side peer")
}
}
// TestApplyResolvedRule_UnidirectionalSSHSkipsSourcePeer is the negative
// counterpart: a unidirectional SSH rule must not enable SSH for a peer
// that appears only in sources.
func TestApplyResolvedRule_UnidirectionalSSHSkipsSourcePeer(t *testing.T) {
collected := false
cb := ruleAuthCallbacks{
collectSSHUsers: func(_ *PolicyRule, _ map[string]map[string]struct{}) {
collected = true
},
}
rule := &PolicyRule{
Protocol: PolicyRuleProtocolNetbirdSSH,
Bidirectional: false,
}
state := &peerConnResolveState{
authorizedUsers: make(map[string]map[string]struct{}),
vncAuthorizedUsers: make(map[string]map[string]struct{}),
}
applyResolvedRuleToState(rule, nil, nil, true /*peerInSources*/, false /*peerInDestinations*/, false, func(*PolicyRule, []*nbpeer.Peer, int) {}, cb, state)
if state.sshEnabled {
t.Fatal("expected SSH NOT enabled on source-only peer of unidirectional SSH rule")
}
if collected {
t.Fatal("expected NO authorized users collected on source-only peer of unidirectional SSH rule")
}
}