mirror of
https://github.com/netbirdio/netbird.git
synced 2026-08-26 09:31:30 +02:00
Merge branch 'main' of github.com:netbirdio/netbird into feature/shared-service-config-loader
This commit is contained in:
18
client/ui/frontend/src/components/ReadySignal.tsx
Normal file
18
client/ui/frontend/src/components/ReadySignal.tsx
Normal file
@@ -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;
|
||||
};
|
||||
@@ -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 = () => {
|
||||
<DebugBundleProvider>
|
||||
<ClientVersionProvider>
|
||||
<Outlet />
|
||||
<ReadySignal />
|
||||
</ClientVersionProvider>
|
||||
</DebugBundleProvider>
|
||||
</RestrictionsProvider>
|
||||
|
||||
@@ -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()
|
||||
})
|
||||
}
|
||||
|
||||
@@ -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).
|
||||
|
||||
@@ -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) ||
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
2
go.mod
2
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
|
||||
|
||||
4
go.sum
4
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=
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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".
|
||||
|
||||
@@ -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{}{}
|
||||
}
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user