mirror of
https://github.com/netbirdio/netbird.git
synced 2026-09-23 23:29:08 +02:00
feat(private-service): expose NetBird-only services over tunnel peers
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,163 @@
|
||||
package roundtrip
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"crypto/tls"
|
||||
"io"
|
||||
"net"
|
||||
"net/http"
|
||||
"sort"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
log "github.com/sirupsen/logrus"
|
||||
)
|
||||
|
||||
// upstreamLogBodyMax caps the request body bytes copied into the
|
||||
// debug log line so a giant prompt or streamed payload doesn't fill the
|
||||
// log. The body itself is always restored to the request unchanged.
|
||||
const upstreamLogBodyMax = 4096
|
||||
|
||||
// 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.
|
||||
type MultiTransport struct {
|
||||
embedded http.RoundTripper
|
||||
direct *http.Transport
|
||||
insecure *http.Transport
|
||||
logger *log.Logger
|
||||
}
|
||||
|
||||
// NewMultiTransport wires both branches. embedded is the existing NetBird
|
||||
// roundtripper; the direct branches are constructed here with sensible
|
||||
// defaults that mirror Go's stdlib defaults plus a dial-timeout wrapper
|
||||
// honouring the per-request value attached via types.WithDialTimeout.
|
||||
// Pass embedded=nil to disable the WG branch entirely (every request
|
||||
// will route direct, regardless of the context flag). logger may be
|
||||
// nil; when nil the transport falls back to the logrus default
|
||||
// instance.
|
||||
func NewMultiTransport(embedded http.RoundTripper, logger *log.Logger) *MultiTransport {
|
||||
dialer := &net.Dialer{
|
||||
Timeout: 30 * time.Second,
|
||||
KeepAlive: 30 * time.Second,
|
||||
}
|
||||
direct := &http.Transport{
|
||||
DialContext: dialWithTimeout(dialer.DialContext),
|
||||
ForceAttemptHTTP2: true,
|
||||
MaxIdleConns: 100,
|
||||
IdleConnTimeout: 90 * time.Second,
|
||||
TLSHandshakeTimeout: 10 * time.Second,
|
||||
ExpectContinueTimeout: 1 * time.Second,
|
||||
}
|
||||
insecure := direct.Clone()
|
||||
insecure.TLSClientConfig = &tls.Config{InsecureSkipVerify: true} //nolint:gosec // matches the embedded NetBird transport's per-target opt-in
|
||||
|
||||
if logger == nil {
|
||||
logger = log.StandardLogger()
|
||||
}
|
||||
|
||||
return &MultiTransport{
|
||||
embedded: embedded,
|
||||
direct: direct,
|
||||
insecure: insecure,
|
||||
logger: logger,
|
||||
}
|
||||
}
|
||||
|
||||
// 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) {
|
||||
m.logUpstreamRequest(req)
|
||||
if DirectUpstreamFromContext(req.Context()) || m.embedded == nil {
|
||||
if skipTLSVerifyFromContext(req.Context()) {
|
||||
return m.insecure.RoundTrip(req)
|
||||
}
|
||||
return m.direct.RoundTrip(req)
|
||||
}
|
||||
return m.embedded.RoundTrip(req)
|
||||
}
|
||||
|
||||
// logUpstreamRequest emits the outbound request method, URL, headers,
|
||||
// and a (capped) body snippet at info level for debugging. The body is
|
||||
// read, copied into a snippet, and restored on the request so the
|
||||
// actual upstream call sees it unchanged.
|
||||
func (m *MultiTransport) logUpstreamRequest(req *http.Request) {
|
||||
if req == nil {
|
||||
return
|
||||
}
|
||||
body := snapshotRequestBody(req)
|
||||
m.logger.Debugf("upstream request: method=%s url=%s host=%s body_length=%d headers=%s body=%s",
|
||||
req.Method, req.URL.String(), req.Host, req.ContentLength, formatHeaders(req.Header), body)
|
||||
}
|
||||
|
||||
// formatHeaders renders the headers as a deterministic single-line
|
||||
// string. Multi-valued headers are joined with commas. Sensitive
|
||||
// header values (the upstream Authorization NetBird just stamped, plus
|
||||
// any cookie jar that survived) are redacted so logs don't leak the
|
||||
// provider API key.
|
||||
func formatHeaders(h http.Header) string {
|
||||
if len(h) == 0 {
|
||||
return "{}"
|
||||
}
|
||||
keys := make([]string, 0, len(h))
|
||||
for k := range h {
|
||||
keys = append(keys, k)
|
||||
}
|
||||
sort.Strings(keys)
|
||||
var sb strings.Builder
|
||||
sb.WriteByte('{')
|
||||
for i, k := range keys {
|
||||
if i > 0 {
|
||||
sb.WriteByte(' ')
|
||||
}
|
||||
sb.WriteString(k)
|
||||
sb.WriteByte('=')
|
||||
sb.WriteString(redactHeaderValue(k, h.Values(k)))
|
||||
}
|
||||
sb.WriteByte('}')
|
||||
return sb.String()
|
||||
}
|
||||
|
||||
// redactHeaderValue replaces sensitive credentials with a placeholder.
|
||||
// All other header values are joined with commas verbatim.
|
||||
func redactHeaderValue(name string, values []string) string {
|
||||
switch strings.ToLower(name) {
|
||||
case "authorization", "proxy-authorization", "x-api-key", "api-key", "cookie":
|
||||
return "[redacted]"
|
||||
}
|
||||
return strings.Join(values, ",")
|
||||
}
|
||||
|
||||
// snapshotRequestBody returns a printable snippet of the request body
|
||||
// (capped to upstreamLogBodyMax) and restores the body so downstream
|
||||
// transports can still read it. Returns the empty string when there's
|
||||
// no body or it can't be read.
|
||||
func snapshotRequestBody(req *http.Request) string {
|
||||
if req.Body == nil || req.Body == http.NoBody {
|
||||
return ""
|
||||
}
|
||||
raw, err := io.ReadAll(req.Body)
|
||||
if err != nil {
|
||||
return ""
|
||||
}
|
||||
req.Body = io.NopCloser(bytes.NewReader(raw))
|
||||
// Restore GetBody so transports performing redirects or retries
|
||||
// still get a fresh reader.
|
||||
req.GetBody = func() (io.ReadCloser, error) {
|
||||
return io.NopCloser(bytes.NewReader(raw)), nil
|
||||
}
|
||||
if len(raw) > upstreamLogBodyMax {
|
||||
return string(raw[:upstreamLogBodyMax]) + "...[truncated]"
|
||||
}
|
||||
return string(raw)
|
||||
}
|
||||
@@ -0,0 +1,79 @@
|
||||
package roundtrip
|
||||
|
||||
import (
|
||||
"context"
|
||||
"io"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"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")
|
||||
})
|
||||
}
|
||||
|
||||
func TestMultiTransport_NilEmbeddedAlwaysDirects(t *testing.T) {
|
||||
mt := NewMultiTransport(nil, nil)
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
|
||||
_, _ = io.WriteString(w, "ok")
|
||||
}))
|
||||
defer srv.Close()
|
||||
|
||||
req, err := http.NewRequest(http.MethodGet, srv.URL, nil)
|
||||
require.NoError(t, err)
|
||||
resp, err := mt.RoundTrip(req)
|
||||
require.NoError(t, err, "nil embedded must fall through to direct without panic")
|
||||
_ = resp.Body.Close()
|
||||
assert.Equal(t, http.StatusOK, resp.StatusCode)
|
||||
}
|
||||
@@ -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,11 +162,14 @@ type NetBird struct {
|
||||
clients map[types.AccountID]*clientEntry
|
||||
initLogOnce sync.Once
|
||||
statusNotifier statusNotifier
|
||||
|
||||
// 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
|
||||
// was reused). The duration covers keygen + gRPC CreateProxyPeer + embed.New.
|
||||
OnAddPeer func(d time.Duration, err error)
|
||||
// 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)
|
||||
}
|
||||
|
||||
// ClientDebugInfo contains debug information about a client.
|
||||
@@ -167,9 +190,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 +197,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 +217,15 @@ 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)
|
||||
if n.OnAddPeer != nil {
|
||||
n.OnAddPeer(time.Since(createStart), err)
|
||||
}
|
||||
entry, err := n.createClientEntry(ctx, accountID, key, authToken, si)
|
||||
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 +233,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,7 +297,7 @@ func (n *NetBird) createClientEntry(ctx context.Context, accountID types.Account
|
||||
ManagementURL: n.clientCfg.MgmtAddr,
|
||||
PrivateKey: privateKey.String(),
|
||||
LogLevel: log.WarnLevel.String(),
|
||||
BlockInbound: true,
|
||||
BlockInbound: n.clientCfg.BlockInbound,
|
||||
WireguardPort: &wgPort,
|
||||
PreSharedKey: n.clientCfg.PreSharedKey,
|
||||
})
|
||||
@@ -385,8 +364,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 +428,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 +450,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 +539,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 +597,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 +665,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 +719,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