mirror of
https://github.com/netbirdio/netbird.git
synced 2026-10-01 19:19:07 +02:00
Merge remote-tracking branch 'origin/main' into fix_debug_upload_url_from_mgmt
# Conflicts: # management/server/activity/codes.go # management/server/store/sql_store.go # management/server/store/sql_store_test.go # upload-server/server/server.go
This commit is contained in:
@@ -0,0 +1,57 @@
|
||||
package lifecycle
|
||||
|
||||
import (
|
||||
"runtime/debug"
|
||||
"sync"
|
||||
|
||||
log "github.com/sirupsen/logrus"
|
||||
)
|
||||
|
||||
// StopHandlers collects functions to run once when their owner exits. Embed it
|
||||
// in a server type to expose OnStop and RunStopHandlers.
|
||||
type StopHandlers struct {
|
||||
mu sync.Mutex
|
||||
stopped bool
|
||||
handlers []func()
|
||||
}
|
||||
|
||||
// OnStop registers fn to run once when the owner stops. Handlers run in
|
||||
// reverse registration order. A handler registered after the owner has
|
||||
// stopped runs immediately.
|
||||
func (h *StopHandlers) OnStop(fn func()) {
|
||||
h.mu.Lock()
|
||||
stopped := h.stopped
|
||||
if !stopped {
|
||||
h.handlers = append(h.handlers, fn)
|
||||
}
|
||||
h.mu.Unlock()
|
||||
|
||||
if stopped {
|
||||
runStopHandler(fn)
|
||||
}
|
||||
}
|
||||
|
||||
// RunStopHandlers runs every registered handler once, last registered first.
|
||||
// Later calls are no-ops, so it can be wired to several exit paths at once.
|
||||
func (h *StopHandlers) RunStopHandlers() {
|
||||
h.mu.Lock()
|
||||
handlers := h.handlers
|
||||
h.handlers = nil
|
||||
h.stopped = true
|
||||
h.mu.Unlock()
|
||||
|
||||
for i := len(handlers) - 1; i >= 0; i-- {
|
||||
runStopHandler(handlers[i])
|
||||
}
|
||||
}
|
||||
|
||||
// runStopHandler keeps one panicking handler from skipping the ones still
|
||||
// pending; on the shutdown path there is no second chance to run them.
|
||||
func runStopHandler(fn func()) {
|
||||
defer func() {
|
||||
if r := recover(); r != nil {
|
||||
log.Errorf("stop handler panicked: %v\n%s", r, debug.Stack())
|
||||
}
|
||||
}()
|
||||
fn()
|
||||
}
|
||||
@@ -0,0 +1,43 @@
|
||||
package lifecycle
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
)
|
||||
|
||||
func TestStopHandlers_RunOnceInReverseOrder(t *testing.T) {
|
||||
var h StopHandlers
|
||||
var order []string
|
||||
h.OnStop(func() { order = append(order, "first") })
|
||||
h.OnStop(func() { order = append(order, "second") })
|
||||
|
||||
h.RunStopHandlers()
|
||||
h.RunStopHandlers()
|
||||
|
||||
assert.Equal(t, []string{"second", "first"}, order, "handlers must run once, last registered first")
|
||||
}
|
||||
|
||||
func TestStopHandlers_PanicDoesNotSkipRemainingHandlers(t *testing.T) {
|
||||
var h StopHandlers
|
||||
var order []string
|
||||
h.OnStop(func() { order = append(order, "first") })
|
||||
h.OnStop(func() { panic("boom") })
|
||||
h.OnStop(func() { order = append(order, "third") })
|
||||
|
||||
h.RunStopHandlers()
|
||||
|
||||
assert.Equal(t, []string{"third", "first"}, order, "handlers around a panicking one must still run")
|
||||
}
|
||||
|
||||
func TestStopHandlers_LateRegistrationRunsImmediately(t *testing.T) {
|
||||
var h StopHandlers
|
||||
h.RunStopHandlers()
|
||||
|
||||
runs := 0
|
||||
h.OnStop(func() { runs++ })
|
||||
assert.Equal(t, 1, runs, "a handler registered after the stop must run right away")
|
||||
|
||||
h.RunStopHandlers()
|
||||
assert.Equal(t, 1, runs, "later runs must stay no-ops and must not repeat the handler")
|
||||
}
|
||||
@@ -35,7 +35,12 @@ type EnvelopeResult struct {
|
||||
//
|
||||
// dnsName is the account's DNS domain ("netbird.cloud" etc.); used when
|
||||
// rebuilding the per-peer FQDNs that proto.RemotePeerConfig carries.
|
||||
func EnvelopeToNetworkMap(ctx context.Context, env *proto.NetworkMapEnvelope, localPeerKey, dnsName string) (*EnvelopeResult, error) {
|
||||
//
|
||||
// skipRouteFirewallRules leaves RoutesFirewallRules empty. Callers that have
|
||||
// no firewall to program pass true: the rules are the most expensive part of
|
||||
// Calculate on a peer that routes many network resources, and nothing reads
|
||||
// them afterwards.
|
||||
func EnvelopeToNetworkMap(ctx context.Context, env *proto.NetworkMapEnvelope, localPeerKey, dnsName string, skipRouteFirewallRules bool) (*EnvelopeResult, error) {
|
||||
components, err := DecodeEnvelope(ctx, env)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("decode envelope: %w", err)
|
||||
@@ -53,6 +58,7 @@ func EnvelopeToNetworkMap(ctx context.Context, env *proto.NetworkMapEnvelope, lo
|
||||
return nil, fmt.Errorf("receiving peer (wg_key prefix %q) not found among %d decoded peers — components have no PeerID, Calculate would return empty", trimKey(localPeerKey), len(components.Peers))
|
||||
}
|
||||
components.PeerID = canonicalKey
|
||||
components.SkipRouteFirewallRules = skipRouteFirewallRules
|
||||
|
||||
includeIPv6 := localPeer.SupportsIPv6() && localPeer.IPv6.IsValid()
|
||||
useSourcePrefixes := localPeer.SupportsSourcePrefixes()
|
||||
|
||||
@@ -9,6 +9,7 @@ import (
|
||||
"net/netip"
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
goproto "google.golang.org/protobuf/proto"
|
||||
|
||||
@@ -37,7 +38,7 @@ func TestEnvelopeToNetworkMap_RoundTrip(t *testing.T) {
|
||||
var decoded proto.NetworkMapEnvelope
|
||||
require.NoError(t, goproto.Unmarshal(wire, &decoded), "unmarshal envelope")
|
||||
|
||||
result, err := nbnetworkmap.EnvelopeToNetworkMap(context.Background(), &decoded, localPeerKey, "netbird.cloud")
|
||||
result, err := nbnetworkmap.EnvelopeToNetworkMap(context.Background(), &decoded, localPeerKey, "netbird.cloud", false)
|
||||
require.NoError(t, err, "EnvelopeToNetworkMap")
|
||||
require.NotNil(t, result)
|
||||
require.NotNil(t, result.NetworkMap, "decoded NetworkMap must be non-nil")
|
||||
@@ -78,7 +79,7 @@ func TestCalculate_FirewallRuleProtocol_NeverNetbirdSSH(t *testing.T) {
|
||||
var decoded proto.NetworkMapEnvelope
|
||||
require.NoError(t, goproto.Unmarshal(wire, &decoded))
|
||||
|
||||
result, err := nbnetworkmap.EnvelopeToNetworkMap(context.Background(), &decoded, localPeerKey, "netbird.cloud")
|
||||
result, err := nbnetworkmap.EnvelopeToNetworkMap(context.Background(), &decoded, localPeerKey, "netbird.cloud", false)
|
||||
require.NoError(t, err)
|
||||
require.NotEmpty(t, result.NetworkMap.FirewallRules, "ssh policy should produce firewall rules")
|
||||
for i, fr := range result.NetworkMap.FirewallRules {
|
||||
@@ -88,13 +89,13 @@ func TestCalculate_FirewallRuleProtocol_NeverNetbirdSSH(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestEnvelopeToNetworkMap_NilEnvelope(t *testing.T) {
|
||||
_, err := nbnetworkmap.EnvelopeToNetworkMap(context.Background(), nil, "key", "netbird.cloud")
|
||||
_, err := nbnetworkmap.EnvelopeToNetworkMap(context.Background(), nil, "key", "netbird.cloud", false)
|
||||
require.Error(t, err, "nil envelope must produce an error rather than panic")
|
||||
}
|
||||
|
||||
func TestEnvelopeToNetworkMap_FullPayloadMissing(t *testing.T) {
|
||||
env := &proto.NetworkMapEnvelope{}
|
||||
_, err := nbnetworkmap.EnvelopeToNetworkMap(context.Background(), env, "key", "netbird.cloud")
|
||||
_, err := nbnetworkmap.EnvelopeToNetworkMap(context.Background(), env, "key", "netbird.cloud", false)
|
||||
require.Error(t, err, "envelope with no Full payload must produce an error")
|
||||
}
|
||||
|
||||
@@ -126,7 +127,7 @@ func TestDecodeEnvelope_MalformedWgKeyPeerSkipped(t *testing.T) {
|
||||
var decoded proto.NetworkMapEnvelope
|
||||
require.NoError(t, goproto.Unmarshal(wire, &decoded), "unmarshal envelope")
|
||||
|
||||
result, err := nbnetworkmap.EnvelopeToNetworkMap(context.Background(), &decoded, localPeerKey, "netbird.cloud")
|
||||
result, err := nbnetworkmap.EnvelopeToNetworkMap(context.Background(), &decoded, localPeerKey, "netbird.cloud", false)
|
||||
require.NoError(t, err, "EnvelopeToNetworkMap must tolerate one bad peer key")
|
||||
require.NotNil(t, result)
|
||||
require.NotNil(t, result.Components)
|
||||
@@ -195,7 +196,7 @@ func TestEnvelopeRoundTrip_AllGroupShortCircuitParity(t *testing.T) {
|
||||
var decodedEnv proto.NetworkMapEnvelope
|
||||
require.NoError(t, goproto.Unmarshal(wire, &decodedEnv), "unmarshal envelope")
|
||||
|
||||
result, err := nbnetworkmap.EnvelopeToNetworkMap(ctx, &decodedEnv, peers["peer-T"].Key, "netbird.cloud")
|
||||
result, err := nbnetworkmap.EnvelopeToNetworkMap(ctx, &decodedEnv, peers["peer-T"].Key, "netbird.cloud", false)
|
||||
require.NoError(t, err, "EnvelopeToNetworkMap")
|
||||
clientNM := result.NetworkMap
|
||||
|
||||
@@ -253,7 +254,7 @@ func TestEnvelopeToNetworkMap_EmptyComponents(t *testing.T) {
|
||||
var decoded proto.NetworkMapEnvelope
|
||||
require.NoError(t, goproto.Unmarshal(wire, &decoded), "unmarshal envelope")
|
||||
|
||||
result, err := nbnetworkmap.EnvelopeToNetworkMap(context.Background(), &decoded, localPeerKey, "netbird.cloud")
|
||||
result, err := nbnetworkmap.EnvelopeToNetworkMap(context.Background(), &decoded, localPeerKey, "netbird.cloud", false)
|
||||
require.NoError(t, err, "EnvelopeToNetworkMap must degrade gracefully on empty components")
|
||||
require.Equal(t, uint64(7), result.NetworkMap.Serial)
|
||||
require.Empty(t, result.NetworkMap.RemotePeers, "unvalidated peer connects to nobody")
|
||||
@@ -276,7 +277,7 @@ func TestEnvelopeToNetworkMap_MissingNetwork(t *testing.T) {
|
||||
var decoded proto.NetworkMapEnvelope
|
||||
require.NoError(t, goproto.Unmarshal(wire, &decoded), "unmarshal envelope")
|
||||
|
||||
result, err := nbnetworkmap.EnvelopeToNetworkMap(context.Background(), &decoded, localPeerKey, "netbird.cloud")
|
||||
result, err := nbnetworkmap.EnvelopeToNetworkMap(context.Background(), &decoded, localPeerKey, "netbird.cloud", false)
|
||||
require.NoError(t, err, "a missing AccountNetwork must not panic the client")
|
||||
require.NotNil(t, result.Components.Network)
|
||||
require.NotEmpty(t, result.NetworkMap.RemotePeers, "the rest of the snapshot stays usable")
|
||||
@@ -353,3 +354,110 @@ func randomWgKey(t *testing.T) string {
|
||||
require.NoError(t, err)
|
||||
return base64.StdEncoding.EncodeToString(raw[:])
|
||||
}
|
||||
|
||||
// TestEnvelopeToNetworkMap_SkipRouteFirewallRules covers the flag end to end,
|
||||
// through the envelope rather than by poking Calculate directly. The
|
||||
// RoutesFirewallRulesIsEmpty derivation is the part that matters: the client's
|
||||
// legacy-management probe reads an empty rule list together with that bit, so
|
||||
// skipping the rules must set it rather than leave it false.
|
||||
func TestEnvelopeToNetworkMap_SkipRouteFirewallRules(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
c, routerKey := buildRoutedResourceComponents(t)
|
||||
|
||||
envelope := mgmtgrpc.EncodeNetworkMapEnvelope(mgmtgrpc.ComponentsEnvelopeInput{
|
||||
Components: c,
|
||||
DNSDomain: "netbird.cloud",
|
||||
})
|
||||
wire, err := goproto.Marshal(envelope)
|
||||
require.NoError(t, err, "marshal envelope")
|
||||
var decoded proto.NetworkMapEnvelope
|
||||
require.NoError(t, goproto.Unmarshal(wire, &decoded), "unmarshal envelope")
|
||||
|
||||
full, err := nbnetworkmap.EnvelopeToNetworkMap(ctx, &decoded, routerKey, "netbird.cloud", false)
|
||||
require.NoError(t, err, "EnvelopeToNetworkMap without skip")
|
||||
require.NotEmpty(t, full.NetworkMap.RoutesFirewallRules,
|
||||
"baseline: the router peer must receive route firewall rules")
|
||||
require.False(t, full.NetworkMap.RoutesFirewallRulesIsEmpty,
|
||||
"baseline: the empty bit must be false when rules are present")
|
||||
|
||||
var decodedSkip proto.NetworkMapEnvelope
|
||||
require.NoError(t, goproto.Unmarshal(wire, &decodedSkip), "unmarshal envelope")
|
||||
skipped, err := nbnetworkmap.EnvelopeToNetworkMap(ctx, &decodedSkip, routerKey, "netbird.cloud", true)
|
||||
require.NoError(t, err, "EnvelopeToNetworkMap with skip")
|
||||
|
||||
assert.Empty(t, skipped.NetworkMap.RoutesFirewallRules,
|
||||
"route firewall rules must not be computed when skipped")
|
||||
assert.True(t, skipped.NetworkMap.RoutesFirewallRulesIsEmpty,
|
||||
"the empty bit must be derived from the skipped list, or the client misreads it as legacy management")
|
||||
assert.Len(t, skipped.NetworkMap.Routes, len(full.NetworkMap.Routes),
|
||||
"skipping route firewall rules must not change the routes")
|
||||
assert.Len(t, skipped.NetworkMap.RemotePeers, len(full.NetworkMap.RemotePeers),
|
||||
"skipping route firewall rules must not change the remote peers")
|
||||
}
|
||||
|
||||
// buildRoutedResourceComponents returns components in which the local peer is
|
||||
// the routing peer for one enabled network resource, reachable by a second
|
||||
// peer through a resource policy — the minimum shape that yields a non-empty
|
||||
// RoutesFirewallRules. It also returns the local peer's WG key.
|
||||
func buildRoutedResourceComponents(t *testing.T) (*types.NetworkMapComponents, string) {
|
||||
t.Helper()
|
||||
|
||||
routerKey := randomWgKey(t)
|
||||
peers := map[string]*nmdata.Peer{
|
||||
"peer-R": {
|
||||
ID: "peer-R", Key: routerKey, DNSLabel: "router",
|
||||
IP: netip.AddrFrom4([4]byte{100, 64, 0, 1}),
|
||||
Meta: nmdata.PeerSystemMeta{WtVersion: "0.40.0"},
|
||||
},
|
||||
"peer-S": {
|
||||
ID: "peer-S", Key: randomWgKey(t), DNSLabel: "source",
|
||||
IP: netip.AddrFrom4([4]byte{100, 64, 0, 2}),
|
||||
Meta: nmdata.PeerSystemMeta{WtVersion: "0.40.0"},
|
||||
},
|
||||
}
|
||||
|
||||
resourcePolicy := &nmdata.Policy{
|
||||
ID: "pol-res", PublicID: "10", Enabled: true,
|
||||
Rules: []*nmdata.PolicyRule{{
|
||||
ID: "rule-res",
|
||||
Enabled: true,
|
||||
Action: string(types.PolicyTrafficActionAccept),
|
||||
Protocol: string(types.PolicyRuleProtocolALL),
|
||||
Sources: []string{"g-src"},
|
||||
}},
|
||||
}
|
||||
|
||||
c := &types.NetworkMapComponents{
|
||||
PeerID: "peer-R",
|
||||
Network: &nmdata.Network{
|
||||
Identifier: "net-routed-resource",
|
||||
Net: net.IPNet{IP: net.IP{100, 64, 0, 0}, Mask: net.CIDRMask(10, 32)},
|
||||
Serial: 1,
|
||||
},
|
||||
AccountSettings: &nmdata.AccountSettingsInfo{},
|
||||
DNSSettings: &nmdata.DNSSettings{},
|
||||
Peers: peers,
|
||||
Groups: map[string]*nmdata.Group{
|
||||
"g-src": {PublicID: "1", Name: "sources", Peers: []string{"peer-S"}},
|
||||
"g-routers": {PublicID: "2", Name: "routers", Peers: []string{"peer-R"}},
|
||||
},
|
||||
NetworkResources: []*nmdata.NetworkResource{{
|
||||
ID: "res-1", NetworkID: "netid-1", PublicID: "100", Name: "res1",
|
||||
Type: "subnet",
|
||||
Prefix: netip.MustParsePrefix("10.200.0.0/24"),
|
||||
Enabled: true,
|
||||
}},
|
||||
RoutersMap: map[string]map[string]*nmdata.NetworkRouter{
|
||||
"netid-1": {"peer-R": {
|
||||
PublicID: "200", PeerGroups: []string{"g-routers"}, Metric: 9999, Enabled: true,
|
||||
}},
|
||||
},
|
||||
ResourcePoliciesMap: map[string][]*nmdata.Policy{
|
||||
"res-1": {resourcePolicy},
|
||||
},
|
||||
Policies: []*nmdata.Policy{resourcePolicy},
|
||||
NetworkXIDToPublicID: map[string]string{"netid-1": "1"},
|
||||
}
|
||||
|
||||
return c, routerKey
|
||||
}
|
||||
|
||||
@@ -135,6 +135,11 @@ func NewUserPendingApprovalError() error {
|
||||
return Errorf(PermissionDenied, "user is pending approval")
|
||||
}
|
||||
|
||||
// NewUserPendingApprovalByOwnerError creates a new Error with PermissionDenied type for a blocked user pending approval, naming the masked address of the owner who can approve them
|
||||
func NewUserPendingApprovalByOwnerError(ownerEmail string) error {
|
||||
return Errorf(PermissionDenied, "user is pending approval by owner %s", ownerEmail)
|
||||
}
|
||||
|
||||
// NewPeerNotRegisteredError creates a new Error with Unauthenticated type unregistered peer
|
||||
func NewPeerNotRegisteredError() error {
|
||||
return Errorf(Unauthenticated, "peer is not registered")
|
||||
|
||||
@@ -58,6 +58,13 @@ type NetworkMapComponents struct {
|
||||
// domain targets.
|
||||
ForceRoutingPeerDNSResolution bool
|
||||
|
||||
// SkipRouteFirewallRules drops the route firewall rule computation from
|
||||
// Calculate. A receiver without a firewall manager never reads
|
||||
// RoutesFirewallRules, and on a routing peer with many network resources
|
||||
// building them dominates the cost of a sync. Defaults to false so the
|
||||
// management server keeps producing them.
|
||||
SkipRouteFirewallRules bool
|
||||
|
||||
routesByPeerOnce sync.Once
|
||||
routesByPeerIdx map[string][]routeIndexEntry
|
||||
|
||||
@@ -149,11 +156,15 @@ func (c *NetworkMapComponents) Calculate(ctx context.Context) *NetworkMap {
|
||||
includeIPv6 = p.SupportsIPv6() && p.IPv6.IsValid()
|
||||
}
|
||||
routesUpdate := filterAndExpandRoutes(c.getRoutesToSync(targetPeerID, peersToConnect, peerGroups), includeIPv6)
|
||||
routesFirewallRules := c.getPeerRoutesFirewallRules(ctx, targetPeerID, includeIPv6)
|
||||
|
||||
var routesFirewallRules []*RouteFirewallRule
|
||||
if !c.SkipRouteFirewallRules {
|
||||
routesFirewallRules = c.getPeerRoutesFirewallRules(ctx, targetPeerID, includeIPv6)
|
||||
}
|
||||
|
||||
isRouter, networkResourcesRoutes, sourcePeers := c.getNetworkResourcesRoutesToSync(targetPeerID)
|
||||
var networkResourcesFirewallRules []*RouteFirewallRule
|
||||
if isRouter {
|
||||
if isRouter && !c.SkipRouteFirewallRules {
|
||||
networkResourcesFirewallRules = c.getPeerNetworkResourceFirewallRules(ctx, targetPeerID, networkResourcesRoutes, includeIPv6)
|
||||
}
|
||||
|
||||
|
||||
@@ -0,0 +1,127 @@
|
||||
package profiling
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"fmt"
|
||||
"net/netip"
|
||||
"net/url"
|
||||
"os"
|
||||
"strings"
|
||||
"sync/atomic"
|
||||
|
||||
"github.com/caarlos0/env/v11"
|
||||
"github.com/grafana/pyroscope-go"
|
||||
log "github.com/sirupsen/logrus"
|
||||
)
|
||||
|
||||
var errNotConfigured = errors.New("pyroscope not configured")
|
||||
|
||||
var started atomic.Bool
|
||||
|
||||
type config struct {
|
||||
Address string `env:"NB_PYROSCOPE_ADDRESS"`
|
||||
User string `env:"NB_PYROSCOPE_USER,notEmpty"`
|
||||
Password string `env:"NB_PYROSCOPE_PASSWORD,notEmpty"`
|
||||
}
|
||||
|
||||
func Start(applicationName string) func() {
|
||||
noop := func() {}
|
||||
|
||||
cfg, err := loadConfig()
|
||||
switch {
|
||||
case errors.Is(err, errNotConfigured):
|
||||
log.Info("pyroscope not configured, continuous profiling disabled")
|
||||
return noop
|
||||
case err != nil:
|
||||
log.Errorf("failed to load pyroscope config: %v", err)
|
||||
return noop
|
||||
}
|
||||
|
||||
// pprof allows one CPU profile per process, so a second profiler (e.g. the
|
||||
// signal server inside the combined binary) would only log errors.
|
||||
if !started.CompareAndSwap(false, true) {
|
||||
log.Warnf("continuous profiling already running in this process, not starting it for %s", applicationName)
|
||||
return noop
|
||||
}
|
||||
|
||||
tags := map[string]string{}
|
||||
if hostname, err := os.Hostname(); err == nil {
|
||||
tags["instance"] = hostname
|
||||
} else {
|
||||
log.Warnf("failed to resolve hostname for profile tags: %v", err)
|
||||
}
|
||||
|
||||
profiler, err := pyroscope.Start(pyroscope.Config{
|
||||
ApplicationName: applicationName,
|
||||
ServerAddress: cfg.Address,
|
||||
BasicAuthUser: cfg.User,
|
||||
BasicAuthPassword: cfg.Password,
|
||||
Logger: log.StandardLogger(),
|
||||
Tags: tags,
|
||||
ProfileTypes: []pyroscope.ProfileType{
|
||||
pyroscope.ProfileCPU,
|
||||
pyroscope.ProfileAllocObjects,
|
||||
pyroscope.ProfileAllocSpace,
|
||||
pyroscope.ProfileInuseObjects,
|
||||
pyroscope.ProfileInuseSpace,
|
||||
},
|
||||
})
|
||||
if err != nil {
|
||||
started.Store(false)
|
||||
log.Errorf("failed to start continuous profiling: %v", err)
|
||||
return noop
|
||||
}
|
||||
|
||||
return func() {
|
||||
_ = profiler.Stop()
|
||||
started.Store(false)
|
||||
}
|
||||
}
|
||||
|
||||
func loadConfig() (config, error) {
|
||||
var cfg config
|
||||
if err := env.Parse(&cfg); err != nil {
|
||||
if cfg.Address == "" {
|
||||
return cfg, errNotConfigured
|
||||
}
|
||||
return cfg, fmt.Errorf("failed to parse pyroscope config: %w", err)
|
||||
}
|
||||
|
||||
if cfg.Address == "" {
|
||||
return cfg, errNotConfigured
|
||||
}
|
||||
if err := validateAddress(cfg.Address); err != nil {
|
||||
return cfg, err
|
||||
}
|
||||
|
||||
return cfg, nil
|
||||
}
|
||||
|
||||
// validateAddress refuses to send the basic-auth credentials in plaintext to
|
||||
// anything but a loopback or private endpoint.
|
||||
func validateAddress(address string) error {
|
||||
u, err := url.Parse(address)
|
||||
if err != nil {
|
||||
return fmt.Errorf("invalid pyroscope address %q: %w", address, err)
|
||||
}
|
||||
|
||||
switch u.Scheme {
|
||||
case "https":
|
||||
return nil
|
||||
case "http":
|
||||
if isLocalOrPrivate(u.Hostname()) {
|
||||
return nil
|
||||
}
|
||||
return fmt.Errorf("insecure pyroscope address %q: use https for non-local endpoints", address)
|
||||
default:
|
||||
return fmt.Errorf("pyroscope address %q must use http or https", address)
|
||||
}
|
||||
}
|
||||
|
||||
func isLocalOrPrivate(host string) bool {
|
||||
if host == "localhost" || strings.HasSuffix(host, ".localhost") {
|
||||
return true
|
||||
}
|
||||
ip, err := netip.ParseAddr(host)
|
||||
return err == nil && (ip.IsLoopback() || ip.IsPrivate())
|
||||
}
|
||||
@@ -0,0 +1,202 @@
|
||||
package profiling
|
||||
|
||||
import (
|
||||
"os"
|
||||
"testing"
|
||||
|
||||
log "github.com/sirupsen/logrus"
|
||||
logtest "github.com/sirupsen/logrus/hooks/test"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
func TestStartSkipsSecondProfilerInProcess(t *testing.T) {
|
||||
clearEnv(t)
|
||||
t.Setenv("NB_PYROSCOPE_ADDRESS", "http://127.0.0.1:1")
|
||||
t.Setenv("NB_PYROSCOPE_USER", "user")
|
||||
t.Setenv("NB_PYROSCOPE_PASSWORD", "token")
|
||||
|
||||
started.Store(true)
|
||||
t.Cleanup(func() { started.Store(false) })
|
||||
hook := logtest.NewGlobal()
|
||||
t.Cleanup(hook.Reset)
|
||||
|
||||
stop := Start("netbird-second")
|
||||
stop()
|
||||
|
||||
assert.True(t, started.Load(), "the running profiler must stay marked as started")
|
||||
entry := hook.LastEntry()
|
||||
require.NotNil(t, entry, "the skipped start must be logged")
|
||||
assert.Equal(t, log.WarnLevel, entry.Level)
|
||||
assert.Contains(t, entry.Message, "already running")
|
||||
}
|
||||
|
||||
func TestLoadConfig(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
env map[string]string
|
||||
expected config
|
||||
errIs error
|
||||
wantErr bool
|
||||
}{
|
||||
{
|
||||
name: "address unset disables profiling",
|
||||
errIs: errNotConfigured,
|
||||
},
|
||||
{
|
||||
name: "empty address disables profiling",
|
||||
env: map[string]string{"NB_PYROSCOPE_ADDRESS": ""},
|
||||
errIs: errNotConfigured,
|
||||
},
|
||||
{
|
||||
name: "credentials without address disable profiling",
|
||||
env: map[string]string{
|
||||
"NB_PYROSCOPE_USER": "123456",
|
||||
"NB_PYROSCOPE_PASSWORD": "token",
|
||||
},
|
||||
errIs: errNotConfigured,
|
||||
},
|
||||
{
|
||||
name: "address without credentials fails",
|
||||
env: map[string]string{
|
||||
"NB_PYROSCOPE_ADDRESS": "https://profiles-prod-001.grafana.net",
|
||||
},
|
||||
wantErr: true,
|
||||
},
|
||||
{
|
||||
name: "address with empty credentials fails",
|
||||
env: map[string]string{
|
||||
"NB_PYROSCOPE_ADDRESS": "https://profiles-prod-001.grafana.net",
|
||||
"NB_PYROSCOPE_USER": "",
|
||||
"NB_PYROSCOPE_PASSWORD": "",
|
||||
},
|
||||
wantErr: true,
|
||||
},
|
||||
{
|
||||
name: "address without password fails",
|
||||
env: map[string]string{
|
||||
"NB_PYROSCOPE_ADDRESS": "https://profiles-prod-001.grafana.net",
|
||||
"NB_PYROSCOPE_USER": "123456",
|
||||
},
|
||||
wantErr: true,
|
||||
},
|
||||
{
|
||||
name: "full configuration",
|
||||
env: map[string]string{
|
||||
"NB_PYROSCOPE_ADDRESS": "https://profiles-prod-001.grafana.net",
|
||||
"NB_PYROSCOPE_USER": "123456",
|
||||
"NB_PYROSCOPE_PASSWORD": "token",
|
||||
},
|
||||
expected: config{
|
||||
Address: "https://profiles-prod-001.grafana.net",
|
||||
User: "123456",
|
||||
Password: "token",
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "http to loopback is allowed",
|
||||
env: map[string]string{
|
||||
"NB_PYROSCOPE_ADDRESS": "http://127.0.0.1:4040",
|
||||
"NB_PYROSCOPE_USER": "123456",
|
||||
"NB_PYROSCOPE_PASSWORD": "token",
|
||||
},
|
||||
expected: config{
|
||||
Address: "http://127.0.0.1:4040",
|
||||
User: "123456",
|
||||
Password: "token",
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "http to localhost is allowed",
|
||||
env: map[string]string{
|
||||
"NB_PYROSCOPE_ADDRESS": "http://localhost:4040",
|
||||
"NB_PYROSCOPE_USER": "123456",
|
||||
"NB_PYROSCOPE_PASSWORD": "token",
|
||||
},
|
||||
expected: config{
|
||||
Address: "http://localhost:4040",
|
||||
User: "123456",
|
||||
Password: "token",
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "http to private network is allowed",
|
||||
env: map[string]string{
|
||||
"NB_PYROSCOPE_ADDRESS": "http://10.0.0.5:4040",
|
||||
"NB_PYROSCOPE_USER": "123456",
|
||||
"NB_PYROSCOPE_PASSWORD": "token",
|
||||
},
|
||||
expected: config{
|
||||
Address: "http://10.0.0.5:4040",
|
||||
User: "123456",
|
||||
Password: "token",
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "http to public host is rejected",
|
||||
env: map[string]string{
|
||||
"NB_PYROSCOPE_ADDRESS": "http://pyroscope.example.com",
|
||||
"NB_PYROSCOPE_USER": "123456",
|
||||
"NB_PYROSCOPE_PASSWORD": "token",
|
||||
},
|
||||
wantErr: true,
|
||||
},
|
||||
{
|
||||
name: "http to public address is rejected",
|
||||
env: map[string]string{
|
||||
"NB_PYROSCOPE_ADDRESS": "http://203.0.113.10:4040",
|
||||
"NB_PYROSCOPE_USER": "123456",
|
||||
"NB_PYROSCOPE_PASSWORD": "token",
|
||||
},
|
||||
wantErr: true,
|
||||
},
|
||||
{
|
||||
name: "address without scheme is rejected",
|
||||
env: map[string]string{
|
||||
"NB_PYROSCOPE_ADDRESS": "pyroscope.example.com:4040",
|
||||
"NB_PYROSCOPE_USER": "123456",
|
||||
"NB_PYROSCOPE_PASSWORD": "token",
|
||||
},
|
||||
wantErr: true,
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
clearEnv(t)
|
||||
for k, v := range tt.env {
|
||||
t.Setenv(k, v)
|
||||
}
|
||||
|
||||
cfg, err := loadConfig()
|
||||
|
||||
switch {
|
||||
case tt.errIs != nil:
|
||||
require.ErrorIs(t, err, tt.errIs)
|
||||
case tt.wantErr:
|
||||
require.Error(t, err)
|
||||
require.NotErrorIs(t, err, errNotConfigured)
|
||||
default:
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, tt.expected, cfg)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestStartWithoutConfigurationIsNoop(t *testing.T) {
|
||||
clearEnv(t)
|
||||
|
||||
stop := Start("netbird-test")
|
||||
require.NotNil(t, stop)
|
||||
stop()
|
||||
}
|
||||
|
||||
func clearEnv(t *testing.T) {
|
||||
t.Helper()
|
||||
|
||||
for _, k := range []string{"NB_PYROSCOPE_ADDRESS", "NB_PYROSCOPE_USER", "NB_PYROSCOPE_PASSWORD"} {
|
||||
t.Setenv(k, "")
|
||||
require.NoError(t, os.Unsetenv(k))
|
||||
}
|
||||
}
|
||||
@@ -30,6 +30,12 @@ const (
|
||||
|
||||
var (
|
||||
ErrConnAlreadyExists = fmt.Errorf("connection already exists")
|
||||
// ErrServerDisconnected is the cancellation cause of a relayed Conn when the
|
||||
// client lost the connection to the relay server.
|
||||
ErrServerDisconnected = fmt.Errorf("relay server disconnected")
|
||||
// ErrPeerDisconnected is the cancellation cause of a relayed Conn when the
|
||||
// remote peer went offline.
|
||||
ErrPeerDisconnected = fmt.Errorf("remote peer disconnected")
|
||||
)
|
||||
|
||||
type internalStopFlag struct {
|
||||
@@ -74,16 +80,17 @@ type connContainer struct {
|
||||
msgChanLock sync.Mutex
|
||||
closed bool // flag to check if channel is closed
|
||||
ctx context.Context
|
||||
cancel context.CancelFunc
|
||||
cancel context.CancelCauseFunc
|
||||
}
|
||||
|
||||
func newConnContainer(log *log.Entry, c *Client, peerID messages.PeerID, instanceURL *RelayAddr) *connContainer {
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
ctx, cancel := context.WithCancelCause(context.Background())
|
||||
msgChan := make(chan Msg, connChannelSize)
|
||||
cn := &Conn{
|
||||
dstID: peerID,
|
||||
messageChan: msgChan,
|
||||
instanceURL: instanceURL,
|
||||
ctx: ctx,
|
||||
}
|
||||
cc := &connContainer{
|
||||
log: log,
|
||||
@@ -106,10 +113,6 @@ func newConnContainer(log *log.Entry, c *Client, peerID messages.PeerID, instanc
|
||||
return cc
|
||||
}
|
||||
|
||||
func (cc *connContainer) netConn() net.Conn {
|
||||
return cc.conn
|
||||
}
|
||||
|
||||
func (cc *connContainer) writeMsg(msg Msg) {
|
||||
cc.msgChanLock.Lock()
|
||||
defer cc.msgChanLock.Unlock()
|
||||
@@ -128,8 +131,8 @@ func (cc *connContainer) writeMsg(msg Msg) {
|
||||
}
|
||||
}
|
||||
|
||||
func (cc *connContainer) close() {
|
||||
cc.cancel()
|
||||
func (cc *connContainer) close(cause error) {
|
||||
cc.cancel(cause)
|
||||
|
||||
cc.msgChanLock.Lock()
|
||||
defer cc.msgChanLock.Unlock()
|
||||
@@ -279,7 +282,7 @@ func (c *Client) Connect(ctx context.Context) error {
|
||||
c.stateSubscription = NewPeersStateSubscription(c.log, c.relayConn, c.closeConnsByPeerID)
|
||||
|
||||
c.log = c.log.WithField("relay", instanceURL.String())
|
||||
c.log.Infof("relay connection established")
|
||||
c.log.Infof("relay connection established, server IP: %s", connectedIP(c.relayConn))
|
||||
|
||||
c.serviceIsRunning = true
|
||||
|
||||
@@ -293,12 +296,12 @@ func (c *Client) Connect(ctx context.Context) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
// OpenConn create a new net.Conn for the destination peer ID. In case if the connection is in progress
|
||||
// OpenConn create a new Conn for the destination peer ID. In case if the connection is in progress
|
||||
// to the relay server, the function will block until the connection is established or timed out. Otherwise,
|
||||
// it will return immediately.
|
||||
// It block until the server confirm the peer is online.
|
||||
// todo: what should happen if call with the same peerID with multiple times?
|
||||
func (c *Client) OpenConn(ctx context.Context, dstPeerID string) (net.Conn, error) {
|
||||
func (c *Client) OpenConn(ctx context.Context, dstPeerID string) (*Conn, error) {
|
||||
peerID := messages.HashID(dstPeerID)
|
||||
|
||||
c.mu.Lock()
|
||||
@@ -335,7 +338,7 @@ func (c *Client) OpenConn(ctx context.Context, dstPeerID string) (net.Conn, erro
|
||||
delete(c.conns, peerID)
|
||||
}
|
||||
c.mu.Unlock()
|
||||
container.close()
|
||||
container.close(err)
|
||||
return nil, err
|
||||
}
|
||||
|
||||
@@ -345,13 +348,13 @@ func (c *Client) OpenConn(ctx context.Context, dstPeerID string) (net.Conn, erro
|
||||
delete(c.conns, peerID)
|
||||
}
|
||||
c.mu.Unlock()
|
||||
container.close()
|
||||
container.close(ErrServerDisconnected)
|
||||
return nil, fmt.Errorf("relay connection is not established")
|
||||
}
|
||||
c.mu.Unlock()
|
||||
|
||||
c.log.Infof("remote peer is available: %s", peerID)
|
||||
return container.netConn(), nil
|
||||
return container.conn, nil
|
||||
}
|
||||
|
||||
// ServerInstanceURL returns the address of the relay server. It could change after the close and reopen the connection.
|
||||
@@ -364,23 +367,6 @@ func (c *Client) ServerInstanceURL() (string, error) {
|
||||
return c.instanceURL.String(), nil
|
||||
}
|
||||
|
||||
// ConnectedIP returns the IP address of the live relay-server connection,
|
||||
// extracted from the underlying socket's RemoteAddr. Zero value if not
|
||||
// connected or if the address is not an IP literal.
|
||||
func (c *Client) ConnectedIP() netip.Addr {
|
||||
c.mu.Lock()
|
||||
conn := c.relayConn
|
||||
c.mu.Unlock()
|
||||
if conn == nil {
|
||||
return netip.Addr{}
|
||||
}
|
||||
addr := conn.RemoteAddr()
|
||||
if addr == nil {
|
||||
return netip.Addr{}
|
||||
}
|
||||
return extractIPLiteral(addr.String())
|
||||
}
|
||||
|
||||
// SetOnDisconnectListener sets a function that will be called when the connection to the relay server is closed.
|
||||
func (c *Client) SetOnDisconnectListener(fn func(string)) {
|
||||
c.listenerMutex.Lock()
|
||||
@@ -777,9 +763,20 @@ func (c *Client) listenForStopEvents(ctx context.Context, hc *healthcheck.Receiv
|
||||
}
|
||||
}
|
||||
|
||||
func (c *Client) serverInstanceAddress() (string, netip.Addr, error) {
|
||||
c.mu.Lock()
|
||||
defer c.mu.Unlock()
|
||||
|
||||
addr, err := c.ServerInstanceURL()
|
||||
if err != nil {
|
||||
return "", netip.Addr{}, err
|
||||
}
|
||||
return addr, connectedIP(c.relayConn), nil
|
||||
}
|
||||
|
||||
func (c *Client) closeAllConns() {
|
||||
for _, container := range c.conns {
|
||||
container.close()
|
||||
container.close(ErrServerDisconnected)
|
||||
}
|
||||
c.conns = make(map[messages.PeerID]*connContainer)
|
||||
|
||||
@@ -799,7 +796,7 @@ func (c *Client) closeConnsByPeerID(peerIDs []messages.PeerID) {
|
||||
}
|
||||
|
||||
container.log.Infof("remote peer has been disconnected, free up connection: %s", peerID)
|
||||
container.close()
|
||||
container.close(ErrPeerDisconnected)
|
||||
delete(c.conns, peerID)
|
||||
}
|
||||
|
||||
@@ -827,7 +824,7 @@ func (c *Client) closeConn(containerRef *connContainer, id messages.PeerID) erro
|
||||
|
||||
c.log.Infof("free up connection to peer: %s", id)
|
||||
delete(c.conns, id)
|
||||
current.close()
|
||||
current.close(net.ErrClosed)
|
||||
|
||||
return nil
|
||||
}
|
||||
@@ -923,6 +920,17 @@ func (c *Client) handlePeersWentOfflineMsg(buf []byte) {
|
||||
c.stateSubscription.OnPeersWentOffline(peersID)
|
||||
}
|
||||
|
||||
func connectedIP(conn net.Conn) netip.Addr {
|
||||
if conn == nil {
|
||||
return netip.Addr{}
|
||||
}
|
||||
addr := conn.RemoteAddr()
|
||||
if addr == nil {
|
||||
return netip.Addr{}
|
||||
}
|
||||
return extractIPLiteral(addr.String())
|
||||
}
|
||||
|
||||
// extractIPLiteral returns the IP from address forms produced by the relay
|
||||
// dialers (URL or host:port). Zero value if the host is not an IP.
|
||||
func extractIPLiteral(s string) netip.Addr {
|
||||
|
||||
@@ -8,6 +8,8 @@ import (
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
"go.opentelemetry.io/otel"
|
||||
|
||||
"github.com/netbirdio/netbird/client/iface"
|
||||
@@ -68,18 +70,17 @@ func TestClient_ServerIPRecoversFromUnresolvableFQDN(t *testing.T) {
|
||||
if !c.Ready() {
|
||||
t.Fatalf("client not ready after connect")
|
||||
}
|
||||
if got := c.ConnectedIP(); got.String() != "127.0.0.1" {
|
||||
t.Fatalf("ConnectedIP = %q, want 127.0.0.1", got)
|
||||
}
|
||||
url, ip, err := c.serverInstanceAddress()
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, srvCfg.ExposedAddress, url, "relay URL must come from the handshake")
|
||||
assert.Equal(t, netip.MustParseAddr("127.0.0.1"), ip, "relay IP must come from the connection")
|
||||
})
|
||||
}
|
||||
|
||||
// TestClient_ConnectedIPAfterFQDNDial verifies ConnectedIP returns the
|
||||
// resolved IP after a successful FQDN-based dial. The underlying socket's
|
||||
// RemoteAddr must be exposed through the dialer wrappers; if it returns
|
||||
// the dial-time URL instead, ConnectedIP returns empty and the dial
|
||||
// IP we advertise to peers is empty too.
|
||||
func TestClient_ConnectedIPAfterFQDNDial(t *testing.T) {
|
||||
// TestClient_ServerInstanceAddressAfterFQDNDial verifies the relay address
|
||||
// includes the resolved IP after an FQDN dial. The dialer wrappers must expose
|
||||
// the socket's RemoteAddr; returning the dial-time URL would lose the IP.
|
||||
func TestClient_ServerInstanceAddressAfterFQDNDial(t *testing.T) {
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 15*time.Second)
|
||||
defer cancel()
|
||||
|
||||
@@ -111,10 +112,10 @@ func TestClient_ConnectedIPAfterFQDNDial(t *testing.T) {
|
||||
}
|
||||
t.Cleanup(func() { _ = c.Close() })
|
||||
|
||||
got := c.ConnectedIP().String()
|
||||
if got != "127.0.0.1" && got != "::1" {
|
||||
t.Fatalf("ConnectedIP after FQDN dial = %q, want 127.0.0.1 or ::1", got)
|
||||
}
|
||||
url, ip, err := c.serverInstanceAddress()
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, srvCfg.ExposedAddress, url, "relay URL must come from the handshake")
|
||||
assert.Contains(t, []string{"127.0.0.1", "::1"}, ip.String(), "relay IP must resolve to localhost")
|
||||
}
|
||||
|
||||
func TestSubstituteHost(t *testing.T) {
|
||||
@@ -214,15 +215,12 @@ func TestSubstituteHost(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestClient_ConnectedIPEmptyWhenNotConnected(t *testing.T) {
|
||||
c := NewClient("rel://example.invalid:80", hmacTokenStore, "x", iface.DefaultMTU)
|
||||
if got := c.ConnectedIP(); got.IsValid() {
|
||||
t.Fatalf("ConnectedIP on disconnected client = %q, want zero", got)
|
||||
}
|
||||
func TestConnectedIPNilConnection(t *testing.T) {
|
||||
assert.False(t, connectedIP(nil).IsValid(), "missing connection must not provide an IP")
|
||||
}
|
||||
|
||||
// staticAddr is a net.Addr that returns a fixed string. Used to verify
|
||||
// ConnectedIP parses RemoteAddr correctly.
|
||||
// connectedIP parses RemoteAddr correctly.
|
||||
type staticAddr struct{ s string }
|
||||
|
||||
func (a staticAddr) Network() string { return "tcp" }
|
||||
@@ -235,7 +233,7 @@ type stubConn struct {
|
||||
|
||||
func (s stubConn) RemoteAddr() net.Addr { return s.remote }
|
||||
|
||||
func TestClient_ConnectedIPParsesRemoteAddr(t *testing.T) {
|
||||
func TestConnectedIPParsesRemoteAddr(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
s string
|
||||
@@ -252,15 +250,12 @@ func TestClient_ConnectedIPParsesRemoteAddr(t *testing.T) {
|
||||
}
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
c := &Client{relayConn: stubConn{remote: staticAddr{s: tt.s}}}
|
||||
got := c.ConnectedIP()
|
||||
got := connectedIP(stubConn{remote: staticAddr{s: tt.s}})
|
||||
var gotStr string
|
||||
if got.IsValid() {
|
||||
gotStr = got.String()
|
||||
}
|
||||
if gotStr != tt.want {
|
||||
t.Errorf("ConnectedIP(%q) = %q, want %q", tt.s, gotStr, tt.want)
|
||||
}
|
||||
assert.Equal(t, tt.want, gotStr, "IP extracted from RemoteAddr %q", tt.s)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
package client
|
||||
|
||||
import (
|
||||
"context"
|
||||
"net"
|
||||
"time"
|
||||
|
||||
@@ -12,11 +13,20 @@ type Conn struct {
|
||||
dstID messages.PeerID
|
||||
messageChan chan Msg
|
||||
instanceURL *RelayAddr
|
||||
ctx context.Context
|
||||
writeFn func(messages.PeerID, []byte) (int, error)
|
||||
closeFn func(messages.PeerID) error
|
||||
localAddrFn func() net.Addr
|
||||
}
|
||||
|
||||
// Context returns a context that is cancelled when the connection is torn down,
|
||||
// either by Close or by the relay client losing the server connection. The
|
||||
// cancellation cause carries the reason, see ErrServerDisconnected and
|
||||
// ErrPeerDisconnected.
|
||||
func (c *Conn) Context() context.Context {
|
||||
return c.ctx
|
||||
}
|
||||
|
||||
func (c *Conn) Write(p []byte) (n int, err error) {
|
||||
return c.writeFn(c.dstID, p)
|
||||
}
|
||||
|
||||
@@ -1,12 +1,9 @@
|
||||
package client
|
||||
|
||||
import (
|
||||
"container/list"
|
||||
"context"
|
||||
"fmt"
|
||||
"net"
|
||||
"net/netip"
|
||||
"reflect"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
@@ -43,8 +40,6 @@ func NewRelayTrack() *RelayTrack {
|
||||
}
|
||||
}
|
||||
|
||||
type OnServerCloseListener func()
|
||||
|
||||
// ManagerOption configures a Manager at construction time.
|
||||
type ManagerOption func(*Manager)
|
||||
|
||||
@@ -91,7 +86,6 @@ type Manager struct {
|
||||
relayClients map[string]*RelayTrack
|
||||
relayClientsMutex sync.RWMutex
|
||||
|
||||
onDisconnectedListeners map[string]*list.List
|
||||
onReconnectedListenerFn func()
|
||||
listenerLock sync.Mutex
|
||||
|
||||
@@ -126,10 +120,9 @@ func NewManager(ctx context.Context, serverURLs []string, peerID string, mtu uin
|
||||
ConnectionTimeout: defaultConnectionTimeout,
|
||||
TransportFallback: tf,
|
||||
},
|
||||
relayClients: make(map[string]*RelayTrack),
|
||||
onDisconnectedListeners: make(map[string]*list.List),
|
||||
cleanupInterval: relayCleanupInterval,
|
||||
keepUnusedServerTime: keepUnusedServerTime,
|
||||
relayClients: make(map[string]*RelayTrack),
|
||||
cleanupInterval: relayCleanupInterval,
|
||||
keepUnusedServerTime: keepUnusedServerTime,
|
||||
}
|
||||
for _, opt := range opts {
|
||||
opt(m)
|
||||
@@ -168,11 +161,11 @@ func (m *Manager) Serve() error {
|
||||
|
||||
// OpenConn opens a connection to the given peer key. If the peer is on the same relay server, the connection will be
|
||||
// established via the relay server. If the peer is on a different relay server, the manager will establish a new
|
||||
// connection to the relay server. It returns back with a net.Conn what represent the remote peer connection.
|
||||
// connection to the relay server. It returns the relayed connection to the remote peer.
|
||||
//
|
||||
// serverIP, when valid and serverAddress is foreign, is used as a dial target if the FQDN-based dial fails.
|
||||
// Ignored for the local home-server path. TLS verification still uses the FQDN via SNI.
|
||||
func (m *Manager) OpenConn(ctx context.Context, serverAddress, peerKey string, serverIP netip.Addr) (net.Conn, error) {
|
||||
func (m *Manager) OpenConn(ctx context.Context, serverAddress, peerKey string, serverIP netip.Addr) (*Conn, error) {
|
||||
m.relayClientMu.RLock()
|
||||
defer m.relayClientMu.RUnlock()
|
||||
|
||||
@@ -185,9 +178,7 @@ func (m *Manager) OpenConn(ctx context.Context, serverAddress, peerKey string, s
|
||||
return nil, err
|
||||
}
|
||||
|
||||
var (
|
||||
netConn net.Conn
|
||||
)
|
||||
var netConn *Conn
|
||||
if !foreign {
|
||||
log.Debugf("open peer connection via permanent server: %s", peerKey)
|
||||
netConn, err = m.relayClient.OpenConn(ctx, peerKey)
|
||||
@@ -220,31 +211,6 @@ func (m *Manager) SetOnReconnectedListener(f func()) {
|
||||
m.onReconnectedListenerFn = f
|
||||
}
|
||||
|
||||
// AddCloseListener adds a listener to the given server instance address. The listener will be called if the connection
|
||||
// closed.
|
||||
func (m *Manager) AddCloseListener(serverAddress string, onClosedListener OnServerCloseListener) error {
|
||||
m.relayClientMu.RLock()
|
||||
defer m.relayClientMu.RUnlock()
|
||||
|
||||
if m.relayClient == nil {
|
||||
return ErrRelayClientNotConnected
|
||||
}
|
||||
|
||||
foreign, err := m.isForeignServer(serverAddress)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
var listenerAddr string
|
||||
if foreign {
|
||||
listenerAddr = serverAddress
|
||||
} else {
|
||||
listenerAddr = m.relayClient.connectionURL
|
||||
}
|
||||
m.addListener(listenerAddr, onClosedListener)
|
||||
return nil
|
||||
}
|
||||
|
||||
// RelayInstanceAddress returns the address and resolved IP of the permanent relay server. It could change if the
|
||||
// network connection is lost. The address is sent to the target peer to choose the common relay server for the
|
||||
// communication; the IP is sent alongside so remote peers can dial directly without their own DNS lookup. Both
|
||||
@@ -256,11 +222,7 @@ func (m *Manager) RelayInstanceAddress() (string, netip.Addr, error) {
|
||||
if m.relayClient == nil {
|
||||
return "", netip.Addr{}, ErrRelayClientNotConnected
|
||||
}
|
||||
addr, err := m.relayClient.ServerInstanceURL()
|
||||
if err != nil {
|
||||
return "", netip.Addr{}, err
|
||||
}
|
||||
return addr, m.relayClient.ConnectedIP(), nil
|
||||
return m.relayClient.serverInstanceAddress()
|
||||
}
|
||||
|
||||
// ServerURLs returns the addresses of the relay servers.
|
||||
@@ -334,7 +296,7 @@ func (m *Manager) UpdateToken(token *relayAuth.Token) error {
|
||||
return m.tokenStore.UpdateToken(token)
|
||||
}
|
||||
|
||||
func (m *Manager) openConnVia(ctx context.Context, serverAddress, peerKey string, serverIP netip.Addr) (net.Conn, error) {
|
||||
func (m *Manager) openConnVia(ctx context.Context, serverAddress, peerKey string, serverIP netip.Addr) (*Conn, error) {
|
||||
// check if already has a connection to the desired relay server
|
||||
m.relayClientsMutex.RLock()
|
||||
rt, ok := m.relayClients[serverAddress]
|
||||
@@ -387,7 +349,7 @@ func (m *Manager) openConnVia(ctx context.Context, serverAddress, peerKey string
|
||||
// waiting for the dial started by another openConnVia call to finish. It waits
|
||||
// on rt.ready rather than the track lock, so it neither holds nor contends the
|
||||
// track lock across the dial.
|
||||
func (m *Manager) openConnOnTrack(ctx context.Context, rt *RelayTrack, peerKey string) (net.Conn, error) {
|
||||
func (m *Manager) openConnOnTrack(ctx context.Context, rt *RelayTrack, peerKey string) (*Conn, error) {
|
||||
select {
|
||||
case <-rt.ready:
|
||||
case <-ctx.Done():
|
||||
@@ -432,8 +394,6 @@ func (m *Manager) onServerDisconnected(serverAddress string) {
|
||||
if !isHome {
|
||||
m.evictForeignRelay(serverAddress)
|
||||
}
|
||||
|
||||
m.notifyOnDisconnectListeners(serverAddress)
|
||||
}
|
||||
|
||||
func (m *Manager) evictForeignRelay(serverAddress string) {
|
||||
@@ -527,36 +487,6 @@ func (m *Manager) cleanUpUnusedRelays() {
|
||||
}
|
||||
}
|
||||
|
||||
func (m *Manager) addListener(serverAddress string, onClosedListener OnServerCloseListener) {
|
||||
m.listenerLock.Lock()
|
||||
defer m.listenerLock.Unlock()
|
||||
l, ok := m.onDisconnectedListeners[serverAddress]
|
||||
if !ok {
|
||||
l = list.New()
|
||||
}
|
||||
for e := l.Front(); e != nil; e = e.Next() {
|
||||
if reflect.ValueOf(e.Value).Pointer() == reflect.ValueOf(onClosedListener).Pointer() {
|
||||
return
|
||||
}
|
||||
}
|
||||
l.PushBack(onClosedListener)
|
||||
m.onDisconnectedListeners[serverAddress] = l
|
||||
}
|
||||
|
||||
func (m *Manager) notifyOnDisconnectListeners(serverAddress string) {
|
||||
m.listenerLock.Lock()
|
||||
defer m.listenerLock.Unlock()
|
||||
|
||||
l, ok := m.onDisconnectedListeners[serverAddress]
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
for e := l.Front(); e != nil; e = e.Next() {
|
||||
go e.Value.(OnServerCloseListener)()
|
||||
}
|
||||
delete(m.onDisconnectedListeners, serverAddress)
|
||||
}
|
||||
|
||||
func relayConnState(c *Client) RelayConnState {
|
||||
addr, err := c.ServerInstanceURL()
|
||||
if err != nil {
|
||||
|
||||
@@ -0,0 +1,103 @@
|
||||
package client
|
||||
|
||||
import (
|
||||
"net/netip"
|
||||
"runtime"
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
func TestManager_RelayInstanceAddressAcrossReconnect(t *testing.T) {
|
||||
relays := []struct {
|
||||
url *RelayAddr
|
||||
conn stubConn
|
||||
ip netip.Addr
|
||||
}{
|
||||
{
|
||||
url: &RelayAddr{addr: "rels://relay-a.example:443"},
|
||||
conn: stubConn{remote: staticAddr{s: "192.0.2.1:443"}},
|
||||
ip: netip.MustParseAddr("192.0.2.1"),
|
||||
},
|
||||
{
|
||||
url: &RelayAddr{addr: "rels://relay-b.example:443"},
|
||||
conn: stubConn{remote: staticAddr{s: "192.0.2.2:443"}},
|
||||
ip: netip.MustParseAddr("192.0.2.2"),
|
||||
},
|
||||
}
|
||||
c := &Client{
|
||||
instanceURL: relays[0].url,
|
||||
relayConn: relays[0].conn,
|
||||
serviceIsRunning: true,
|
||||
}
|
||||
m := &Manager{relayClient: c}
|
||||
started := make(chan struct{})
|
||||
stop := make(chan struct{})
|
||||
done := make(chan struct{})
|
||||
t.Cleanup(func() {
|
||||
close(stop)
|
||||
<-done
|
||||
})
|
||||
go func() {
|
||||
defer close(done)
|
||||
for i := 0; ; i++ {
|
||||
select {
|
||||
case <-stop:
|
||||
return
|
||||
default:
|
||||
}
|
||||
// Publish successive connection states using the lifecycle locks.
|
||||
// Yield before publication so a getter using only muInstanceURL
|
||||
// can read the old URL while waiting for the new connection's IP.
|
||||
c.mu.Lock()
|
||||
runtime.Gosched()
|
||||
relay := relays[i%len(relays)]
|
||||
c.muInstanceURL.Lock()
|
||||
c.instanceURL = relay.url
|
||||
c.muInstanceURL.Unlock()
|
||||
c.relayConn = relay.conn
|
||||
c.mu.Unlock()
|
||||
if i == 0 {
|
||||
close(started)
|
||||
}
|
||||
}
|
||||
}()
|
||||
<-started
|
||||
|
||||
for range 1000 {
|
||||
url, ip, err := m.RelayInstanceAddress()
|
||||
require.NoError(t, err)
|
||||
wantIP := relays[0].ip
|
||||
if url == relays[1].url.String() {
|
||||
wantIP = relays[1].ip
|
||||
}
|
||||
if !assert.Equal(t, wantIP, ip, "advertised IP must belong to relay %s", url) {
|
||||
return
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestManager_RelayInstanceAddressDisconnected(t *testing.T) {
|
||||
for _, tt := range []struct {
|
||||
name string
|
||||
client *Client
|
||||
}{
|
||||
{name: "no client"},
|
||||
{name: "not connected", client: &Client{}},
|
||||
{
|
||||
name: "closed connection",
|
||||
client: &Client{
|
||||
relayConn: stubConn{remote: staticAddr{s: "192.0.2.1:443"}},
|
||||
},
|
||||
},
|
||||
} {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
m := &Manager{relayClient: tt.client}
|
||||
url, ip, err := m.RelayInstanceAddress()
|
||||
assert.Error(t, err)
|
||||
assert.Empty(t, url, "disconnected relay must not advertise a URL")
|
||||
assert.False(t, ip.IsValid(), "disconnected relay must not advertise a stale IP")
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -2,7 +2,9 @@ package client
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"net"
|
||||
"net/netip"
|
||||
"testing"
|
||||
"time"
|
||||
@@ -291,35 +293,29 @@ func TestForeignAutoClose(t *testing.T) {
|
||||
t.Fatalf("failed to serve manager: %s", err)
|
||||
}
|
||||
|
||||
// Set up a disconnect listener to track when foreign server disconnects
|
||||
foreignServerURL := toURL(srvCfg2)[0]
|
||||
disconnected := make(chan struct{})
|
||||
onDisconnect := func() {
|
||||
select {
|
||||
case disconnected <- struct{}{}:
|
||||
default:
|
||||
}
|
||||
}
|
||||
|
||||
t.Log("open connection to another peer")
|
||||
if _, err = mgr.OpenConn(ctx, foreignServerURL, "anotherpeer", netip.Addr{}); err == nil {
|
||||
t.Fatalf("should have failed to open connection to another peer")
|
||||
}
|
||||
|
||||
// Add the disconnect listener after the connection attempt
|
||||
if err := mgr.AddCloseListener(foreignServerURL, onDisconnect); err != nil {
|
||||
t.Logf("failed to add close listener (expected if connection failed): %s", err)
|
||||
}
|
||||
|
||||
// Wait for cleanup to happen
|
||||
timeout := relayCleanupInterval + keepUnusedServerTime + 2*time.Second
|
||||
t.Logf("waiting for relay cleanup: %s", timeout)
|
||||
|
||||
select {
|
||||
case <-disconnected:
|
||||
t.Log("foreign relay connection cleaned up successfully")
|
||||
case <-time.After(timeout):
|
||||
t.Log("timeout waiting for cleanup - this might be expected if connection never established")
|
||||
deadline := time.After(timeout)
|
||||
for {
|
||||
mgr.relayClientsMutex.RLock()
|
||||
_, tracked := mgr.relayClients[foreignServerURL]
|
||||
mgr.relayClientsMutex.RUnlock()
|
||||
if !tracked {
|
||||
t.Log("foreign relay connection cleaned up successfully")
|
||||
break
|
||||
}
|
||||
select {
|
||||
case <-deadline:
|
||||
t.Fatal("foreign relay was not cleaned up")
|
||||
case <-time.After(200 * time.Millisecond):
|
||||
}
|
||||
}
|
||||
|
||||
t.Logf("closing manager")
|
||||
@@ -413,23 +409,24 @@ func waitForReady(ctx context.Context, m *Manager, timeout time.Duration) error
|
||||
return fmt.Errorf("manager not ready within %s", timeout)
|
||||
}
|
||||
|
||||
func TestNotifierDoubleAdd(t *testing.T) {
|
||||
func toURL(address server.ListenerConfig) []string {
|
||||
return []string{"rel://" + address.Address}
|
||||
}
|
||||
|
||||
func TestConnContextCancelledOnServerDisconnect(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
|
||||
listenerCfg1 := server.ListenerConfig{
|
||||
Address: "localhost:52501",
|
||||
}
|
||||
srv, err := server.NewServer(newManagerTestServerConfig(listenerCfg1.Address))
|
||||
srvCfg := server.ListenerConfig{Address: "localhost:52601"}
|
||||
srv, err := server.NewServer(newManagerTestServerConfig(srvCfg.Address))
|
||||
if err != nil {
|
||||
t.Fatalf("failed to create server: %s", err)
|
||||
}
|
||||
errChan := make(chan error, 1)
|
||||
go func() {
|
||||
if err := srv.Listen(listenerCfg1); err != nil {
|
||||
if err := srv.Listen(srvCfg); err != nil {
|
||||
errChan <- err
|
||||
}
|
||||
}()
|
||||
|
||||
defer func() {
|
||||
if err := srv.Shutdown(ctx); err != nil {
|
||||
t.Errorf("failed to close server: %s", err)
|
||||
@@ -440,46 +437,106 @@ func TestNotifierDoubleAdd(t *testing.T) {
|
||||
t.Fatalf("failed to start server: %s", err)
|
||||
}
|
||||
|
||||
log.Debugf("connect by alice")
|
||||
mCtx, cancel := context.WithCancel(ctx)
|
||||
defer cancel()
|
||||
|
||||
clientBob := NewManager(mCtx, toURL(listenerCfg1), "bob", iface.DefaultMTU)
|
||||
if err = clientBob.Serve(); err != nil {
|
||||
mgrBob := NewManager(mCtx, toURL(srvCfg), "bob", iface.DefaultMTU)
|
||||
if err := mgrBob.Serve(); err != nil {
|
||||
t.Fatalf("failed to serve bob manager: %s", err)
|
||||
}
|
||||
|
||||
mgr := NewManager(mCtx, toURL(srvCfg), "alice", iface.DefaultMTU)
|
||||
if err := mgr.Serve(); err != nil {
|
||||
t.Fatalf("failed to serve manager: %s", err)
|
||||
}
|
||||
|
||||
clientAlice := NewManager(mCtx, toURL(listenerCfg1), "alice", iface.DefaultMTU)
|
||||
if err = clientAlice.Serve(); err != nil {
|
||||
ra, _, err := mgr.RelayInstanceAddress()
|
||||
if err != nil {
|
||||
t.Fatalf("failed to get relay address: %s", err)
|
||||
}
|
||||
|
||||
relayedConn, err := mgr.OpenConn(ctx, ra, "bob", netip.Addr{})
|
||||
if err != nil {
|
||||
t.Fatalf("failed to open conn: %s", err)
|
||||
}
|
||||
|
||||
select {
|
||||
case <-relayedConn.Context().Done():
|
||||
t.Fatal("conn context cancelled while the relay is still up")
|
||||
default:
|
||||
}
|
||||
|
||||
_ = mgr.relayClient.relayConn.Close()
|
||||
|
||||
select {
|
||||
case <-relayedConn.Context().Done():
|
||||
case <-time.After(15 * time.Second):
|
||||
t.Fatal("conn context was not cancelled after the relay connection dropped")
|
||||
}
|
||||
|
||||
if cause := context.Cause(relayedConn.Context()); !errors.Is(cause, ErrServerDisconnected) {
|
||||
t.Errorf("unexpected cancellation cause: %v, want %v", cause, ErrServerDisconnected)
|
||||
}
|
||||
}
|
||||
|
||||
func TestConnContextCauseOnLocalClose(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
|
||||
srvCfg := server.ListenerConfig{Address: "localhost:52602"}
|
||||
srv, err := server.NewServer(newManagerTestServerConfig(srvCfg.Address))
|
||||
if err != nil {
|
||||
t.Fatalf("failed to create server: %s", err)
|
||||
}
|
||||
errChan := make(chan error, 1)
|
||||
go func() {
|
||||
if err := srv.Listen(srvCfg); err != nil {
|
||||
errChan <- err
|
||||
}
|
||||
}()
|
||||
defer func() {
|
||||
if err := srv.Shutdown(ctx); err != nil {
|
||||
t.Errorf("failed to close server: %s", err)
|
||||
}
|
||||
}()
|
||||
|
||||
if err := waitForServerToStart(errChan); err != nil {
|
||||
t.Fatalf("failed to start server: %s", err)
|
||||
}
|
||||
|
||||
mCtx, cancel := context.WithCancel(ctx)
|
||||
defer cancel()
|
||||
|
||||
mgrBob := NewManager(mCtx, toURL(srvCfg), "bob", iface.DefaultMTU)
|
||||
if err := mgrBob.Serve(); err != nil {
|
||||
t.Fatalf("failed to serve bob manager: %s", err)
|
||||
}
|
||||
|
||||
mgr := NewManager(mCtx, toURL(srvCfg), "alice", iface.DefaultMTU)
|
||||
if err := mgr.Serve(); err != nil {
|
||||
t.Fatalf("failed to serve manager: %s", err)
|
||||
}
|
||||
|
||||
conn1, err := clientAlice.OpenConn(ctx, clientAlice.ServerURLs()[0], "bob", netip.Addr{})
|
||||
ra, _, err := mgr.RelayInstanceAddress()
|
||||
if err != nil {
|
||||
t.Fatalf("failed to bind channel: %s", err)
|
||||
t.Fatalf("failed to get relay address: %s", err)
|
||||
}
|
||||
|
||||
fnCloseListener := OnServerCloseListener(func() {
|
||||
log.Infof("close listener")
|
||||
})
|
||||
|
||||
err = clientAlice.AddCloseListener(clientAlice.ServerURLs()[0], fnCloseListener)
|
||||
relayedConn, err := mgr.OpenConn(ctx, ra, "bob", netip.Addr{})
|
||||
if err != nil {
|
||||
t.Fatalf("failed to add close listener: %s", err)
|
||||
t.Fatalf("failed to open conn: %s", err)
|
||||
}
|
||||
|
||||
err = clientAlice.AddCloseListener(clientAlice.ServerURLs()[0], fnCloseListener)
|
||||
if err != nil {
|
||||
t.Fatalf("failed to add close listener: %s", err)
|
||||
if err := relayedConn.Close(); err != nil {
|
||||
t.Fatalf("failed to close conn: %s", err)
|
||||
}
|
||||
|
||||
err = conn1.Close()
|
||||
if err != nil {
|
||||
t.Errorf("failed to close connection: %s", err)
|
||||
select {
|
||||
case <-relayedConn.Context().Done():
|
||||
case <-time.After(5 * time.Second):
|
||||
t.Fatal("conn context was not cancelled after a local close")
|
||||
}
|
||||
|
||||
}
|
||||
|
||||
func toURL(address server.ListenerConfig) []string {
|
||||
return []string{"rel://" + address.Address}
|
||||
if cause := context.Cause(relayedConn.Context()); !errors.Is(cause, net.ErrClosed) {
|
||||
t.Errorf("unexpected cancellation cause after a local close: %v, want %v", cause, net.ErrClosed)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -63,7 +63,12 @@ func (sp *ServerPicker) PickServer(parentCtx context.Context) (*Client, error) {
|
||||
if !ok {
|
||||
return nil, <-errChan
|
||||
}
|
||||
log.Infof("chosen home Relay server: %s", cr.Url)
|
||||
instanceURL, serverIP, err := cr.RelayClient.serverInstanceAddress()
|
||||
if err != nil {
|
||||
log.Infof("chosen home Relay server: %s, instance address unavailable: %v", cr.Url, err)
|
||||
return cr.RelayClient, nil
|
||||
}
|
||||
log.Infof("chosen home Relay server: %s, instance URL: %s, server IP: %s", cr.Url, instanceURL, serverIP)
|
||||
return cr.RelayClient, nil
|
||||
case <-ctx.Done():
|
||||
return nil, fmt.Errorf("connect to relay server: %w", ctx.Err())
|
||||
|
||||
Reference in New Issue
Block a user