mirror of
https://github.com/netbirdio/netbird.git
synced 2026-09-20 21:59:07 +02:00
[management, client, proxy] add expose NetBird-only services over tunnel peers (#6226)
Adds a new "private" service mode for the reverse proxy: services reachable exclusively over the embedded WireGuard tunnel, gated by per-peer group membership instead of operator auth schemes. Wire contract - ProxyMapping.private (field 13): the proxy MUST call ValidateTunnelPeer and fail closed; operator schemes are bypassed. - ProxyCapabilities.private (4) + supports_private_service (5): capability gate. Management never streams private mappings to proxies that don't claim the capability; the broadcast path applies the same filter via filterMappingsForProxy. - ValidateTunnelPeer RPC: resolves an inbound tunnel IP to a peer, checks the peer's groups against service.AccessGroups, and mints a session JWT on success. checkPeerGroupAccess fails closed when a private service has empty AccessGroups. - ValidateSession/ValidateTunnelPeer responses now carry peer_group_ids + peer_group_names so the proxy can authorise policy-aware middlewares without an extra management round-trip. - ProxyInboundListener + SendStatusUpdate.inbound_listener: per-account inbound listener state surfaced to dashboards. - PathTargetOptions.direct_upstream (11): bypass the embedded NetBird client and dial the target via the proxy host's network stack for upstreams reachable without WireGuard. Data model - Service.Private (bool) + Service.AccessGroups ([]string, JSON- serialised). Validate() rejects bearer auth on private services. Copy() deep-copies AccessGroups. pgx getServices loads the columns. - DomainConfig.Private threaded into the proxy auth middleware. Request handler routes private services through forwardWithTunnelPeer and returns 403 on validation failure. - Account-level SynthesizePrivateServiceZones (synthetic DNS) and injectPrivateServicePolicies (synthetic ACL) gate on len(svc.AccessGroups) > 0. Proxy - /netbird proxy --private (embedded mode) flag; Config.Private in proxy/lifecycle.go. - Per-account inbound listener (proxy/inbound.go) binding HTTP/HTTPS on the embedded NetBird client's WireGuard tunnel netstack. - proxy/internal/auth/tunnel_cache: ValidateTunnelPeer response cache with single-flight de-duplication and per-account eviction. - Local peerstore short-circuit: when the inbound IP isn't in the account roster, deny fast without an RPC. - proxy/server.go reports SupportsPrivateService=true and redacts the full ProxyMapping JSON from info logs (auth_token + header-auth hashed values now only at debug level). Identity forwarding - ValidateSessionJWT returns user_id, email, method, groups, group_names. sessionkey.Claims carries Email + Groups + GroupNames so the proxy can stamp identity onto upstream requests without an extra management round-trip on every cookie-bearing request. - CapturedData carries userEmail / userGroups / userGroupNames; the proxy stamps X-NetBird-User and X-NetBird-Groups on r.Out from the authenticated identity (strips client-supplied values first to prevent spoofing). - AccessLog.UserGroups: access-log enrichment captures the user's group memberships at write time so the dashboard can render group context without reverse-resolving stale memberships. OpenAPI/dashboard surface - ReverseProxyService gains private + access_groups; ReverseProxyCluster gains private + supports_private. ReverseProxyTarget target_type enum gains "cluster". ServiceTargetOptions gains direct_upstream. ProxyAccessLog gains user_groups.
This commit is contained in:
@@ -0,0 +1,112 @@
|
||||
package roundtrip
|
||||
|
||||
import (
|
||||
"crypto/tls"
|
||||
"errors"
|
||||
"net"
|
||||
"net/http"
|
||||
"time"
|
||||
|
||||
log "github.com/sirupsen/logrus"
|
||||
)
|
||||
|
||||
// MultiTransport dispatches each request to either the embedded NetBird
|
||||
// http.RoundTripper or a stdlib http.Transport based on a per-request
|
||||
// context flag set by the reverse-proxy rewrite step. When the flag is
|
||||
// absent (the default for every existing target), requests follow the
|
||||
// embedded NetBird path — current behaviour, preserved.
|
||||
//
|
||||
// The stdlib branch is used when a target was configured with
|
||||
// direct_upstream=true. It dials via the host's network stack, which is
|
||||
// what private (`netbird proxy`) deployments and centralised proxies
|
||||
// fronting host-reachable upstreams (public APIs, LAN services,
|
||||
// localhost sidecars) want.
|
||||
//
|
||||
// An embedded roundtripper is required. To run direct-only (no WG
|
||||
// branch at all), construct the MultiTransport via NewDirectOnly.
|
||||
type MultiTransport struct {
|
||||
embedded http.RoundTripper
|
||||
direct *http.Transport
|
||||
insecure *http.Transport
|
||||
}
|
||||
|
||||
// errNoEmbeddedTransport is returned when a request reaches the
|
||||
// embedded branch on a MultiTransport that wasn't given one. Surfaces
|
||||
// the misconfiguration to the caller instead of silently routing to
|
||||
// the direct branch, which would bypass the WG tunnel.
|
||||
var errNoEmbeddedTransport = errors.New("multitransport: embedded roundtripper not configured")
|
||||
|
||||
// NewMultiTransport wires both branches. embedded is the existing NetBird
|
||||
// roundtripper and must not be nil — pass to NewDirectOnly for a
|
||||
// MultiTransport that only ever uses the direct branch. The direct
|
||||
// branches honour the same NB_PROXY_* tuning env vars as the embedded
|
||||
// transport (see loadTransportConfig) plus a dial-timeout wrapper that
|
||||
// respects types.WithDialTimeout.
|
||||
func NewMultiTransport(embedded http.RoundTripper, logger *log.Logger) *MultiTransport {
|
||||
if logger == nil {
|
||||
logger = log.StandardLogger()
|
||||
}
|
||||
cfg := loadTransportConfig(logger)
|
||||
dialer := &net.Dialer{
|
||||
Timeout: 30 * time.Second,
|
||||
KeepAlive: 30 * time.Second,
|
||||
}
|
||||
direct := &http.Transport{
|
||||
DialContext: dialWithTimeout(dialer.DialContext),
|
||||
ForceAttemptHTTP2: true,
|
||||
MaxIdleConns: cfg.maxIdleConns,
|
||||
MaxIdleConnsPerHost: cfg.maxIdleConnsPerHost,
|
||||
MaxConnsPerHost: cfg.maxConnsPerHost,
|
||||
IdleConnTimeout: cfg.idleConnTimeout,
|
||||
TLSHandshakeTimeout: cfg.tlsHandshakeTimeout,
|
||||
ExpectContinueTimeout: cfg.expectContinueTimeout,
|
||||
ResponseHeaderTimeout: cfg.responseHeaderTimeout,
|
||||
WriteBufferSize: cfg.writeBufferSize,
|
||||
ReadBufferSize: cfg.readBufferSize,
|
||||
DisableCompression: cfg.disableCompression,
|
||||
}
|
||||
insecure := direct.Clone()
|
||||
insecure.TLSClientConfig = &tls.Config{InsecureSkipVerify: true} //nolint:gosec // matches the embedded NetBird transport's per-target opt-in
|
||||
|
||||
return &MultiTransport{
|
||||
embedded: embedded,
|
||||
direct: direct,
|
||||
insecure: insecure,
|
||||
}
|
||||
}
|
||||
|
||||
// NewDirectOnly returns a MultiTransport with no embedded branch.
|
||||
// Every request goes through the direct branch regardless of the
|
||||
// per-request flag, so the embedded path can never be reached
|
||||
// silently — wiring code that needs WG must use NewMultiTransport.
|
||||
func NewDirectOnly(logger *log.Logger) *MultiTransport {
|
||||
return NewMultiTransport(noEmbeddedRoundTripper{}, logger)
|
||||
}
|
||||
|
||||
// noEmbeddedRoundTripper is the sentinel embedded transport for
|
||||
// direct-only MultiTransports. RoundTrip is never called in practice
|
||||
// because the direct branch matches every request, but if anything
|
||||
// ever did reach this path it would fail loudly instead of falling
|
||||
// back to direct.
|
||||
type noEmbeddedRoundTripper struct{}
|
||||
|
||||
func (noEmbeddedRoundTripper) RoundTrip(*http.Request) (*http.Response, error) {
|
||||
return nil, errNoEmbeddedTransport
|
||||
}
|
||||
|
||||
// RoundTrip dispatches by reading the direct-upstream flag from the request
|
||||
// context. When set, the request is forwarded via the stdlib transport,
|
||||
// honouring the existing per-request skip-TLS-verify flag. Otherwise it
|
||||
// goes through the embedded NetBird roundtripper.
|
||||
func (m *MultiTransport) RoundTrip(req *http.Request) (*http.Response, error) {
|
||||
if DirectUpstreamFromContext(req.Context()) {
|
||||
if skipTLSVerifyFromContext(req.Context()) {
|
||||
return m.insecure.RoundTrip(req)
|
||||
}
|
||||
return m.direct.RoundTrip(req)
|
||||
}
|
||||
if m.embedded == nil {
|
||||
return nil, errNoEmbeddedTransport
|
||||
}
|
||||
return m.embedded.RoundTrip(req)
|
||||
}
|
||||
@@ -0,0 +1,134 @@
|
||||
package roundtrip
|
||||
|
||||
import (
|
||||
"context"
|
||||
"io"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
// stubRoundTripper records whether RoundTrip was called and returns a
|
||||
// canned response so tests can assert the dispatch decision without
|
||||
// running a real network.
|
||||
type stubRoundTripper struct {
|
||||
called bool
|
||||
body string
|
||||
}
|
||||
|
||||
func (s *stubRoundTripper) RoundTrip(_ *http.Request) (*http.Response, error) {
|
||||
s.called = true
|
||||
return &http.Response{
|
||||
StatusCode: http.StatusOK,
|
||||
Body: io.NopCloser(strings.NewReader(s.body)),
|
||||
Header: http.Header{},
|
||||
}, nil
|
||||
}
|
||||
|
||||
func TestMultiTransport_DispatchesByContextFlag(t *testing.T) {
|
||||
embedded := &stubRoundTripper{body: "embedded"}
|
||||
mt := NewMultiTransport(embedded, nil)
|
||||
|
||||
t.Run("default routes to embedded", func(t *testing.T) {
|
||||
embedded.called = false
|
||||
req := httptest.NewRequest(http.MethodGet, "http://example.invalid", nil)
|
||||
resp, err := mt.RoundTrip(req)
|
||||
require.NoError(t, err, "embedded path must not error on stubbed transport")
|
||||
require.NotNil(t, resp)
|
||||
_ = resp.Body.Close()
|
||||
assert.True(t, embedded.called, "request without WithDirectUpstream must hit the embedded transport")
|
||||
})
|
||||
|
||||
t.Run("WithDirectUpstream skips embedded", func(t *testing.T) {
|
||||
embedded.called = false
|
||||
// Hit a server we control to verify the stdlib transport is used.
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
|
||||
_, _ = io.WriteString(w, "direct")
|
||||
}))
|
||||
defer srv.Close()
|
||||
|
||||
req, err := http.NewRequestWithContext(WithDirectUpstream(context.Background()), http.MethodGet, srv.URL, nil)
|
||||
require.NoError(t, err)
|
||||
resp, err := mt.RoundTrip(req)
|
||||
require.NoError(t, err, "direct path must dial via stdlib transport")
|
||||
body, err := io.ReadAll(resp.Body)
|
||||
_ = resp.Body.Close()
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, "direct", string(body), "stdlib transport must reach the test server")
|
||||
assert.False(t, embedded.called, "WithDirectUpstream must bypass the embedded transport")
|
||||
})
|
||||
}
|
||||
|
||||
// TestMultiTransport_AppliesEnvOverridesToDirect verifies that the
|
||||
// NB_PROXY_* env vars consumed by loadTransportConfig flow into the
|
||||
// direct branches (previously they only applied to the embedded
|
||||
// roundtripper, so direct-upstream traffic ignored operator tuning).
|
||||
func TestMultiTransport_AppliesEnvOverridesToDirect(t *testing.T) {
|
||||
t.Setenv(EnvMaxIdleConns, "42")
|
||||
t.Setenv(EnvIdleConnTimeout, "11s")
|
||||
t.Setenv(EnvTLSHandshakeTimeout, "7s")
|
||||
|
||||
mt := NewMultiTransport(&stubRoundTripper{body: "embedded"}, nil)
|
||||
|
||||
assert.Equal(t, 42, mt.direct.MaxIdleConns,
|
||||
"NB_PROXY_MAX_IDLE_CONNS must propagate to the direct transport")
|
||||
assert.Equal(t, 11*time.Second, mt.direct.IdleConnTimeout,
|
||||
"NB_PROXY_IDLE_CONN_TIMEOUT must propagate to the direct transport")
|
||||
assert.Equal(t, 7*time.Second, mt.direct.TLSHandshakeTimeout,
|
||||
"NB_PROXY_TLS_HANDSHAKE_TIMEOUT must propagate to the direct transport")
|
||||
assert.Equal(t, 42, mt.insecure.MaxIdleConns,
|
||||
"env tuning must also apply to the insecure-skip-verify direct transport")
|
||||
}
|
||||
|
||||
// TestMultiTransport_NilEmbeddedErrorsWhenWGPathRequested guards
|
||||
// against the previous silent fallback: a MultiTransport constructed
|
||||
// without an embedded transport must reject requests that don't
|
||||
// explicitly opt into the direct branch, rather than routing them
|
||||
// over the host stack and bypassing WireGuard.
|
||||
func TestMultiTransport_NilEmbeddedErrorsWhenWGPathRequested(t *testing.T) {
|
||||
mt := NewMultiTransport(nil, nil)
|
||||
|
||||
req := httptest.NewRequest(http.MethodGet, "http://example.invalid", nil)
|
||||
resp, err := mt.RoundTrip(req)
|
||||
if resp != nil {
|
||||
_ = resp.Body.Close()
|
||||
}
|
||||
require.Error(t, err, "nil embedded must surface as an explicit error, not a silent direct dispatch")
|
||||
assert.Nil(t, resp)
|
||||
assert.ErrorIs(t, err, errNoEmbeddedTransport,
|
||||
"the error must be the sentinel so callers can distinguish misconfiguration from network failures")
|
||||
}
|
||||
|
||||
// TestMultiTransport_DirectOnlyServesDirectBranch verifies NewDirectOnly
|
||||
// constructs a MultiTransport whose direct branch handles requests with
|
||||
// the direct-upstream flag set, and surfaces the explicit sentinel
|
||||
// when the embedded path is reached.
|
||||
func TestMultiTransport_DirectOnlyServesDirectBranch(t *testing.T) {
|
||||
mt := NewDirectOnly(nil)
|
||||
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
|
||||
_, _ = io.WriteString(w, "ok")
|
||||
}))
|
||||
defer srv.Close()
|
||||
|
||||
req, err := http.NewRequestWithContext(WithDirectUpstream(context.Background()), http.MethodGet, srv.URL, nil)
|
||||
require.NoError(t, err)
|
||||
resp, err := mt.RoundTrip(req)
|
||||
require.NoError(t, err, "direct-only must serve requests that opt into the direct branch")
|
||||
_ = resp.Body.Close()
|
||||
assert.Equal(t, http.StatusOK, resp.StatusCode)
|
||||
|
||||
wgReq := httptest.NewRequest(http.MethodGet, "http://example.invalid", nil)
|
||||
resp, err = mt.RoundTrip(wgReq)
|
||||
if resp != nil {
|
||||
_ = resp.Body.Close()
|
||||
}
|
||||
require.Error(t, err, "direct-only must refuse requests that didn't opt into the direct branch")
|
||||
assert.Nil(t, resp)
|
||||
assert.ErrorIs(t, err, errNoEmbeddedTransport)
|
||||
}
|
||||
@@ -7,6 +7,7 @@ import (
|
||||
"fmt"
|
||||
"net"
|
||||
"net/http"
|
||||
"net/netip"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
@@ -76,11 +77,11 @@ type clientEntry struct {
|
||||
services map[ServiceKey]serviceInfo
|
||||
createdAt time.Time
|
||||
started bool
|
||||
// ready is closed once the client has been fully initialized.
|
||||
// Callers that find a pending entry wait on this channel before
|
||||
// accessing the client. A nil initErr means success.
|
||||
ready chan struct{}
|
||||
initErr error
|
||||
// inbound is opaque per-account state owned by the NetBird parent's
|
||||
// ReadyHandler. The roundtrip package never inspects this value; it
|
||||
// only stores it so RemovePeer / StopAll can hand it back to the
|
||||
// matching StopHandler. Nil when no inbound integration is active.
|
||||
inbound any
|
||||
// Per-backend in-flight limiting keyed by target host:port.
|
||||
// TODO: clean up stale entries when backend targets change.
|
||||
inflightMu sync.Mutex
|
||||
@@ -88,6 +89,19 @@ type clientEntry struct {
|
||||
maxInflight int
|
||||
}
|
||||
|
||||
// IdentityForIP resolves a tunnel IP to the peer identity locally known by
|
||||
// this account's embedded client. Returns (pubKey, fqdn) on success.
|
||||
// ok=false means the IP is not in the account's roster — callers can use
|
||||
// that as a fast deny without round-tripping management. The returned
|
||||
// strings carry only what the embedded peerstore exposes; user identity
|
||||
// (UserID / Email / Groups) still flows through ValidateTunnelPeer.
|
||||
func (e *clientEntry) IdentityForIP(ip netip.Addr) (pubKey, fqdn string, ok bool) {
|
||||
if e == nil || e.client == nil || !ip.IsValid() {
|
||||
return "", "", false
|
||||
}
|
||||
return e.client.IdentityForIP(ip)
|
||||
}
|
||||
|
||||
// acquireInflight attempts to acquire an in-flight slot for the given backend.
|
||||
// It returns a release function that must always be called, and true on success.
|
||||
func (e *clientEntry) acquireInflight(backend backendKey) (release func(), ok bool) {
|
||||
@@ -117,6 +131,12 @@ type ClientConfig struct {
|
||||
MgmtAddr string
|
||||
WGPort uint16
|
||||
PreSharedKey string
|
||||
// BlockInbound mirrors embed.Options.BlockInbound. Set to true on the
|
||||
// standalone proxy where the embedded client never accepts inbound;
|
||||
// set to false on the private/embedded proxy so the engine creates
|
||||
// the ACL manager and applies management's per-policy firewall rules
|
||||
// (which is what gates per-account inbound listeners on the netstack).
|
||||
BlockInbound bool
|
||||
}
|
||||
|
||||
type statusNotifier interface {
|
||||
@@ -142,6 +162,14 @@ type NetBird struct {
|
||||
clients map[types.AccountID]*clientEntry
|
||||
initLogOnce sync.Once
|
||||
statusNotifier statusNotifier
|
||||
// readyHandler runs after the embedded client for an account reports
|
||||
// Ready. The opaque return value is stored on clientEntry and handed
|
||||
// back to stopHandler when the entry is torn down. Nil disables the
|
||||
// hook entirely (default for the standalone proxy).
|
||||
readyHandler func(ctx context.Context, accountID types.AccountID, client *embed.Client) any
|
||||
// stopHandler runs when an account's last service is removed (or the
|
||||
// transport is shutting down). Receives whatever readyHandler returned.
|
||||
stopHandler func(accountID types.AccountID, state any)
|
||||
|
||||
// OnAddPeer, when set, is called after AddPeer completes for a new account
|
||||
// (i.e. when a new client was actually created, not when an existing one
|
||||
@@ -167,9 +195,6 @@ type skipTLSVerifyContextKey struct{}
|
||||
// AddPeer registers a service for an account. If the account doesn't have a client yet,
|
||||
// one is created by authenticating with the management server using the provided token.
|
||||
// Multiple services can share the same client.
|
||||
//
|
||||
// Client creation (WG keygen, gRPC, embed.New) runs without holding clientsMux
|
||||
// so that concurrent AddPeer calls for different accounts execute in parallel.
|
||||
func (n *NetBird) AddPeer(ctx context.Context, accountID types.AccountID, key ServiceKey, authToken string, serviceID types.ServiceID) error {
|
||||
si := serviceInfo{serviceID: serviceID}
|
||||
|
||||
@@ -177,23 +202,10 @@ func (n *NetBird) AddPeer(ctx context.Context, accountID types.AccountID, key Se
|
||||
|
||||
entry, exists := n.clients[accountID]
|
||||
if exists {
|
||||
ready := entry.ready
|
||||
entry.services[key] = si
|
||||
started := entry.started
|
||||
n.clientsMux.Unlock()
|
||||
|
||||
// If the entry is still being initialized by another goroutine, wait.
|
||||
if ready != nil {
|
||||
select {
|
||||
case <-ready:
|
||||
case <-ctx.Done():
|
||||
return ctx.Err()
|
||||
}
|
||||
if entry.initErr != nil {
|
||||
return fmt.Errorf("peer initialization failed: %w", entry.initErr)
|
||||
}
|
||||
}
|
||||
|
||||
n.logger.WithFields(log.Fields{
|
||||
"account_id": accountID,
|
||||
"service_key": key,
|
||||
@@ -210,43 +222,19 @@ func (n *NetBird) AddPeer(ctx context.Context, accountID types.AccountID, key Se
|
||||
return nil
|
||||
}
|
||||
|
||||
// Insert a placeholder so other goroutines calling AddPeer for the same
|
||||
// account will wait on the ready channel instead of starting a second
|
||||
// client creation.
|
||||
entry = &clientEntry{
|
||||
services: map[ServiceKey]serviceInfo{key: si},
|
||||
ready: make(chan struct{}),
|
||||
}
|
||||
n.clients[accountID] = entry
|
||||
n.clientsMux.Unlock()
|
||||
|
||||
createStart := time.Now()
|
||||
created, err := n.createClientEntry(ctx, accountID, key, authToken, si)
|
||||
entry, err := n.createClientEntry(ctx, accountID, key, authToken, si)
|
||||
if n.OnAddPeer != nil {
|
||||
n.OnAddPeer(time.Since(createStart), err)
|
||||
}
|
||||
if err != nil {
|
||||
entry.initErr = err
|
||||
close(entry.ready)
|
||||
|
||||
n.clientsMux.Lock()
|
||||
delete(n.clients, accountID)
|
||||
n.clientsMux.Unlock()
|
||||
return err
|
||||
}
|
||||
|
||||
// Transfer any services that were registered by concurrent AddPeer calls
|
||||
// while we were creating the client.
|
||||
n.clientsMux.Lock()
|
||||
for k, v := range entry.services {
|
||||
created.services[k] = v
|
||||
}
|
||||
created.ready = nil
|
||||
n.clients[accountID] = created
|
||||
n.clients[accountID] = entry
|
||||
n.clientsMux.Unlock()
|
||||
|
||||
close(entry.ready)
|
||||
|
||||
n.logger.WithFields(log.Fields{
|
||||
"account_id": accountID,
|
||||
"service_key": key,
|
||||
@@ -254,13 +242,13 @@ func (n *NetBird) AddPeer(ctx context.Context, accountID types.AccountID, key Se
|
||||
|
||||
// Attempt to start the client in the background; if this fails we will
|
||||
// retry on the first request via RoundTrip.
|
||||
go n.runClientStartup(ctx, accountID, created.client)
|
||||
go n.runClientStartup(ctx, accountID, entry.client)
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// createClientEntry generates a WireGuard keypair, authenticates with management,
|
||||
// and creates an embedded NetBird client.
|
||||
// and creates an embedded NetBird client. Must be called with clientsMux held.
|
||||
func (n *NetBird) createClientEntry(ctx context.Context, accountID types.AccountID, key ServiceKey, authToken string, si serviceInfo) (*clientEntry, error) {
|
||||
serviceID := si.serviceID
|
||||
n.logger.WithFields(log.Fields{
|
||||
@@ -318,9 +306,15 @@ func (n *NetBird) createClientEntry(ctx context.Context, accountID types.Account
|
||||
ManagementURL: n.clientCfg.MgmtAddr,
|
||||
PrivateKey: privateKey.String(),
|
||||
LogLevel: log.WarnLevel.String(),
|
||||
BlockInbound: true,
|
||||
WireguardPort: &wgPort,
|
||||
PreSharedKey: n.clientCfg.PreSharedKey,
|
||||
BlockInbound: n.clientCfg.BlockInbound,
|
||||
// The embedded proxy peer must never be a stepping stone into
|
||||
// the proxy host's LAN: it only exists to reach NetBird mesh
|
||||
// targets or, when direct_upstream is set, the host network
|
||||
// stack via the MultiTransport's direct branch (which bypasses
|
||||
// the engine routing entirely).
|
||||
BlockLANAccess: true,
|
||||
WireguardPort: &wgPort,
|
||||
PreSharedKey: n.clientCfg.PreSharedKey,
|
||||
})
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("create netbird client: %w", err)
|
||||
@@ -385,8 +379,25 @@ func (n *NetBird) runClientStartup(ctx context.Context, accountID types.AccountI
|
||||
toNotify = append(toNotify, serviceNotification{key: key, serviceID: info.serviceID})
|
||||
}
|
||||
}
|
||||
readyHandler := n.readyHandler
|
||||
n.clientsMux.Unlock()
|
||||
|
||||
if readyHandler != nil {
|
||||
state := readyHandler(ctx, accountID, client)
|
||||
n.clientsMux.Lock()
|
||||
if e, ok := n.clients[accountID]; ok {
|
||||
e.inbound = state
|
||||
} else if state != nil && n.stopHandler != nil {
|
||||
// Account was removed while readyHandler ran; tear down the
|
||||
// resources it just brought up.
|
||||
stop := n.stopHandler
|
||||
n.clientsMux.Unlock()
|
||||
stop(accountID, state)
|
||||
n.clientsMux.Lock()
|
||||
}
|
||||
n.clientsMux.Unlock()
|
||||
}
|
||||
|
||||
if n.statusNotifier == nil {
|
||||
return
|
||||
}
|
||||
@@ -432,11 +443,15 @@ func (n *NetBird) RemovePeer(ctx context.Context, accountID types.AccountID, key
|
||||
stopClient := len(entry.services) == 0
|
||||
var client *embed.Client
|
||||
var transport, insecureTransport *http.Transport
|
||||
var inbound any
|
||||
var stopHandler func(types.AccountID, any)
|
||||
if stopClient {
|
||||
n.logger.WithField("account_id", accountID).Info("stopping client, no more services")
|
||||
client = entry.client
|
||||
transport = entry.transport
|
||||
insecureTransport = entry.insecureTransport
|
||||
inbound = entry.inbound
|
||||
stopHandler = n.stopHandler
|
||||
delete(n.clients, accountID)
|
||||
} else {
|
||||
n.logger.WithFields(log.Fields{
|
||||
@@ -450,6 +465,9 @@ func (n *NetBird) RemovePeer(ctx context.Context, accountID types.AccountID, key
|
||||
n.notifyDisconnect(ctx, accountID, key, si.serviceID)
|
||||
|
||||
if stopClient {
|
||||
if inbound != nil && stopHandler != nil {
|
||||
stopHandler(accountID, inbound)
|
||||
}
|
||||
transport.CloseIdleConnections()
|
||||
insecureTransport.CloseIdleConnections()
|
||||
if err := client.Stop(ctx); err != nil {
|
||||
@@ -536,8 +554,12 @@ func (n *NetBird) StopAll(ctx context.Context) error {
|
||||
n.clientsMux.Lock()
|
||||
defer n.clientsMux.Unlock()
|
||||
|
||||
stopHandler := n.stopHandler
|
||||
var merr *multierror.Error
|
||||
for accountID, entry := range n.clients {
|
||||
if entry.inbound != nil && stopHandler != nil {
|
||||
stopHandler(accountID, entry.inbound)
|
||||
}
|
||||
entry.transport.CloseIdleConnections()
|
||||
entry.insecureTransport.CloseIdleConnections()
|
||||
if err := entry.client.Stop(ctx); err != nil {
|
||||
@@ -590,6 +612,19 @@ func (n *NetBird) GetClient(accountID types.AccountID) (*embed.Client, bool) {
|
||||
return entry.client, true
|
||||
}
|
||||
|
||||
// IdentityForIP resolves a tunnel IP to a peer identity local to the given
|
||||
// account. Delegates to clientEntry.IdentityForIP. Returns ok=false when
|
||||
// the account has no client or the IP is not in its peerstore.
|
||||
func (n *NetBird) IdentityForIP(accountID types.AccountID, ip netip.Addr) (pubKey, fqdn string, ok bool) {
|
||||
n.clientsMux.RLock()
|
||||
entry, exists := n.clients[accountID]
|
||||
n.clientsMux.RUnlock()
|
||||
if !exists {
|
||||
return "", "", false
|
||||
}
|
||||
return entry.IdentityForIP(ip)
|
||||
}
|
||||
|
||||
// ListClientsForDebug returns information about all clients for debug purposes.
|
||||
func (n *NetBird) ListClientsForDebug() map[types.AccountID]ClientDebugInfo {
|
||||
n.clientsMux.RLock()
|
||||
@@ -645,6 +680,18 @@ func NewNetBird(proxyID, proxyAddr string, clientCfg ClientConfig, logger *log.L
|
||||
}
|
||||
}
|
||||
|
||||
// SetClientLifecycle registers callbacks that run when an embedded
|
||||
// client becomes ready and when its entry is torn down. The opaque value
|
||||
// returned by ready is stored on the entry and handed back to stop on
|
||||
// cleanup. Must be called before AddPeer. A nil pair leaves the
|
||||
// outbound-only behaviour intact.
|
||||
func (n *NetBird) SetClientLifecycle(ready func(ctx context.Context, accountID types.AccountID, client *embed.Client) any, stop func(accountID types.AccountID, state any)) {
|
||||
n.clientsMux.Lock()
|
||||
defer n.clientsMux.Unlock()
|
||||
n.readyHandler = ready
|
||||
n.stopHandler = stop
|
||||
}
|
||||
|
||||
// dialWithTimeout wraps a DialContext function so that any dial timeout
|
||||
// stored in the context (via types.WithDialTimeout) is applied only to
|
||||
// the connection establishment phase, not the full request lifetime.
|
||||
@@ -687,3 +734,22 @@ func skipTLSVerifyFromContext(ctx context.Context) bool {
|
||||
v, _ := ctx.Value(skipTLSVerifyContextKey{}).(bool)
|
||||
return v
|
||||
}
|
||||
|
||||
// directUpstreamContextKey signals that the request should bypass the embedded
|
||||
// NetBird WireGuard client and dial via the host's network stack instead.
|
||||
// Set by the reverse-proxy rewrite step when the matched target carries
|
||||
// PathTarget.DirectUpstream; consumed by MultiTransport.
|
||||
type directUpstreamContextKey struct{}
|
||||
|
||||
// WithDirectUpstream marks the context so MultiTransport routes the request
|
||||
// through its stdlib transport instead of the embedded NetBird roundtripper.
|
||||
func WithDirectUpstream(ctx context.Context) context.Context {
|
||||
return context.WithValue(ctx, directUpstreamContextKey{}, true)
|
||||
}
|
||||
|
||||
// DirectUpstreamFromContext reports whether the context has been marked to
|
||||
// bypass the embedded NetBird client.
|
||||
func DirectUpstreamFromContext(ctx context.Context) bool {
|
||||
v, _ := ctx.Value(directUpstreamContextKey{}).(bool)
|
||||
return v
|
||||
}
|
||||
|
||||
@@ -3,6 +3,7 @@ package roundtrip
|
||||
import (
|
||||
"context"
|
||||
"net/http"
|
||||
"net/netip"
|
||||
"sync"
|
||||
"testing"
|
||||
|
||||
@@ -305,6 +306,36 @@ func TestNetBird_AddPeer_ExistingStartedClient_NotifiesStatus(t *testing.T) {
|
||||
assert.True(t, calls[0].connected)
|
||||
}
|
||||
|
||||
// TestNetBird_IdentityForIP_UnknownAccountReturnsFalse confirms that the
|
||||
// public lookup short-circuits when no client has been registered for
|
||||
// the queried account. The auth middleware uses ok=false as a fast deny.
|
||||
func TestNetBird_IdentityForIP_UnknownAccountReturnsFalse(t *testing.T) {
|
||||
nb := mockNetBird()
|
||||
_, _, ok := nb.IdentityForIP("acct-missing", netip.MustParseAddr("100.64.0.10"))
|
||||
assert.False(t, ok, "unknown account must yield ok=false")
|
||||
}
|
||||
|
||||
// TestClientEntry_IdentityForIP_NilClientGuard ensures the receiver
|
||||
// methods stay safe when called on partially-initialized state, which
|
||||
// can happen briefly during AddPeer setup or test fixtures.
|
||||
func TestClientEntry_IdentityForIP_NilClientGuard(t *testing.T) {
|
||||
var e *clientEntry
|
||||
_, _, ok := e.IdentityForIP(netip.MustParseAddr("100.64.0.10"))
|
||||
assert.False(t, ok, "nil clientEntry must yield ok=false")
|
||||
|
||||
e = &clientEntry{}
|
||||
_, _, ok = e.IdentityForIP(netip.MustParseAddr("100.64.0.10"))
|
||||
assert.False(t, ok, "clientEntry with nil embed.Client must yield ok=false")
|
||||
}
|
||||
|
||||
// TestClientEntry_IdentityForIP_InvalidIPReturnsFalse covers the input
|
||||
// guard so callers don't have to repeat the check.
|
||||
func TestClientEntry_IdentityForIP_InvalidIPReturnsFalse(t *testing.T) {
|
||||
e := &clientEntry{}
|
||||
_, _, ok := e.IdentityForIP(netip.Addr{})
|
||||
assert.False(t, ok, "invalid IP must yield ok=false")
|
||||
}
|
||||
|
||||
func TestNetBird_RemovePeer_NotifiesDisconnection(t *testing.T) {
|
||||
notifier := &mockStatusNotifier{}
|
||||
nb := NewNetBird("test-proxy", "invalid.test", ClientConfig{
|
||||
|
||||
Reference in New Issue
Block a user