mirror of
https://github.com/netbirdio/netbird.git
synced 2026-10-01 02:59:08 +02:00
Fix review findings for embedded VNC server
This commit is contained in:
@@ -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{
|
||||
|
||||
@@ -138,7 +138,6 @@ type Flags struct {
|
||||
DisableIPv6 bool
|
||||
|
||||
LazyConnectionEnabled bool
|
||||
|
||||
}
|
||||
|
||||
// PeerSystemMeta is a metadata of a Peer machine system
|
||||
|
||||
@@ -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
|
||||
})
|
||||
|
||||
@@ -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")
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user