diff --git a/client/ui/frontend/src/components/ReadySignal.tsx b/client/ui/frontend/src/components/ReadySignal.tsx new file mode 100644 index 000000000..0d040cabc --- /dev/null +++ b/client/ui/frontend/src/components/ReadySignal.tsx @@ -0,0 +1,18 @@ +import { useEffect, useRef } from "react"; +import { Events } from "@wailsio/runtime"; +import { useStatus } from "@/contexts/StatusContext.tsx"; + +const EVENT_WINDOW_PAINTED = "netbird:window-painted"; + +export const ReadySignal = () => { + const { isReady } = useStatus(); + const sent = useRef(false); + + useEffect(() => { + if (!isReady || sent.current) return; + sent.current = true; + void Events.Emit(EVENT_WINDOW_PAINTED); + }, [isReady]); + + return null; +}; diff --git a/client/ui/frontend/src/layouts/AppLayout.tsx b/client/ui/frontend/src/layouts/AppLayout.tsx index 1588d9d08..0c2837b53 100644 --- a/client/ui/frontend/src/layouts/AppLayout.tsx +++ b/client/ui/frontend/src/layouts/AppLayout.tsx @@ -5,6 +5,7 @@ import { DebugBundleProvider } from "@/contexts/DebugBundleContext.tsx"; import { ProfileProvider } from "@/contexts/ProfileContext.tsx"; import { DialogProvider } from "@/contexts/DialogContext.tsx"; import { RestrictionsProvider } from "@/contexts/RestrictionsContext.tsx"; +import { ReadySignal } from "@/components/ReadySignal.tsx"; export const AppLayout = () => { return ( @@ -16,6 +17,7 @@ export const AppLayout = () => { + diff --git a/client/ui/main.go b/client/ui/main.go index 5f740f5ec..e20bfe074 100644 --- a/client/ui/main.go +++ b/client/ui/main.go @@ -139,13 +139,11 @@ func main() { prefStore: prefStore, }) - window := newMainWindow(app, prefStore) - - // Settings is created eagerly (hidden) so the first gear click paints - // instantly and React keeps per-tab state across reopens. The other - // auxiliary windows stay lazy + destroy-on-close so Wails's macOS - // dock-reopen handler can't resurrect them. - windowManager := services.NewWindowManager(app, window, bundle, prefStore, iconWindow) + windowManager := services.NewWindowManager(app, nil, bundle, prefStore, iconWindow) + windowManager.SetMainFactory(func(startURL string) *application.WebviewWindow { + return newMainWindow(app, prefStore, windowManager, startURL) + }) + registerDockReopenHook(app, windowManager) // Minimal WMs (XEmbed-tray path) neither center small windows nor restore // position across hide -> show, dropping them top-left. Gate Go-side // re-centering on that environment; nil leaves placement to the WM on full @@ -168,7 +166,7 @@ func main() { // RegisterStatusNotifierItem hits a watcher we control. startStatusNotifierWatcher() - tray = NewTray(app, window, TrayServices{ + tray = NewTray(app, nil, TrayServices{ Connection: connection, Settings: settings, Profiles: profiles, @@ -279,10 +277,12 @@ func newApplication(onSecondInstance func()) *application.App { ActivationPolicy: application.ActivationPolicyAccessory, }, Linux: application.LinuxOptions{ - ProgramName: "netbird", + ProgramName: "netbird", + DisableQuitOnLastWindowClosed: true, }, Windows: application.WindowsOptions{ - WndProcInterceptor: endSessionInterceptor(), + WndProcInterceptor: endSessionInterceptor(), + DisableQuitOnLastWindowClosed: true, }, SingleInstance: &application.SingleInstanceOptions{ UniqueID: "io.netbird.ui", @@ -338,9 +338,7 @@ func registerServices(app *application.App, conn *Conn, s registeredServices) { app.RegisterService(application.NewService(s.compat)) } -// newMainWindow creates the hidden main window, sized to the user's last view -// mode, and installs the hide-on-close and macOS dock-reopen hooks. -func newMainWindow(app *application.App, prefStore *preferences.Store) *application.WebviewWindow { +func newMainWindow(app *application.App, prefStore *preferences.Store, wm *services.WindowManager, startURL string) *application.WebviewWindow { // Width matches the last view mode so Advanced-mode users don't see the // window pop from 380px to 900px on launch. Height is mode-agnostic. initialWidth := 380 @@ -357,7 +355,7 @@ func newMainWindow(app *application.App, prefStore *preferences.Store) *applicat InitialPosition: application.WindowCentered, Hidden: true, BackgroundColour: services.WindowBackgroundColour, - URL: "/", + URL: startURL, DisableResize: true, MinimiseButtonState: application.ButtonHidden, MaximiseButtonState: application.ButtonHidden, @@ -368,29 +366,25 @@ func newMainWindow(app *application.App, prefStore *preferences.Store) *applicat }, }) - // Hide instead of quit on close; "really quit" is reached via tray -> Quit. - window.RegisterHook(events.Common.WindowClosing, func(e *application.WindowEvent) { + window.RegisterHook(events.Common.WindowClosing, func(_ *application.WindowEvent) { if services.ShuttingDown() { return } - e.Cancel() - window.Hide() + wm.ForgetMain() }) - // On macOS, Wails' default applicationShouldHandleReopen handler Show()s - // every hidden window on dock-icon click, resurrecting hide-on-close - // surfaces like Settings. Cancel it in a hook (hooks run before listeners) - // and show only the main window. No-op elsewhere — the event never fires. - if runtime.GOOS == "darwin" { - app.Event.RegisterApplicationEventHook(events.Mac.ApplicationShouldHandleReopen, func(e *application.ApplicationEvent) { - e.Cancel() - if e.Context().HasVisibleWindows() { - return - } - window.Show() - window.Focus() - }) - } - return window } + +func registerDockReopenHook(app *application.App, wm *services.WindowManager) { + if runtime.GOOS != "darwin" { + return + } + app.Event.RegisterApplicationEventHook(events.Mac.ApplicationShouldHandleReopen, func(e *application.ApplicationEvent) { + if e.Context().HasVisibleWindows() { + return + } + e.Cancel() + wm.ShowMain() + }) +} diff --git a/client/ui/services/windowmanager.go b/client/ui/services/windowmanager.go index 5f7aaa7bd..4930ce22b 100644 --- a/client/ui/services/windowmanager.go +++ b/client/ui/services/windowmanager.go @@ -8,6 +8,7 @@ import ( "sync" "time" + log "github.com/sirupsen/logrus" "github.com/wailsapp/wails/v3/pkg/application" "github.com/wailsapp/wails/v3/pkg/events" @@ -29,6 +30,12 @@ const EventBrowserLoginCancel = "browser-login:cancel" // EventSettingsOpen tells the mounted settings window which tab to show. const EventSettingsOpen = "netbird:settings:open" +const EventWindowPainted = "netbird:window-painted" + +const paintedFallback = 2 * time.Second + +const headlessTeardownDelay = 2 * time.Second + var WindowBackgroundColour = application.NewRGB(24, 26, 29) // bg-nb-gray-950 // WindowHeight is shared by the main and Settings windows. @@ -94,9 +101,6 @@ func DialogWindowOptions(name, title, url string, linuxIcon []byte) application. } } -// WindowManager owns the auxiliary windows (main is created in main.go). Settings is created -// eagerly and hidden on close to keep React state; the rest are created on open, destroyed on -// close, so the macOS dock-reopen handler finds no hidden window to resurrect. type WindowManager struct { app *application.App mainWindow *application.WebviewWindow @@ -112,15 +116,35 @@ type WindowManager struct { // hiddenForLogin holds windows hidden while the BrowserLogin popup is open, restored on close. hiddenForLogin []application.Window mu sync.Mutex + createMu sync.Mutex + newMain func(startURL string) *application.WebviewWindow + ready map[uint]bool + showPending map[uint]bool + pendingTab map[uint]string + pendingEmits map[uint][]string + fallbackTimers map[uint]*time.Timer + headlessMain bool + headlessTimer *time.Timer // recenterOnShow is set only on the minimal-WM/XEmbed path, where the WM neither centers nor // restores position; nil on full desktops so re-centering can't fight a user-moved window. recenterOnShow func() bool } -// NewWindowManager wires the manager to the main app; translator/prefs may be nil (tests). The -// Settings window is created here (hidden) so the first OpenSettings is instant. func NewWindowManager(app *application.App, mainWindow *application.WebviewWindow, translator ErrorTranslator, prefs LanguagePreference, linuxIcon []byte) *WindowManager { - s := &WindowManager{app: app, mainWindow: mainWindow, translator: translator, prefs: prefs, linuxIcon: linuxIcon} + s := &WindowManager{ + app: app, + mainWindow: mainWindow, + translator: translator, + prefs: prefs, + linuxIcon: linuxIcon, + ready: map[uint]bool{}, + showPending: map[uint]bool{}, + pendingTab: map[uint]string{}, + pendingEmits: map[uint][]string{}, + fallbackTimers: map[uint]*time.Timer{}, + } + s.watchPainted() + s.watchTriggerLogin() // Re-title live windows on language flip. Wired internally so the binding generator // doesn't try to expose the interface param. if sub, ok := prefs.(LanguageSubscriber); ok && sub != nil { @@ -136,7 +160,11 @@ func NewWindowManager(app *application.App, mainWindow *application.WebviewWindo } }() } - s.settings = app.Window.NewWithOptions(application.WebviewWindowOptions{ + return s +} + +func (s *WindowManager) newSettingsWindow() *application.WebviewWindow { + w := s.app.Window.NewWithOptions(application.WebviewWindowOptions{ Name: "settings", Title: s.title("window.title.settings"), Width: 900, @@ -150,18 +178,15 @@ func NewWindowManager(app *application.App, mainWindow *application.WebviewWindo URL: "/#/settings", Mac: AppleMacOSAppearanceOptions(), Windows: MicrosoftWindowsAppearanceOptions(), - Linux: LinuxAppearanceOptions(linuxIcon), + Linux: LinuxAppearanceOptions(s.linuxIcon), }) - // Hide (not destroy) on close to keep React state; reset to General for a flash-free reopen. - s.settings.RegisterHook(events.Common.WindowClosing, func(e *application.WindowEvent) { - if ShuttingDown() { - return - } - e.Cancel() - s.app.Event.Emit(EventSettingsOpen, "general") - s.settings.Hide() + w.RegisterHook(events.Common.WindowClosing, func(_ *application.WindowEvent) { + s.mu.Lock() + s.settings = nil + s.forgetWindowLocked(w) + s.mu.Unlock() }) - return s + return w } // OpenSettings shows the settings window on tab (empty → General), switching tab via @@ -171,11 +196,20 @@ func (s *WindowManager) OpenSettings(tab string) { if target == "" { target = "general" } - s.app.Event.Emit(EventSettingsOpen, target) - s.settings.Show() - s.settings.Focus() - // Re-center (minimal-WM only; see centerWhenReady). - s.centerWhenReady(s.settings) + + w, _ := s.ensureWindow(&s.settings, s.newSettingsWindow) + + s.mu.Lock() + ready := s.ready[w.ID()] + if !ready { + s.pendingTab[w.ID()] = target + } + s.mu.Unlock() + + if ready { + s.app.Event.Emit(EventSettingsOpen, target) + } + s.showWhenReady(w) } // OpenBrowserLogin shows the SSO popup, creating it on first use. @@ -440,13 +474,295 @@ func (s *WindowManager) OpenMain() { // ShowMain brings the main window forward (re-centering on minimal WMs). The single entry // point every surface (tray, SIGUSR1, welcome) should use so centering applies uniformly. func (s *WindowManager) ShowMain() { - if s.mainWindow == nil { + s.showWhenReady(s.MainWindow()) +} + +// ShowMainAndEmit brings the main window forward and emits event once its frontend is ready. +func (s *WindowManager) ShowMainAndEmit(event string) { + w := s.MainWindow() + if w == nil { return } - s.mainWindow.Show() - s.mainWindow.Focus() - // Re-center (minimal-WM only; see centerWhenReady). - s.centerWhenReady(s.mainWindow) + + id := w.ID() + s.mu.Lock() + ready := s.ready[id] + if !ready { + s.pendingEmits[id] = append(s.pendingEmits[id], event) + } + s.mu.Unlock() + + s.showWhenReady(w) + if ready { + s.app.Event.Emit(event) + } +} + +func (s *WindowManager) MainWindow() *application.WebviewWindow { + w, _ := s.ensureMain("/") + return w +} + +func (s *WindowManager) ensureMain(startURL string) (*application.WebviewWindow, bool) { + s.mu.Lock() + factory := s.newMain + s.mu.Unlock() + if factory == nil { + return s.ensureWindow(&s.mainWindow, nil) + } + return s.ensureWindow(&s.mainWindow, func() *application.WebviewWindow { + return factory(startURL) + }) +} + +func (s *WindowManager) ensureWindow(slot **application.WebviewWindow, factory func() *application.WebviewWindow) (*application.WebviewWindow, bool) { + s.createMu.Lock() + defer s.createMu.Unlock() + + s.mu.Lock() + w := *slot + s.mu.Unlock() + if w != nil || factory == nil { + return w, false + } + + w = factory() + s.armReady(w) + + s.mu.Lock() + *slot = w + s.mu.Unlock() + return w, true +} + +func (s *WindowManager) armReady(w *application.WebviewWindow) { + if w == nil { + return + } + w.RegisterHook(events.Common.WindowRuntimeReady, func(_ *application.WindowEvent) { + timer := time.AfterFunc(paintedFallback, func() { + log.Warnf("window %q never reported a first render, showing it anyway", w.Name()) + s.markReady(w) + }) + s.mu.Lock() + s.fallbackTimers[w.ID()] = timer + s.mu.Unlock() + }) +} + +func (s *WindowManager) watchPainted() { + s.app.Event.On(EventWindowPainted, func(e *application.CustomEvent) { + if w := s.windowByName(e.Sender); w != nil { + s.markReady(w) + } + }) +} + +func (s *WindowManager) watchTriggerLogin() { + s.app.Event.On(EventTriggerLogin, func(_ *application.CustomEvent) { + s.mu.Lock() + if s.headlessTimer != nil { + s.headlessTimer.Stop() + s.headlessTimer = nil + } + w := s.mainWindow + ready := w != nil && s.ready[w.ID()] + s.mu.Unlock() + if ready { + return + } + + w, created := s.ensureMain("/") + if w == nil { + return + } + + s.mu.Lock() + if created { + s.headlessMain = true + } + pending := !s.ready[w.ID()] + if pending { + s.pendingEmits[w.ID()] = append(s.pendingEmits[w.ID()], EventTriggerLogin) + } + s.mu.Unlock() + + if !pending { + s.app.Event.Emit(EventTriggerLogin) + } + }) + + s.app.Event.On(EventBrowserLoginCancel, func(_ *application.CustomEvent) { + s.scheduleHeadlessTeardown() + }) + + s.app.Event.On(EventStatusSnapshot, func(e *application.CustomEvent) { + st, ok := e.Data.(Status) + if !ok { + return + } + switch st.Status { + case StatusConnected, StatusLoginFailed, StatusDaemonUnavailable: + s.scheduleHeadlessTeardown() + } + }) +} + +func (s *WindowManager) scheduleHeadlessTeardown() { + s.mu.Lock() + defer s.mu.Unlock() + if !s.headlessMain || s.mainWindow == nil { + return + } + if s.headlessTimer != nil { + s.headlessTimer.Stop() + } + s.headlessTimer = time.AfterFunc(headlessTeardownDelay, s.closeHeadlessMain) +} + +func (s *WindowManager) closeHeadlessMain() { + s.mu.Lock() + w := s.mainWindow + headless := s.headlessMain + s.headlessTimer = nil + s.mu.Unlock() + if !headless || w == nil { + return + } + w.Close() +} + +func (s *WindowManager) forgetWindowLocked(w *application.WebviewWindow) { + if w == nil { + return + } + + id := w.ID() + if timer := s.fallbackTimers[id]; timer != nil { + timer.Stop() + } + delete(s.fallbackTimers, id) + delete(s.ready, id) + delete(s.showPending, id) + delete(s.pendingTab, id) + delete(s.pendingEmits, id) + + kept := s.hiddenForLogin[:0] + for _, hidden := range s.hiddenForLogin { + if hidden != application.Window(w) { + kept = append(kept, hidden) + } + } + s.hiddenForLogin = kept +} + +func (s *WindowManager) windowByName(name string) *application.WebviewWindow { + s.mu.Lock() + defer s.mu.Unlock() + switch name { + case "main": + return s.mainWindow + case "settings": + return s.settings + default: + return nil + } +} + +func (s *WindowManager) markReady(w *application.WebviewWindow) { + id := w.ID() + s.mu.Lock() + already := s.ready[id] + s.ready[id] = true + wanted := s.showPending[id] + tab, hasTab := s.pendingTab[id] + emits := s.pendingEmits[id] + if timer := s.fallbackTimers[id]; timer != nil { + timer.Stop() + delete(s.fallbackTimers, id) + } + delete(s.showPending, id) + delete(s.pendingTab, id) + delete(s.pendingEmits, id) + s.mu.Unlock() + + if already { + return + } + + if hasTab { + s.app.Event.Emit(EventSettingsOpen, tab) + } + + if wanted { + s.showNow(w) + } + + for _, event := range emits { + s.app.Event.Emit(event) + } +} + +func (s *WindowManager) showWhenReady(w *application.WebviewWindow) { + if w == nil { + return + } + + id := w.ID() + s.mu.Lock() + ready := s.ready[id] + if !ready { + s.showPending[id] = true + } + s.mu.Unlock() + + if ready { + s.showNow(w) + } +} + +func (s *WindowManager) showNow(w *application.WebviewWindow) { + s.mu.Lock() + if w == s.mainWindow { + s.headlessMain = false + if s.headlessTimer != nil { + s.headlessTimer.Stop() + s.headlessTimer = nil + } + } + s.mu.Unlock() + w.Show() + w.Focus() + s.centerWhenReady(w) +} + +func (s *WindowManager) ShowMainAt(url string) { + w, created := s.ensureMain(url) + if w == nil { + return + } + if !created { + w.SetURL(url) + } + s.showWhenReady(w) +} + +func (s *WindowManager) SetMainFactory(f func(startURL string) *application.WebviewWindow) { + s.mu.Lock() + defer s.mu.Unlock() + s.newMain = f +} + +func (s *WindowManager) ForgetMain() { + s.mu.Lock() + defer s.mu.Unlock() + s.forgetWindowLocked(s.mainWindow) + s.mainWindow = nil + s.headlessMain = false + if s.headlessTimer != nil { + s.headlessTimer.Stop() + s.headlessTimer = nil + } } // SetRecenterOnShow installs the recenterOnShow predicate (see the field). diff --git a/client/ui/tray.go b/client/ui/tray.go index 148dd50b3..c392a0b62 100644 --- a/client/ui/tray.go +++ b/client/ui/tray.go @@ -174,7 +174,7 @@ func NewTray(app *application.App, window *application.WebviewWindow, svc TraySe // in the right locale — no English flash then re-paint. loc: svc.Localizer, } - t.updater = newTrayUpdater(app, window, svc.Update, svc.Notifier, t.loc, func() { t.applyIcon() }, func() { t.relayoutMenu() }) + t.updater = newTrayUpdater(app, t.showMainAt, svc.Update, svc.Notifier, t.loc, func() { t.applyIcon() }, func() { t.relayoutMenu() }) t.tray = app.SystemTray.New() // Seed panel-theme detection before the first paint so the initial icon // matches the panel's light/dark scheme (Linux only). @@ -241,9 +241,6 @@ func (t *Tray) ShowWindow() { w.Focus() return } - if t.window == nil { - return - } // Route through WindowManager so the main window is centered on first // show — minimal WMs (fluxbox, the XEmbed tray path) otherwise drop it in // the top-left corner. @@ -251,8 +248,49 @@ func (t *Tray) ShowWindow() { t.svc.WindowManager.ShowMain() return } - t.window.Show() - t.window.Focus() + if w := t.mainWindow(); w != nil { + w.Show() + w.Focus() + } +} + +func (t *Tray) mainWindow() *application.WebviewWindow { + if t.svc.WindowManager == nil { + return t.window + } + return t.svc.WindowManager.MainWindow() +} + +func (t *Tray) showMainAt(url string) { + if t.svc.WindowManager != nil { + t.svc.WindowManager.ShowMainAt(url) + return + } + if w := t.mainWindow(); w != nil { + w.SetURL(url) + w.Show() + w.Focus() + } +} + +func (t *Tray) showMain() { + if t.svc.WindowManager != nil { + t.svc.WindowManager.ShowMain() + return + } + if w := t.mainWindow(); w != nil { + w.Show() + w.Focus() + } +} + +func (t *Tray) showMainAndEmit(event string) { + if t.svc.WindowManager != nil { + t.svc.WindowManager.ShowMainAndEmit(event) + return + } + t.showMain() + t.app.Event.Emit(event) } // applyLanguage re-renders every translated surface in the Localizer's current @@ -479,7 +517,8 @@ func (t *Tray) handleConnect(upItem *application.MenuItem) { // NeedsLogin/SessionExpired/LoginFailed won't honor a plain Up RPC — they // need the Login → WaitSSOLogin → Up sequence. Emit EventTriggerLogin so // the React startLogin() (which owns the BrowserLogin popup) drives it; - // the hidden main webview is alive and subscribed, so only the popup shows. + // the WindowManager materialises a hidden main webview when none is live, + // so only the popup shows. t.statusMu.Lock() needsLogin := strings.EqualFold(t.lastStatus, services.StatusNeedsLogin) || strings.EqualFold(t.lastStatus, services.StatusSessionExpired) || diff --git a/client/ui/tray_session.go b/client/ui/tray_session.go index f25419894..6e5d07740 100644 --- a/client/ui/tray_session.go +++ b/client/ui/tray_session.go @@ -30,10 +30,7 @@ const ( // handleSessionExpired notifies and brings the window forward so the user can reconnect. func (t *Tray) handleSessionExpired() { t.notify(t.loc.T("notify.sessionExpired.title"), t.loc.T("notify.sessionExpired.body"), notifyIDSessionExpired) - if t.window != nil { - t.window.Show() - t.window.Focus() - } + t.showMain() } // applySessionExpiry refreshes the cached SSO deadline and reports whether it changed. @@ -307,7 +304,7 @@ func (t *Tray) openSessionExtendFlow() { } seconds := int(time.Until(deadline).Seconds()) if seconds <= 0 { - t.app.Event.Emit(services.EventTriggerLogin) + t.showMainAndEmit(services.EventTriggerLogin) return } if t.svc.WindowManager == nil { diff --git a/client/ui/tray_update.go b/client/ui/tray_update.go index 27037eccb..3ce1f9600 100644 --- a/client/ui/tray_update.go +++ b/client/ui/tray_update.go @@ -4,6 +4,7 @@ package main import ( "context" + neturl "net/url" "sync" "time" @@ -19,7 +20,7 @@ import ( // trayUpdater owns the tray UI that reacts to auto-update. Composed inside Tray. type trayUpdater struct { app *application.App - window *application.WebviewWindow + showMainAt func(url string) update *services.Update notifier *Notifier loc *Localizer @@ -36,10 +37,10 @@ type trayUpdater struct { progressWindowOpen bool } -func newTrayUpdater(app *application.App, window *application.WebviewWindow, update *services.Update, notifier *Notifier, loc *Localizer, onIconChange func(), onMenuChange func()) *trayUpdater { +func newTrayUpdater(app *application.App, showMainAt func(url string), update *services.Update, notifier *Notifier, loc *Localizer, onIconChange func(), onMenuChange func()) *trayUpdater { u := &trayUpdater{ app: app, - window: window, + showMainAt: showMainAt, update: update, notifier: notifier, loc: loc, @@ -185,14 +186,12 @@ func (u *trayUpdater) sendUpdateNotification(st updater.State) { // openProgressWindow points the main window at the /update progress page and // brings it forward. func (u *trayUpdater) openProgressWindow(version string) { - if u.window == nil { + if u.showMainAt == nil { return } url := "/#/update" if version != "" { - url += "?version=" + version + url += "?version=" + neturl.QueryEscape(version) } - u.window.SetURL(url) - u.window.Show() - u.window.Focus() + u.showMainAt(url) } diff --git a/go.mod b/go.mod index ec1557096..d8f11cc7b 100644 --- a/go.mod +++ b/go.mod @@ -346,6 +346,6 @@ replace github.com/dexidp/dex/api/v2 => github.com/netbirdio/dex/api/v2 v2.0.0-2 replace github.com/mailru/easyjson => github.com/netbirdio/easyjson v0.9.0 -replace github.com/wailsapp/wails/v3 => github.com/netbirdio/wails/v3 v3.0.0-beta.3.0.20260810103952-24e716aea4db +replace github.com/wailsapp/wails/v3 => github.com/netbirdio/wails/v3 v3.0.0-beta.3.0.20260825085513-5f07a01f7a78 tool go.uber.org/mock/mockgen diff --git a/go.sum b/go.sum index 4db0c0806..2365bbb7b 100644 --- a/go.sum +++ b/go.sum @@ -488,8 +488,8 @@ github.com/netbirdio/service v0.0.0-20240911161631-f62744f42502 h1:3tHlFmhTdX9ax github.com/netbirdio/service v0.0.0-20240911161631-f62744f42502/go.mod h1:CIMRFEJVL+0DS1a3Nx06NaMn4Dz63Ng6O7dl0qH0zVM= github.com/netbirdio/signal-dispatcher/dispatcher v0.0.0-20250805121659-6b4ac470ca45 h1:ujgviVYmx243Ksy7NdSwrdGPSRNE3pb8kEDSpH0QuAQ= github.com/netbirdio/signal-dispatcher/dispatcher v0.0.0-20250805121659-6b4ac470ca45/go.mod h1:5/sjFmLb8O96B5737VCqhHyGRzNFIaN/Bu7ZodXc3qQ= -github.com/netbirdio/wails/v3 v3.0.0-beta.3.0.20260810103952-24e716aea4db h1:gBOE2r4AW1soSmpYJC5/n9/1L8UQ8+HLjed8CY/TzZY= -github.com/netbirdio/wails/v3 v3.0.0-beta.3.0.20260810103952-24e716aea4db/go.mod h1:bsdahLwBQxXjlmdPPeQyrTcDJfcqAr/ymFj0RXhwtWI= +github.com/netbirdio/wails/v3 v3.0.0-beta.3.0.20260825085513-5f07a01f7a78 h1:B/jRv24jnFeoA+VccxoCx6K94PUgsqR9wnshpeu9M+8= +github.com/netbirdio/wails/v3 v3.0.0-beta.3.0.20260825085513-5f07a01f7a78/go.mod h1:/6QR46/nhGCSADHbS++XtDb9dkTnenTHlGskTPRo9S0= github.com/netbirdio/wireguard-go v0.0.0-20260628102922-2834bebf6c1a h1:3CWK+yTvRKOcC0Q8VCTGy4l60TEb27CQVS7LkMxwjmw= github.com/netbirdio/wireguard-go v0.0.0-20260628102922-2834bebf6c1a/go.mod h1:rpwXGsirqLqN2L0JDJQlwOboGHmptD5ZD6T2VmcqhTw= github.com/nxadm/tail v1.4.4/go.mod h1:kenIhsEOeOJmVchQTgglprH7qJGnHDVpk1VPCcaMI8A= diff --git a/management/internals/shared/grpc/proxy.go b/management/internals/shared/grpc/proxy.go index 40b0914ef..cee50b270 100644 --- a/management/internals/shared/grpc/proxy.go +++ b/management/internals/shared/grpc/proxy.go @@ -1311,7 +1311,7 @@ func (s *ProxyServiceServer) authenticateHeader(ctx context.Context, serviceID s lastErr = err continue } - return true, "header-user", proxyauth.MethodHeader + return true, proxyauth.HeaderUserID, proxyauth.MethodHeader } if lastErr != nil { diff --git a/proxy/auth/auth.go b/proxy/auth/auth.go index 78f0097d5..5512bf003 100644 --- a/proxy/auth/auth.go +++ b/proxy/auth/auth.go @@ -30,6 +30,12 @@ const ( SessionJWTIssuer = "netbird-management" ) +// HeaderUserID is the synthetic user id recorded for header-authenticated +// requests. Header auth validates a per-service secret and resolves no user +// record, so proxy access logs and management-minted session tokens both +// attribute the request to this id. +const HeaderUserID = "header-user" + // ResolveProto determines the protocol scheme based on the forwarded proto // configuration. When set to "http" or "https" the value is used directly. // Otherwise TLS state is used: if conn is non-nil "https" is returned, else "http". diff --git a/proxy/internal/auth/header.go b/proxy/internal/auth/header.go index 194800a49..64d5da8f1 100644 --- a/proxy/internal/auth/header.go +++ b/proxy/internal/auth/header.go @@ -1,36 +1,33 @@ package auth import ( + "crypto/sha256" "errors" - "fmt" "net/http" + "sync" "github.com/netbirdio/netbird/proxy/auth" - "github.com/netbirdio/netbird/proxy/internal/types" - "github.com/netbirdio/netbird/shared/management/proto" + "github.com/netbirdio/netbird/shared/hash/argon2id" ) -// ErrHeaderAuthFailed indicates that the header was present but the -// credential did not validate. Callers should return 401 instead of -// falling through to other auth schemes. -var ErrHeaderAuthFailed = errors.New("header authentication failed") - -// Header implements header-based authentication. The proxy checks for the -// configured header in each request and validates its value via gRPC. +// Header implements header-based authentication. The service mapping carries +// the argon2id hash of every value accepted for the header, so the proxy +// verifies the credential locally rather than round-tripping to management. type Header struct { - id types.ServiceID - accountId types.AccountID headerName string - client authenticator + hashes []string + verified *verifiedValues } -// NewHeader creates a Header authentication scheme for the given header name. -func NewHeader(client authenticator, id types.ServiceID, accountId types.AccountID, headerName string) Header { +// NewHeader creates a Header authentication scheme accepting any value whose +// argon2id hash appears in hashes. An empty hashes slice rejects every request +// carrying the header, so a mapping that arrived without its hashes fails +// closed instead of leaving the service unprotected. +func NewHeader(headerName string, hashes []string) Header { return Header{ - id: id, - accountId: accountId, - headerName: headerName, - client: client, + headerName: http.CanonicalHeaderKey(headerName), + hashes: hashes, + verified: &verifiedValues{seen: make(map[[32]byte]struct{}, len(hashes))}, } } @@ -39,31 +36,64 @@ func (Header) Type() auth.Method { return auth.MethodHeader } -// Authenticate checks for the configured header in the request. If absent, -// returns empty (unauthenticated). If present, validates via gRPC. -func (h Header) Authenticate(r *http.Request) (string, string, error) { +// Authenticate satisfies Scheme. Header credentials are resolved by Verify +// before the scheme loop runs, so a request that reaches here never carries +// the header and there is no credential to prompt for. +func (Header) Authenticate(*http.Request) (string, string, error) { + return "", "", nil +} + +// Verify reports whether the request carries the configured header and, when +// it does, whether the value matches one of the service's hashes. +// +// A non-nil unusable is a diagnostic rather than a request error: a stored hash +// could not be decoded, so no credential can ever match it and the header stays +// unauthenticatable until the service is saved again. Folding that into an +// ordinary mismatch would hide the misconfiguration behind a permanent 401. +func (h Header) Verify(r *http.Request) (present, matched bool, unusable error) { value := r.Header.Get(h.headerName) if value == "" { - return "", "", nil + return false, false, nil } - res, err := h.client.Authenticate(r.Context(), &proto.AuthenticateRequest{ - Id: string(h.id), - AccountId: string(h.accountId), - Request: &proto.AuthenticateRequest_HeaderAuth{ - HeaderAuth: &proto.HeaderAuthRequest{ - HeaderValue: value, - HeaderName: h.headerName, - }, - }, - }) - if err != nil { - return "", "", fmt.Errorf("authenticate header: %w", err) + digest := sha256.Sum256([]byte(value)) + if h.verified.has(digest) { + return true, true, nil } - if res.GetSuccess() { - return res.GetSessionToken(), "", nil + for _, hash := range h.hashes { + err := argon2id.Verify(value, hash) + if err == nil { + h.verified.add(digest) + return true, true, nil + } + if !errors.Is(err, argon2id.ErrMismatchedHashAndPassword) { + unusable = err + } } - - return "", "", ErrHeaderAuthFailed + return true, false, unusable +} + +// verifiedValues remembers which header values already passed argon2id +// verification. argon2id is deliberately expensive (19 MiB, two passes) and +// header credentials repeat on every request, so re-deriving per request would +// dominate the hot path. The set cannot outgrow the number of configured +// hashes, and a mapping update builds a fresh scheme with an empty set. +// Values are keyed by digest so the plaintext credential is not retained. +type verifiedValues struct { + mu sync.Mutex + seen map[[32]byte]struct{} +} + +func (v *verifiedValues) has(digest [32]byte) bool { + v.mu.Lock() + defer v.mu.Unlock() + _, ok := v.seen[digest] + return ok +} + +func (v *verifiedValues) add(digest [32]byte) { + v.mu.Lock() + defer v.mu.Unlock() + v.seen[digest] = struct{}{} } diff --git a/proxy/internal/auth/middleware.go b/proxy/internal/auth/middleware.go index 72630b085..8abdf2923 100644 --- a/proxy/internal/auth/middleware.go +++ b/proxy/internal/auth/middleware.go @@ -146,7 +146,7 @@ func (mw *Middleware) Protect(next http.Handler) http.Handler { return } - if mw.forwardWithHeaderAuth(w, r, host, config, next) { + if mw.forwardWithHeaderAuth(w, r, config, next) { return } @@ -325,6 +325,16 @@ func (mw *Middleware) forwardWithSessionCookie(w http.ResponseWriter, r *http.Re if err != nil { return false } + + // Header auth is checked per request against the mapping's hashes and mints + // no session, so a header-method token can only predate that. Honouring it + // would keep a rotated credential working until the token expired. + if method == auth.MethodHeader.String() { + mw.logger.WithField("host", host). + Debug("ignoring header-auth session cookie; the header is required on every request") + return false + } + if cd := proxy.CapturedDataFromContext(r.Context()); cd != nil { cd.SetUserID(userID) cd.SetUserEmail(email) @@ -436,73 +446,44 @@ func isTunnelSourceIP(ip netip.Addr) bool { // forwardWithHeaderAuth checks for a Header auth scheme. If the header validates, // the request is forwarded directly (no redirect), which is important for API clients. -func (mw *Middleware) forwardWithHeaderAuth(w http.ResponseWriter, r *http.Request, host string, config DomainConfig, next http.Handler) bool { +func (mw *Middleware) forwardWithHeaderAuth(w http.ResponseWriter, r *http.Request, config DomainConfig, next http.Handler) bool { + var presented []string for _, scheme := range config.Schemes { hdr, ok := scheme.(Header) if !ok { continue } - handled := mw.tryHeaderScheme(w, r, host, config, hdr, next) - if handled { + present, matched, unusable := hdr.Verify(r) + if matched { + if cd := proxy.CapturedDataFromContext(r.Context()); cd != nil { + cd.SetUserID(auth.HeaderUserID) + cd.SetAuthMethod(auth.MethodHeader.String()) + } + next.ServeHTTP(w, r) return true } + if unusable != nil { + mw.logger.WithFields(log.Fields{ + "host": r.Host, + "header": hdr.headerName, + }).WithError(unusable).Error("header auth: a configured hash cannot be decoded, so this header can never authenticate; re-save the service") + } + if present { + presented = append(presented, hdr.headerName) + } } - return false -} -func (mw *Middleware) tryHeaderScheme(w http.ResponseWriter, r *http.Request, host string, config DomainConfig, hdr Header, next http.Handler) bool { - token, _, err := hdr.Authenticate(r) - if err != nil { - return mw.handleHeaderAuthError(w, r, err) - } - if token == "" { + if len(presented) == 0 { return false } - result, err := mw.validateSessionToken(r.Context(), host, token, config.SessionPublicKey, auth.MethodHeader) - if err != nil { - setHeaderCapturedData(r.Context(), "", "", nil, nil) - status := http.StatusBadRequest - msg := "invalid session token" - if errors.Is(err, errValidationUnavailable) { - status = http.StatusBadGateway - msg = "authentication service unavailable" - } - http.Error(w, msg, status) - return true - } - - if !result.Valid { - setHeaderCapturedData(r.Context(), result.UserID, result.UserEmail, result.Groups, result.GroupNames) - http.Error(w, "Unauthorized", http.StatusUnauthorized) - return true - } - - setSessionCookie(w, token, config.SessionExpiration) - if cd := proxy.CapturedDataFromContext(r.Context()); cd != nil { - cd.SetUserID(result.UserID) - cd.SetUserEmail(result.UserEmail) - cd.SetUserGroups(result.Groups) - cd.SetUserGroupNames(result.GroupNames) - cd.SetAuthMethod(auth.MethodHeader.String()) - } - - next.ServeHTTP(w, r) - return true -} - -func (mw *Middleware) handleHeaderAuthError(w http.ResponseWriter, r *http.Request, err error) bool { - if errors.Is(err, ErrHeaderAuthFailed) { - setHeaderCapturedData(r.Context(), "", "", nil, nil) - http.Error(w, "Unauthorized", http.StatusUnauthorized) - return true - } - mw.logger.WithField("scheme", "header").Warnf("header auth infrastructure error: %v", err) - if cd := proxy.CapturedDataFromContext(r.Context()); cd != nil { - cd.SetOrigin(proxy.OriginAuth) - } - http.Error(w, "authentication service unavailable", http.StatusBadGateway) + mw.logger.WithFields(log.Fields{ + "host": r.Host, + "headers": presented, + }).Debug("header auth rejected: no presented header matched a configured hash") + setHeaderCapturedData(r.Context(), "", "", nil, nil) + http.Error(w, "Unauthorized", http.StatusUnauthorized) return true } diff --git a/proxy/internal/auth/middleware_test.go b/proxy/internal/auth/middleware_test.go index 6608c2b22..9220ce790 100644 --- a/proxy/internal/auth/middleware_test.go +++ b/proxy/internal/auth/middleware_test.go @@ -16,6 +16,7 @@ import ( "time" log "github.com/sirupsen/logrus" + logtest "github.com/sirupsen/logrus/hooks/test" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" "google.golang.org/grpc" @@ -25,6 +26,7 @@ import ( "github.com/netbirdio/netbird/proxy/internal/proxy" "github.com/netbirdio/netbird/proxy/internal/restrict" "github.com/netbirdio/netbird/proxy/internal/types" + "github.com/netbirdio/netbird/shared/hash/argon2id" "github.com/netbirdio/netbird/shared/management/proto" ) @@ -1023,38 +1025,24 @@ func TestProtect_OIDCWithOtherMethodShowsLoginPage(t *testing.T) { assert.Equal(t, http.StatusUnauthorized, rec.Code, "should show login page when multiple methods exist") } -// mockAuthenticator is a minimal mock for the authenticator gRPC interface -// used by the Header scheme. -type mockAuthenticator struct { - fn func(ctx context.Context, req *proto.AuthenticateRequest) (*proto.AuthenticateResponse, error) -} - -func (m *mockAuthenticator) Authenticate(ctx context.Context, in *proto.AuthenticateRequest, _ ...grpc.CallOption) (*proto.AuthenticateResponse, error) { - return m.fn(ctx, in) -} - -// newHeaderSchemeWithToken creates a Header scheme backed by a mock that -// returns a signed session token when the expected header value is provided. -func newHeaderSchemeWithToken(t *testing.T, kp *sessionkey.KeyPair, headerName, expectedValue string) Header { +// newHeaderScheme creates a Header scheme accepting each of the given values, +// hashed the way management hashes them before putting them on the mapping. +func newHeaderScheme(t *testing.T, headerName string, acceptedValues ...string) Header { t.Helper() - token, err := sessionkey.SignToken(kp.PrivateKey, "header-user", "", "example.com", auth.MethodHeader, nil, nil, time.Hour) - require.NoError(t, err) - - mock := &mockAuthenticator{fn: func(_ context.Context, req *proto.AuthenticateRequest) (*proto.AuthenticateResponse, error) { - ha := req.GetHeaderAuth() - if ha != nil && ha.GetHeaderValue() == expectedValue { - return &proto.AuthenticateResponse{Success: true, SessionToken: token}, nil - } - return &proto.AuthenticateResponse{Success: false}, nil - }} - return NewHeader(mock, "svc1", "acc1", headerName) + hashes := make([]string, 0, len(acceptedValues)) + for _, v := range acceptedValues { + hash, err := argon2id.Hash(v) + require.NoError(t, err, "hashing an accepted header value must succeed") + hashes = append(hashes, hash) + } + return NewHeader(headerName, hashes) } func TestProtect_HeaderAuth_ForwardsOnSuccess(t *testing.T) { mw := NewMiddleware(log.StandardLogger(), nil, nil) kp := generateTestKeyPair(t) - hdr := newHeaderSchemeWithToken(t, kp, "X-API-Key", "secret-key") + hdr := newHeaderScheme(t, "X-API-Key", "secret-key") require.NoError(t, mw.AddDomain("example.com", []Scheme{hdr}, kp.PublicKey, time.Hour, "acc1", "svc1", nil, false)) var backendCalled bool @@ -1075,19 +1063,12 @@ func TestProtect_HeaderAuth_ForwardsOnSuccess(t *testing.T) { assert.Equal(t, http.StatusOK, rec.Code) assert.Equal(t, "ok", rec.Body.String()) - // Session cookie should be set. - var sessionCookie *http.Cookie + // The credential rides on every request, so no session cookie is issued. for _, c := range rec.Result().Cookies() { - if c.Name == auth.SessionCookieName { - sessionCookie = c - break - } + assert.NotEqual(t, auth.SessionCookieName, c.Name, "header auth must not issue a session cookie") } - require.NotNil(t, sessionCookie, "session cookie should be set after successful header auth") - assert.True(t, sessionCookie.HttpOnly) - assert.True(t, sessionCookie.Secure) - assert.Equal(t, "header-user", capturedData.GetUserID()) + assert.Equal(t, auth.HeaderUserID, capturedData.GetUserID()) assert.Equal(t, "header", capturedData.GetAuthMethod()) } @@ -1095,7 +1076,7 @@ func TestProtect_HeaderAuth_MissingHeaderFallsThrough(t *testing.T) { mw := NewMiddleware(log.StandardLogger(), nil, nil) kp := generateTestKeyPair(t) - hdr := newHeaderSchemeWithToken(t, kp, "X-API-Key", "secret-key") + hdr := newHeaderScheme(t, "X-API-Key", "secret-key") // Also add a PIN scheme so we can verify fallthrough behavior. pinScheme := &stubScheme{method: auth.MethodPIN, promptID: "pin"} require.NoError(t, mw.AddDomain("example.com", []Scheme{hdr, pinScheme}, kp.PublicKey, time.Hour, "acc1", "svc1", nil, false)) @@ -1114,10 +1095,7 @@ func TestProtect_HeaderAuth_WrongValueReturns401(t *testing.T) { mw := NewMiddleware(log.StandardLogger(), nil, nil) kp := generateTestKeyPair(t) - mock := &mockAuthenticator{fn: func(_ context.Context, _ *proto.AuthenticateRequest) (*proto.AuthenticateResponse, error) { - return &proto.AuthenticateResponse{Success: false}, nil - }} - hdr := NewHeader(mock, "svc1", "acc1", "X-API-Key") + hdr := newHeaderScheme(t, "X-API-Key", "secret-key") require.NoError(t, mw.AddDomain("example.com", []Scheme{hdr}, kp.PublicKey, time.Hour, "acc1", "svc1", nil, false)) capturedData := proxy.NewCapturedData("") @@ -1131,93 +1109,282 @@ func TestProtect_HeaderAuth_WrongValueReturns401(t *testing.T) { assert.Equal(t, http.StatusUnauthorized, rec.Code) assert.Equal(t, "header", capturedData.GetAuthMethod()) + assert.Empty(t, hdr.verified.seen, "a rejected value must not be memoized") } -func TestProtect_HeaderAuth_InfraErrorReturns502(t *testing.T) { +// TestProtect_HeaderAuth_MatchesAnyConfiguredHeader covers a client that carries +// a valid credential on one configured header while also sending an unrelated +// value on another — an app-level Authorization alongside an API key, say. +// Schemes OR across header names, so the valid credential admits the request no +// matter which order the mapping happened to list the headers in. +func TestProtect_HeaderAuth_MatchesAnyConfiguredHeader(t *testing.T) { + tests := []struct { + name string + matchedLast bool + }{ + {name: "unmatched header listed first", matchedLast: true}, + {name: "matched header listed first", matchedLast: false}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + mw := NewMiddleware(log.StandardLogger(), nil, nil) + kp := generateTestKeyPair(t) + + authz := newHeaderScheme(t, "Authorization", "Bearer proxy-secret") + apiKey := newHeaderScheme(t, "X-Api-Key", "secret-key") + schemes := []Scheme{apiKey, authz} + if tt.matchedLast { + schemes = []Scheme{authz, apiKey} + } + require.NoError(t, mw.AddDomain("example.com", schemes, kp.PublicKey, time.Hour, "acc1", "svc1", nil, false)) + + var backendCalled bool + handler := mw.Protect(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + backendCalled = true + w.WriteHeader(http.StatusOK) + })) + + req := httptest.NewRequest(http.MethodGet, "http://example.com/", nil) + req.Header.Set("X-Api-Key", "secret-key") + req.Header.Set("Authorization", "Bearer app-level-token") + rec := httptest.NewRecorder() + handler.ServeHTTP(rec, req) + + assert.True(t, backendCalled, "a valid credential on one header must admit the request") + assert.Equal(t, http.StatusOK, rec.Code) + }) + } +} + +// TestProtect_HeaderAuth_RejectsWhenEveryPresentedHeaderFails is the other half +// of the OR: trying all schemes before rejecting must not turn into admitting a +// request that satisfied none of them. +func TestProtect_HeaderAuth_RejectsWhenEveryPresentedHeaderFails(t *testing.T) { mw := NewMiddleware(log.StandardLogger(), nil, nil) kp := generateTestKeyPair(t) - mock := &mockAuthenticator{fn: func(_ context.Context, _ *proto.AuthenticateRequest) (*proto.AuthenticateResponse, error) { - return nil, errors.New("gRPC unavailable") - }} - hdr := NewHeader(mock, "svc1", "acc1", "X-API-Key") - require.NoError(t, mw.AddDomain("example.com", []Scheme{hdr}, kp.PublicKey, time.Hour, "acc1", "svc1", nil, false)) - - handler := mw.Protect(newPassthroughHandler()) - - req := httptest.NewRequest(http.MethodGet, "http://example.com/", nil) - req.Header.Set("X-API-Key", "some-key") - rec := httptest.NewRecorder() - handler.ServeHTTP(rec, req) - - assert.Equal(t, http.StatusBadGateway, rec.Code) -} - -func TestProtect_HeaderAuth_SubsequentRequestUsesSessionCookie(t *testing.T) { - mw := NewMiddleware(log.StandardLogger(), nil, nil) - kp := generateTestKeyPair(t) - - hdr := newHeaderSchemeWithToken(t, kp, "X-API-Key", "secret-key") - require.NoError(t, mw.AddDomain("example.com", []Scheme{hdr}, kp.PublicKey, time.Hour, "acc1", "svc1", nil, false)) + authz := newHeaderScheme(t, "Authorization", "Bearer proxy-secret") + apiKey := newHeaderScheme(t, "X-Api-Key", "secret-key") + require.NoError(t, mw.AddDomain("example.com", []Scheme{authz, apiKey}, kp.PublicKey, time.Hour, "acc1", "svc1", nil, false)) + var backendCalled bool handler := mw.Protect(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + backendCalled = true + w.WriteHeader(http.StatusOK) + })) + + req := httptest.NewRequest(http.MethodGet, "http://example.com/", nil) + req.Header.Set("X-Api-Key", "wrong-key") + req.Header.Set("Authorization", "Bearer wrong-token") + rec := httptest.NewRecorder() + handler.ServeHTTP(rec, req) + + assert.False(t, backendCalled) + assert.Equal(t, http.StatusUnauthorized, rec.Code) +} + +// TestProtect_HeaderAuth_ReportsUndecodableHash covers a stored hash the proxy +// cannot decode. No credential can ever match it, so the header is permanently +// unauthenticatable — an operator fault that has to surface loudly instead of +// hiding behind the same quiet 401 a wrong credential earns. +func TestProtect_HeaderAuth_ReportsUndecodableHash(t *testing.T) { + validHash, err := argon2id.Hash("secret-key") + require.NoError(t, err) + + tests := []struct { + name string + hashes []string + wantErrLog bool + }{ + {name: "stored hash cannot be decoded", hashes: []string{"$argon2id$v=19$garbage"}, wantErrLog: true}, + {name: "wrong credential against a good hash", hashes: []string{validHash}, wantErrLog: false}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + logger, hook := logtest.NewNullLogger() + logger.SetLevel(log.DebugLevel) + mw := NewMiddleware(logger, nil, nil) + kp := generateTestKeyPair(t) + + require.NoError(t, mw.AddDomain("example.com", []Scheme{NewHeader("X-Api-Key", tt.hashes)}, + kp.PublicKey, time.Hour, "acc1", "svc1", nil, false)) + + handler := mw.Protect(newPassthroughHandler()) + + req := httptest.NewRequest(http.MethodGet, "http://example.com/", nil) + req.Header.Set("X-Api-Key", "wrong-key") + rec := httptest.NewRecorder() + handler.ServeHTTP(rec, req) + + require.Equal(t, http.StatusUnauthorized, rec.Code, "either way the request is denied") + + var errored []string + for _, entry := range hook.AllEntries() { + if entry.Level == log.ErrorLevel { + errored = append(errored, entry.Message) + } + } + + if !tt.wantErrLog { + assert.Empty(t, errored, "a wrong credential is not an operator fault") + return + } + require.Len(t, errored, 1, "an undecodable hash must be reported once") + assert.Contains(t, errored[0], "cannot be decoded") + }) + } +} + +// TestProtect_HeaderAuth_NoHashesFailsClosed covers a mapping that names a +// header but carries no hash for it: the check cannot be evaluated, so the +// request must be denied rather than let through unauthenticated. +func TestProtect_HeaderAuth_NoHashesFailsClosed(t *testing.T) { + mw := NewMiddleware(log.StandardLogger(), nil, nil) + kp := generateTestKeyPair(t) + + hdr := NewHeader("X-API-Key", nil) + require.NoError(t, mw.AddDomain("example.com", []Scheme{hdr}, kp.PublicKey, time.Hour, "acc1", "svc1", nil, false)) + + var backendCalled bool + handler := mw.Protect(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + backendCalled = true + w.WriteHeader(http.StatusOK) + })) + + req := httptest.NewRequest(http.MethodGet, "http://example.com/", nil) + req.Header.Set("X-API-Key", "any-key") + rec := httptest.NewRecorder() + handler.ServeHTTP(rec, req) + + assert.Equal(t, http.StatusUnauthorized, rec.Code) + assert.False(t, backendCalled, "a header auth with no hashes must not admit the request") +} + +// TestProtect_HeaderAuth_SubsequentRequestRequiresHeader verifies that header +// auth grants no ambient session: a follow-up request that drops the header is +// treated as unauthenticated. +func TestProtect_HeaderAuth_SubsequentRequestRequiresHeader(t *testing.T) { + mw := NewMiddleware(log.StandardLogger(), nil, nil) + kp := generateTestKeyPair(t) + + hdr := newHeaderScheme(t, "X-API-Key", "secret-key") + require.NoError(t, mw.AddDomain("example.com", []Scheme{hdr}, kp.PublicKey, time.Hour, "acc1", "svc1", nil, false)) + + var backendCalls int + handler := mw.Protect(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + backendCalls++ w.WriteHeader(http.StatusOK) })) - // First request with header auth. req1 := httptest.NewRequest(http.MethodGet, "http://example.com/", nil) req1.Header.Set("X-API-Key", "secret-key") req1 = req1.WithContext(proxy.WithCapturedData(req1.Context(), proxy.NewCapturedData(""))) rec1 := httptest.NewRecorder() handler.ServeHTTP(rec1, req1) require.Equal(t, http.StatusOK, rec1.Code) + require.Equal(t, 1, backendCalls) - // Extract session cookie. - var sessionCookie *http.Cookie - for _, c := range rec1.Result().Cookies() { - if c.Name == auth.SessionCookieName { - sessionCookie = c - break - } - } - require.NotNil(t, sessionCookie) - - // Second request with only the session cookie (no header). - capturedData2 := proxy.NewCapturedData("") + // Same client, second request, header omitted: no cookie was handed out, so + // there is nothing to carry the earlier success forward. req2 := httptest.NewRequest(http.MethodGet, "http://example.com/other", nil) - req2.AddCookie(sessionCookie) - req2 = req2.WithContext(proxy.WithCapturedData(req2.Context(), capturedData2)) + for _, c := range rec1.Result().Cookies() { + req2.AddCookie(c) + } rec2 := httptest.NewRecorder() handler.ServeHTTP(rec2, req2) - assert.Equal(t, http.StatusOK, rec2.Code) - assert.Equal(t, "header-user", capturedData2.GetUserID()) - assert.Equal(t, "header", capturedData2.GetAuthMethod()) + assert.Equal(t, http.StatusUnauthorized, rec2.Code, "dropping the header must revoke access") + assert.Equal(t, 1, backendCalls, "backend must not be reached without the header") } -// TestProtect_HeaderAuth_MultipleValuesSameHeader verifies that the proxy -// correctly handles multiple valid credentials for the same header name. -// In production, the mgmt gRPC authenticateHeader iterates all configured -// header auths and accepts if any hash matches (OR semantics). The proxy -// creates one Header scheme per entry, but a single gRPC call checks all. +// TestProtect_HeaderAuth_LegacySessionCookieIsIgnored covers the upgrade +// window. Header auth used to mint a session token, so cookies with +// method=header survive a proxy upgrade and stay signature-valid for their full +// lifetime. They must not stand in for the header, or a credential rotated +// right after the upgrade would keep working until every such token expired. +func TestProtect_HeaderAuth_LegacySessionCookieIsIgnored(t *testing.T) { + mw := NewMiddleware(log.StandardLogger(), nil, nil) + kp := generateTestKeyPair(t) + + hdr := newHeaderScheme(t, "X-API-Key", "secret-key") + require.NoError(t, mw.AddDomain("example.com", []Scheme{hdr}, kp.PublicKey, time.Hour, "acc1", "svc1", nil, false)) + + // A token management would have minted for header auth before the upgrade. + legacyToken, err := sessionkey.SignToken(kp.PrivateKey, auth.HeaderUserID, "", "example.com", auth.MethodHeader, nil, nil, time.Hour) + require.NoError(t, err) + + var backendCalls int + handler := mw.Protect(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + backendCalls++ + w.WriteHeader(http.StatusOK) + })) + + t.Run("cookie alone is rejected", func(t *testing.T) { + req := httptest.NewRequest(http.MethodGet, "http://example.com/", nil) + req.AddCookie(&http.Cookie{Name: auth.SessionCookieName, Value: legacyToken}) + rec := httptest.NewRecorder() + handler.ServeHTTP(rec, req) + + assert.Equal(t, http.StatusUnauthorized, rec.Code, "a header-auth cookie must not authenticate on its own") + assert.Equal(t, 0, backendCalls, "backend must not be reached without the header") + }) + + t.Run("cookie does not block the header path", func(t *testing.T) { + req := httptest.NewRequest(http.MethodGet, "http://example.com/", nil) + req.AddCookie(&http.Cookie{Name: auth.SessionCookieName, Value: legacyToken}) + req.Header.Set("X-API-Key", "secret-key") + rec := httptest.NewRecorder() + handler.ServeHTTP(rec, req) + + assert.Equal(t, http.StatusOK, rec.Code, "a client sending both must still be admitted by the header") + assert.Equal(t, 1, backendCalls) + }) +} + +// TestProtect_HeaderAuth_RepeatedValueIsMemoized verifies the KDF is run once +// per distinct accepted value. argon2id is deliberately expensive, so a +// credential that repeats on every request must not be re-derived each time. +func TestProtect_HeaderAuth_RepeatedValueIsMemoized(t *testing.T) { + mw := NewMiddleware(log.StandardLogger(), nil, nil) + kp := generateTestKeyPair(t) + + hdr := newHeaderScheme(t, "X-API-Key", "key-a", "key-b") + require.NoError(t, mw.AddDomain("example.com", []Scheme{hdr}, kp.PublicKey, time.Hour, "acc1", "svc1", nil, false)) + + handler := mw.Protect(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + w.WriteHeader(http.StatusOK) + })) + + get := func(value string) int { + req := httptest.NewRequest(http.MethodGet, "http://example.com/", nil) + req.Header.Set("X-API-Key", value) + rec := httptest.NewRecorder() + handler.ServeHTTP(rec, req) + return rec.Code + } + + require.Equal(t, http.StatusOK, get("key-a")) + require.Equal(t, http.StatusOK, get("key-a")) + assert.Len(t, hdr.verified.seen, 1, "the same value must be memoized once") + + require.Equal(t, http.StatusOK, get("key-b")) + assert.Len(t, hdr.verified.seen, 2, "each accepted value gets its own entry") + + require.Equal(t, http.StatusUnauthorized, get("key-c")) + assert.Len(t, hdr.verified.seen, 2, "rejected values must not grow the set") +} + +// TestProtect_HeaderAuth_MultipleValuesSameHeader verifies that a service with +// several accepted credentials for one header name accepts any of them. +// Management applied these OR semantics while it still validated the value; the +// proxy preserves them by carrying every hash for a name on one scheme. func TestProtect_HeaderAuth_MultipleValuesSameHeader(t *testing.T) { mw := NewMiddleware(log.StandardLogger(), nil, nil) kp := generateTestKeyPair(t) - // Mock simulates mgmt behavior: accepts either token-a or token-b. - accepted := map[string]bool{"Bearer token-a": true, "Bearer token-b": true} - mock := &mockAuthenticator{fn: func(_ context.Context, req *proto.AuthenticateRequest) (*proto.AuthenticateResponse, error) { - ha := req.GetHeaderAuth() - if ha != nil && accepted[ha.GetHeaderValue()] { - token, err := sessionkey.SignToken(kp.PrivateKey, "header-user", "", "example.com", auth.MethodHeader, nil, nil, time.Hour) - require.NoError(t, err) - return &proto.AuthenticateResponse{Success: true, SessionToken: token}, nil - } - return &proto.AuthenticateResponse{Success: false}, nil - }} - - // Single Header scheme (as if one entry existed), but the mock checks both values. - hdr := NewHeader(mock, "svc1", "acc1", "Authorization") + hdr := newHeaderScheme(t, "Authorization", "Bearer token-a", "Bearer token-b") require.NoError(t, mw.AddDomain("example.com", []Scheme{hdr}, kp.PublicKey, time.Hour, "acc1", "svc1", nil, false)) var backendCalled bool diff --git a/proxy/server.go b/proxy/server.go index aee748339..38477fb87 100644 --- a/proxy/server.go +++ b/proxy/server.go @@ -20,6 +20,7 @@ import ( "net/url" "path/filepath" "reflect" + "slices" "sync" "time" @@ -2062,9 +2063,7 @@ func (s *Server) updateMapping(ctx context.Context, mapping *proto.ProxyMapping) if mapping.GetAuth().GetOidc() { schemes = append(schemes, auth.NewOIDC(s.mgmtClient, svcID, accountID, s.ForwardedProto)) } - for _, ha := range mapping.GetAuth().GetHeaderAuths() { - schemes = append(schemes, auth.NewHeader(s.mgmtClient, svcID, accountID, ha.GetHeader())) - } + schemes = append(schemes, headerAuthSchemes(mapping.GetAuth().GetHeaderAuths())...) ipRestrictions := s.parseRestrictions(mapping) s.warnIfGeoUnavailable(mapping.GetDomain(), mapping.GetAccessRestrictions()) @@ -2088,6 +2087,32 @@ func (s *Server) updateMapping(ctx context.Context, mapping *proto.ProxyMapping) return nil } +// headerAuthSchemes builds one scheme per canonical header name, carrying every +// hash configured for that name so any of them is accepted — the OR semantics +// management applied while it still validated the credential itself. No entry is +// ever dropped: a name that arrives blank, or without a hash, still yields a +// scheme, because a mapping that lost its only scheme would fall through +// Protect's no-schemes pass-through and serve the domain unauthenticated. +func headerAuthSchemes(headerAuths []*proto.HeaderAuth) []auth.Scheme { + names := make([]string, 0, len(headerAuths)) + hashes := make(map[string][]string, len(headerAuths)) + for _, ha := range headerAuths { + name := http.CanonicalHeaderKey(ha.GetHeader()) + if !slices.Contains(names, name) { + names = append(names, name) + } + if hash := ha.GetHashedValue(); hash != "" { + hashes[name] = append(hashes[name], hash) + } + } + + schemes := make([]auth.Scheme, 0, len(names)) + for _, name := range names { + schemes = append(schemes, auth.NewHeader(name, hashes[name])) + } + return schemes +} + // initMiddlewareManager wires the middleware subsystem at boot. It configures // the per-process FactoryContext concrete middlewares consult, installs the // live-service check, and binds the resolver to the registry concrete diff --git a/proxy/server_test.go b/proxy/server_test.go index f0c4765db..9cef63b95 100644 --- a/proxy/server_test.go +++ b/proxy/server_test.go @@ -6,6 +6,8 @@ import ( "fmt" "io" "net" + "net/http" + "net/http/httptest" "testing" "time" @@ -15,8 +17,10 @@ import ( "go.opentelemetry.io/otel/metric/noop" "google.golang.org/grpc" + "github.com/netbirdio/netbird/proxy/internal/auth" proxymetrics "github.com/netbirdio/netbird/proxy/internal/metrics" "github.com/netbirdio/netbird/proxy/internal/types" + "github.com/netbirdio/netbird/shared/hash/argon2id" "github.com/netbirdio/netbird/shared/management/proto" ) @@ -209,6 +213,62 @@ func TestRedactMappingForLog_HandlesEmptyOrNilFields(t *testing.T) { assert.Empty(t, redacted.Path, "empty Path must remain empty") } +// headerSchemeAccepts reports whether the scheme admits value for headerName. +func headerSchemeAccepts(t *testing.T, scheme auth.Scheme, headerName, value string) bool { + t.Helper() + hdr, ok := scheme.(auth.Header) + require.True(t, ok, "header auths must produce Header schemes") + + req := httptest.NewRequest(http.MethodGet, "http://example.com/", nil) + req.Header.Set(headerName, value) + _, matched, _ := hdr.Verify(req) + return matched +} + +func TestHeaderAuthSchemes_GroupsValuesByCanonicalHeaderName(t *testing.T) { + hashOf := func(v string) string { + hash, err := argon2id.Hash(v) + require.NoError(t, err) + return hash + } + + schemes := headerAuthSchemes([]*proto.HeaderAuth{ + {Header: "Authorization", HashedValue: hashOf("Bearer a")}, + {Header: "authorization", HashedValue: hashOf("Bearer b")}, + {Header: "X-Api-Key", HashedValue: hashOf("key-1")}, + }) + + require.Len(t, schemes, 2, "entries differing only in header-name case must collapse into one scheme") + + assert.True(t, headerSchemeAccepts(t, schemes[0], "Authorization", "Bearer a"), "first value for the header must be accepted") + assert.True(t, headerSchemeAccepts(t, schemes[0], "Authorization", "Bearer b"), "second value for the same header must be accepted") + assert.False(t, headerSchemeAccepts(t, schemes[0], "Authorization", "Bearer c"), "unconfigured value must be rejected") + assert.True(t, headerSchemeAccepts(t, schemes[1], "X-Api-Key", "key-1"), "a second header name keeps its own scheme") +} + +// TestHeaderAuthSchemes_MissingHashFailsClosed covers a mapping that names a +// header but carries no hash for it. Dropping the scheme would leave a service +// whose only auth is that header wide open, so the scheme is kept and denies. +func TestHeaderAuthSchemes_MissingHashFailsClosed(t *testing.T) { + schemes := headerAuthSchemes([]*proto.HeaderAuth{{Header: "X-Api-Key"}}) + + require.Len(t, schemes, 1, "a header without a hash must still register a scheme") + assert.False(t, headerSchemeAccepts(t, schemes[0], "X-Api-Key", "anything"), + "a header auth without a hash must reject every value") +} + +// TestHeaderAuthSchemes_BlankNameFailsClosed covers a mapping row whose header +// name is empty. Skipping it would leave a service whose only auth is that entry +// with no schemes at all, which Protect treats as an unprotected domain, so the +// entry is kept and the domain stays gated. +func TestHeaderAuthSchemes_BlankNameFailsClosed(t *testing.T) { + schemes := headerAuthSchemes([]*proto.HeaderAuth{{Header: "", HashedValue: "$argon2id$not-a-real-hash"}}) + + require.Len(t, schemes, 1, "a blank header name must still register a scheme") + assert.False(t, headerSchemeAccepts(t, schemes[0], "X-Api-Key", "anything"), + "a blank header auth must not admit any request") +} + type statusUpdateOnlyClient struct { proto.ProxyServiceClient }