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:
riccardom
2026-09-28 10:56:27 +02:00
320 changed files with 25425 additions and 12010 deletions
+57
View File
@@ -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()
}
+43
View File
@@ -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")
}
+7 -1
View File
@@ -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()
+116 -8
View File
@@ -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
}
+5
View File
@@ -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)
}
+127
View File
@@ -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())
}
+202
View File
@@ -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))
}
}
+42 -34
View File
@@ -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 {
+20 -25
View File
@@ -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)
})
}
}
+10
View File
@@ -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)
}
+9 -79
View File
@@ -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 {
+103
View File
@@ -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")
})
}
}
+107 -50
View File
@@ -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)
}
}
+6 -1
View File
@@ -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())