Merge remote-tracking branch 'origin/main' into feat-post_quantum_ml_kem

This commit is contained in:
riccardom
2026-10-05 16:03:28 +02:00
85 changed files with 3037 additions and 624 deletions
+30 -9
View File
@@ -213,6 +213,7 @@ func (c *Client) Run(platformFiles PlatformFiles, urlOpener URLOpener, isAndroid
connectClient := internal.NewConnectClient(ctx, cfg, c.recorder,
internal.WithNetEvents(c.netMgr))
c.setState(cfg, cacheDir, cfgFile, connectClient)
connectClient.SetSyncResponsePersistence(true)
// This path runs the interactive SSO flow, so reaching here means the peer
// is authenticated again — release the latch Status() reports from. Clear
// only once the fresh connect client is installed: until then Status()
@@ -256,6 +257,7 @@ func (c *Client) RunWithoutLogin(platformFiles PlatformFiles, dns *DNSList, dnsR
connectClient := internal.NewConnectClient(ctx, cfg, c.recorder,
internal.WithNetEvents(c.netMgr))
c.setState(cfg, cacheDir, cfgFile, connectClient)
connectClient.SetSyncResponsePersistence(true)
return connectClient.RunOnAndroid(c.tunAdapter, c.iFaceDiscover, c.networkChangeListener, slices.Clone(dns.items), dnsReadyListener, stateFile, cacheDir)
}
@@ -327,6 +329,19 @@ func (c *Client) NotifyNetworkChange() {
// or "strict"; strict also anonymizes internal IP ranges, peer names, and
// WireGuard public keys, and implies anonymize.
func (c *Client) DebugBundle(platformFiles PlatformFiles, anonymize bool, anonymizeLevel string) (string, error) {
return c.debugBundle(platformFiles, anonymize, anonymizeLevel, true)
}
// DebugBundleFile generates a debug bundle and returns the path of the zip in
// the cache directory instead of uploading it, so the app can hand the file to
// the user for inspection. The caller owns the file and removes it once done;
// the stale-bundle cleanup of later runs removes it only after a day.
// anonymize and anonymizeLevel behave as in DebugBundle.
func (c *Client) DebugBundleFile(platformFiles PlatformFiles, anonymize bool, anonymizeLevel string) (string, error) {
return c.debugBundle(platformFiles, anonymize, anonymizeLevel, false)
}
func (c *Client) debugBundle(platformFiles PlatformFiles, anonymize bool, anonymizeLevel string, upload bool) (string, error) {
cfg, cacheDir, cc := c.stateSnapshot()
// If the engine hasn't been started, load config from disk
@@ -342,6 +357,11 @@ func (c *Client) DebugBundle(platformFiles PlatformFiles, anonymize bool, anonym
cacheDir = platformFiles.CacheDir()
}
// Clear what an interrupted earlier run may have left in the cache before
// adding to it. Remote debug jobs write to the same directory, so anything
// younger than an hour is treated as possibly still in use.
debug.RemoveStaleBundles(cacheDir, time.Hour)
deps := debug.GeneratorDependencies{
InternalConfig: cfg,
StatusRecorder: c.recorder,
@@ -379,6 +399,9 @@ func (c *Client) DebugBundle(platformFiles PlatformFiles, anonymize bool, anonym
if err != nil {
return "", fmt.Errorf("generate debug bundle: %w", err)
}
if !upload {
return debug.ExportBundle(path)
}
defer func() {
if err := os.Remove(path); err != nil {
log.Errorf("failed to remove debug bundle file: %v", err)
@@ -475,6 +498,7 @@ func (c *Client) Networks() *NetworkArray {
routesMap := routeManager.GetClientRoutesWithNetID()
v6Merged := route.V6ExitMergeSet(routesMap)
resolvedDomains := c.recorder.GetResolvedDomainsStates()
activeRoutePeers := c.recorder.GetActiveRoutePeers()
networkArray := &NetworkArray{
items: make([]Network, 0),
@@ -488,7 +512,7 @@ func (c *Client) Networks() *NetworkArray {
continue
}
network := c.buildNetwork(id, routes, routeSelector.IsSelected(id), resolvedDomains, v6Merged)
network := c.buildNetwork(id, routes, routeSelector.IsSelected(id), resolvedDomains, v6Merged, activeRoutePeers)
if network == nil {
continue
}
@@ -497,14 +521,14 @@ func (c *Client) Networks() *NetworkArray {
return networkArray
}
func (c *Client) buildNetwork(id route.NetID, routes []*route.Route, selected bool, resolvedDomains map[domain.Domain]peer.ResolvedDomainInfo, v6Merged map[route.NetID]struct{}) *Network {
func (c *Client) buildNetwork(id route.NetID, routes []*route.Route, selected bool, resolvedDomains map[domain.Domain]peer.ResolvedDomainInfo, v6Merged map[route.NetID]struct{}, activeRoutePeers map[route.HAUniqueID]string) *Network {
r := routes[0]
netStr := r.Network.String()
if r.IsDynamic() {
netStr = r.Domains.SafeString()
}
routePeer, err := c.findBestRoutePeer(routes)
routePeer, err := c.findBestRoutePeer(routes, activeRoutePeers)
if err != nil {
log.Errorf("could not get peer info for route %s: %v", id, err)
return nil
@@ -528,12 +552,9 @@ func (c *Client) buildNetwork(id route.NetID, routes []*route.Route, selected bo
// findBestRoutePeer returns the peer actively routing traffic for the given
// HA route group. Falls back to the first connected peer, then the first peer.
func (c *Client) findBestRoutePeer(routes []*route.Route) (peer.State, error) {
netStr := routes[0].Network.String()
fullStatus := c.recorder.GetFullStatus()
for _, p := range fullStatus.Peers {
if _, ok := p.GetRoutes()[netStr]; ok {
func (c *Client) findBestRoutePeer(routes []*route.Route, activeRoutePeers map[route.HAUniqueID]string) (peer.State, error) {
if peerKey, ok := activeRoutePeers[routes[0].GetHAUniqueID()]; ok {
if p, err := c.recorder.GetPeer(peerKey); err == nil {
return p, nil
}
}
+16 -36
View File
@@ -40,14 +40,18 @@ func init() {
peerPubKey = peerPrivateKey.PublicKey().String()
}
// testIFaceBlackList mirrors the prefixes profilemanager.DefaultInterfaceBlacklist
// carries for the overlay interface. These tests create their own utun device, and
// stdnet's filter probes with wgctrl every interface it is not told to skip, which
// on a userspace WireGuard platform reaches the UAPI socket of this same process.
// Declared here rather than imported because profilemanager imports this package.
var testIFaceBlackList = []string{"wt", "utun", "tun0"}
func TestWGIface_UpdateAddr(t *testing.T) {
ifaceName := fmt.Sprintf("utun%d", WgIntNumber+4)
addr := "100.64.0.1/8"
wgPort := 33100
newNet, err := stdnet.NewNet(context.Background(), nil)
if err != nil {
t.Fatal(err)
}
newNet := stdnet.NewNet(context.Background(), testIFaceBlackList)
opts := WGIFaceOpts{
IFaceName: ifaceName,
@@ -127,10 +131,7 @@ func getIfaceAddrs(ifaceName string) ([]net.Addr, error) {
func Test_CreateInterface(t *testing.T) {
ifaceName := fmt.Sprintf("utun%d", WgIntNumber+1)
wgIP := "10.99.99.1/32"
newNet, err := stdnet.NewNet(context.Background(), nil)
if err != nil {
t.Fatal(err)
}
newNet := stdnet.NewNet(context.Background(), testIFaceBlackList)
opts := WGIFaceOpts{
IFaceName: ifaceName,
Address: wgaddr.MustParseWGAddress(wgIP),
@@ -170,10 +171,7 @@ func Test_Close(t *testing.T) {
ifaceName := fmt.Sprintf("utun%d", WgIntNumber+2)
wgIP := "10.99.99.2/32"
wgPort := 33100
newNet, err := stdnet.NewNet(context.Background(), nil)
if err != nil {
t.Fatal(err)
}
newNet := stdnet.NewNet(context.Background(), testIFaceBlackList)
opts := WGIFaceOpts{
IFaceName: ifaceName,
@@ -215,10 +213,7 @@ func TestRecreation(t *testing.T) {
ifaceName := fmt.Sprintf("utun%d", WgIntNumber+2)
wgIP := "10.99.99.2/32"
wgPort := 33100
newNet, err := stdnet.NewNet(context.Background(), nil)
if err != nil {
t.Fatal(err)
}
newNet := stdnet.NewNet(context.Background(), testIFaceBlackList)
opts := WGIFaceOpts{
IFaceName: ifaceName,
@@ -288,10 +283,7 @@ func Test_ConfigureInterface(t *testing.T) {
ifaceName := fmt.Sprintf("utun%d", WgIntNumber+3)
wgIP := "10.99.99.5/30"
wgPort := 33100
newNet, err := stdnet.NewNet(context.Background(), nil)
if err != nil {
t.Fatal(err)
}
newNet := stdnet.NewNet(context.Background(), testIFaceBlackList)
opts := WGIFaceOpts{
IFaceName: ifaceName,
Address: wgaddr.MustParseWGAddress(wgIP),
@@ -343,10 +335,7 @@ func Test_ConfigureInterface(t *testing.T) {
func Test_UpdatePeer(t *testing.T) {
ifaceName := fmt.Sprintf("utun%d", WgIntNumber+4)
wgIP := "10.99.99.9/30"
newNet, err := stdnet.NewNet(context.Background(), nil)
if err != nil {
t.Fatal(err)
}
newNet := stdnet.NewNet(context.Background(), testIFaceBlackList)
opts := WGIFaceOpts{
IFaceName: ifaceName,
@@ -413,10 +402,7 @@ func Test_UpdatePeer(t *testing.T) {
func Test_RemovePeer(t *testing.T) {
ifaceName := fmt.Sprintf("utun%d", WgIntNumber+4)
wgIP := "10.99.99.13/30"
newNet, err := stdnet.NewNet(context.Background(), nil)
if err != nil {
t.Fatal(err)
}
newNet := stdnet.NewNet(context.Background(), testIFaceBlackList)
opts := WGIFaceOpts{
IFaceName: ifaceName,
@@ -477,10 +463,7 @@ func Test_ConnectPeers(t *testing.T) {
peer2wgPort := 33200
keepAlive := 1 * time.Second
newNet, err := stdnet.NewNet(context.Background(), nil)
if err != nil {
t.Fatal(err)
}
newNet := stdnet.NewNet(context.Background(), testIFaceBlackList)
guid := fmt.Sprintf("{%s}", uuid.New().String())
device.CustomWindowsGUIDString = strings.ToLower(guid)
@@ -516,10 +499,7 @@ func Test_ConnectPeers(t *testing.T) {
guid = fmt.Sprintf("{%s}", uuid.New().String())
device.CustomWindowsGUIDString = strings.ToLower(guid)
newNet, err = stdnet.NewNet(context.Background(), nil)
if err != nil {
t.Fatal(err)
}
newNet = stdnet.NewNet(context.Background(), testIFaceBlackList)
optsPeer2 := WGIFaceOpts{
IFaceName: peer2ifaceName,
+1 -4
View File
@@ -200,10 +200,7 @@ func (m *SingleSocketUDPMux) updateLocalAddresses() {
}
if len(networks) > 0 {
if m.params.Net == nil {
var err error
if m.params.Net, err = stdnet.NewNet(context.Background(), nil); err != nil {
m.params.Logger.Errorf("failed to get create network: %v", err)
}
m.params.Net = stdnet.NewNet(context.Background(), nil)
}
ips, err := localInterfaces(m.params.Net, m.params.InterfaceFilter, nil, networks, true)
+53 -1
View File
@@ -380,9 +380,38 @@ func NewBundleGenerator(deps GeneratorDependencies, cfg BundleConfig) *BundleGen
}
}
// bundleFilePattern names the bundle zips Generate creates in tempDir; the
// asterisk is filled in by os.CreateTemp.
const bundleFilePattern = "netbird.debug.*.zip"
const exportedBundlePrefix = "netbird.debug-file."
const exportedBundleMaxAge = 24 * time.Hour
// RemoveStaleBundles deletes bundle zips that an interrupted generation or
// upload left behind in dir. Only files older than maxAge go, so a bundle that
// another caller is still writing or uploading in the same directory survives.
// Exported bundles are kept for exportedBundleMaxAge instead.
func RemoveStaleBundles(dir string, maxAge time.Duration) {
removeStaleFiles(dir, bundleFilePattern, maxAge)
removeStaleFiles(dir, exportedBundlePrefix+"*.zip", exportedBundleMaxAge)
}
// ExportBundle renames a generated bundle out of the RemoveStaleBundles pattern
// and returns the new path. The caller owns the file from then on; an export
// abandoned for longer than exportedBundleMaxAge is removed by RemoveStaleBundles.
func ExportBundle(path string) (string, error) {
base := strings.TrimPrefix(filepath.Base(path), strings.SplitN(bundleFilePattern, "*", 2)[0])
exported := filepath.Join(filepath.Dir(path), exportedBundlePrefix+base)
if err := os.Rename(path, exported); err != nil {
return "", fmt.Errorf("export debug bundle: %w", err)
}
return exported, nil
}
// Generate creates a debug bundle and returns the location.
func (g *BundleGenerator) Generate() (resp string, err error) {
bundlePath, err := os.CreateTemp(g.tempDir, "netbird.debug.*.zip")
bundlePath, err := os.CreateTemp(g.tempDir, bundleFilePattern)
if err != nil {
return "", fmt.Errorf("create zip file: %w", err)
}
@@ -1729,3 +1758,26 @@ func anonymizeSlice(v []any, anonymizer *anonymize.Anonymizer) []any {
}
return v
}
func removeStaleFiles(dir, pattern string, maxAge time.Duration) {
matches, err := filepath.Glob(filepath.Join(dir, pattern))
if err != nil {
log.Debugf("glob stale debug bundles in %s: %v", dir, err)
return
}
cutoff := time.Now().Add(-maxAge)
for _, path := range matches {
info, err := os.Stat(path)
if err != nil || info.ModTime().After(cutoff) {
continue
}
if err := os.Remove(path); err != nil {
if !errors.Is(err, fs.ErrNotExist) {
log.Warnf("remove stale debug bundle %s: %v", path, err)
}
continue
}
log.Infof("removed stale debug bundle %s", path)
}
}
+50
View File
@@ -4,6 +4,7 @@ import (
"archive/zip"
"bytes"
"encoding/json"
"fmt"
"net"
"net/netip"
"net/url"
@@ -969,3 +970,52 @@ func renderAddConfigSpecific(g *BundleGenerator) string {
func newAnonymizerForTest() *anonymize.Anonymizer {
return anonymize.NewAnonymizer(anonymize.DefaultAddresses())
}
func TestRemoveStaleBundles(t *testing.T) {
dir := t.TempDir()
stale := filepath.Join(dir, "netbird.debug.111.zip")
fresh := filepath.Join(dir, "netbird.debug.222.zip")
other := filepath.Join(dir, "netbird.debug.333.txt")
owned := filepath.Join(dir, "netbird.debug.444.zip")
abandoned := filepath.Join(dir, "netbird.debug.555.zip")
for _, p := range []string{stale, fresh, other, owned, abandoned} {
require.NoError(t, os.WriteFile(p, []byte("x"), 0o600))
}
exported, err := ExportBundle(owned)
require.NoError(t, err)
exportedAbandoned, err := ExportBundle(abandoned)
require.NoError(t, err)
old := time.Now().Add(-2 * time.Hour)
for _, p := range []string{stale, other, exported} {
require.NoError(t, os.Chtimes(p, old, old))
}
ancient := time.Now().Add(-exportedBundleMaxAge - time.Hour)
require.NoError(t, os.Chtimes(exportedAbandoned, ancient, ancient))
RemoveStaleBundles(dir, time.Hour)
assert.NoFileExists(t, stale, "bundle older than maxAge should be removed")
assert.FileExists(t, fresh, "bundle younger than maxAge must survive, it may still be uploading")
assert.FileExists(t, other, "files outside the bundle pattern must not be touched")
assert.NoFileExists(t, owned)
assert.FileExists(t, exported, "exported bundle is caller-owned and must survive maxAge")
assert.NoFileExists(t, exportedAbandoned, "exported bundle older than exportedBundleMaxAge is abandoned")
}
func TestBundleIncludesNetworkMap(t *testing.T) {
for _, anonymize := range []bool{false, true} {
t.Run(fmt.Sprintf("anonymize=%t", anonymize), func(t *testing.T) {
g := NewBundleGenerator(GeneratorDependencies{
SyncResponse: &mgmProto.SyncResponse{NetworkMap: &mgmProto.NetworkMap{Serial: 1}},
}, BundleConfig{Anonymize: anonymize})
require.Contains(t, bundleEntries(t, g), "network_map.json")
})
}
}
func TestBundleOmitsNetworkMapWithoutSyncResponse(t *testing.T) {
g := NewBundleGenerator(GeneratorDependencies{}, BundleConfig{})
require.NotContains(t, bundleEntries(t, g), "network_map.json")
}
+7 -10
View File
@@ -9,9 +9,9 @@ import (
"os"
"testing"
"go.uber.org/mock/gomock"
"github.com/miekg/dns"
"github.com/stretchr/testify/assert"
"go.uber.org/mock/gomock"
"golang.zx2c4.com/wireguard/wgctrl/wgtypes"
"github.com/netbirdio/netbird/client/iface"
@@ -24,6 +24,10 @@ import (
nbdns "github.com/netbirdio/netbird/dns"
)
// testIFaceBlackList mirrors the overlay prefixes profilemanager.DefaultInterfaceBlacklist
// carries. Declared here rather than imported because profilemanager imports this package.
var testIFaceBlackList = []string{"wt", "utun", "tun0"}
func TestUpdateDNSServer(t *testing.T) {
nameServers := []nbdns.NameServer{
@@ -243,10 +247,7 @@ func TestUpdateDNSServer(t *testing.T) {
for n, testCase := range testCases {
t.Run(testCase.name, func(t *testing.T) {
privKey, _ := wgtypes.GenerateKey()
newNet, err := stdnet.NewNet(context.Background(), nil)
if err != nil {
t.Fatal(err)
}
newNet := stdnet.NewNet(context.Background(), testIFaceBlackList)
opts := iface.WGIFaceOpts{
IFaceName: fmt.Sprintf("utun230%d", n),
@@ -348,11 +349,7 @@ func TestDNSFakeResolverHandleUpdates(t *testing.T) {
defer t.Setenv("NB_WG_KERNEL_DISABLED", ov)
t.Setenv("NB_WG_KERNEL_DISABLED", "true")
newNet, err := stdnet.NewNet(context.Background(), []string{"utun2301"})
if err != nil {
t.Errorf("create stdnet: %v", err)
return
}
newNet := stdnet.NewNet(context.Background(), []string{"utun2301"})
privKey, _ := wgtypes.GeneratePrivateKey()
opts := iface.WGIFaceOpts{
+1 -5
View File
@@ -394,11 +394,7 @@ func createWgInterfaceWithBind(t *testing.T) (*iface.WGIface, error) {
defer t.Setenv("NB_WG_KERNEL_DISABLED", ov)
t.Setenv("NB_WG_KERNEL_DISABLED", "true")
newNet, err := stdnet.NewNet(context.Background(), []string{"utun2301"})
if err != nil {
t.Fatalf("create stdnet: %v", err)
return nil, err
}
newNet := stdnet.NewNet(context.Background(), []string{"utun2301"})
privKey, _ := wgtypes.GeneratePrivateKey()
+1 -4
View File
@@ -2249,10 +2249,7 @@ func (e *Engine) close() {
}
func (e *Engine) newWgIface() (*iface.WGIface, error) {
transportNet, err := e.newStdNet()
if err != nil {
log.Errorf("failed to create pion's stdnet: %s", err)
}
transportNet := e.newStdNet()
opts := iface.WGIFaceOpts{
IFaceName: e.config.WgIfaceName,
+15 -11
View File
@@ -12,12 +12,12 @@ import (
"testing"
"time"
"go.uber.org/mock/gomock"
"github.com/google/uuid"
log "github.com/sirupsen/logrus"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"go.opentelemetry.io/otel"
"go.uber.org/mock/gomock"
"golang.zx2c4.com/wireguard/wgctrl/wgtypes"
"google.golang.org/grpc"
"google.golang.org/grpc/keepalive"
@@ -27,6 +27,7 @@ import (
"github.com/netbirdio/netbird/client/iface/wgaddr"
"github.com/netbirdio/netbird/client/internal/dns"
"github.com/netbirdio/netbird/client/internal/peer"
"github.com/netbirdio/netbird/client/internal/profilemanager"
nbssh "github.com/netbirdio/netbird/client/ssh"
"github.com/netbirdio/netbird/client/system"
nbdns "github.com/netbirdio/netbird/dns"
@@ -81,6 +82,7 @@ func TestEngine_SSH(t *testing.T) {
WgAddr: wgaddr.MustParseWGAddress("100.64.0.1/24"),
WgPrivateKey: key,
WgPort: 33100,
IFaceBlackList: profilemanager.DefaultInterfaceBlacklist,
ServerSSHAllowed: true,
MTU: iface.DefaultMTU,
SSHKey: sshKey,
@@ -204,11 +206,12 @@ func TestEngine_Sync(t *testing.T) {
}
relayMgr := relayClient.NewManager(ctx, nil, key.PublicKey().String(), iface.DefaultMTU)
engine := NewEngine(ctx, cancel, &EngineConfig{
WgIfaceName: "utun103",
WgAddr: wgaddr.MustParseWGAddress("100.64.0.1/24"),
WgPrivateKey: key,
WgPort: 33100,
MTU: iface.DefaultMTU,
WgIfaceName: "utun103",
WgAddr: wgaddr.MustParseWGAddress("100.64.0.1/24"),
WgPrivateKey: key,
WgPort: 33100,
IFaceBlackList: profilemanager.DefaultInterfaceBlacklist,
MTU: iface.DefaultMTU,
}, EngineServices{
SignalClient: &signal.MockClient{},
MgmClient: &mgmt.MockClient{SyncFunc: syncFunc},
@@ -412,11 +415,12 @@ func createEngine(ctx context.Context, cancel context.CancelFunc, setupKey strin
wgPort := 33100 + i
conf := &EngineConfig{
WgIfaceName: ifaceName,
WgAddr: wgaddr.MustParseWGAddress(resp.PeerConfig.Address),
WgPrivateKey: key,
WgPort: wgPort,
MTU: iface.DefaultMTU,
WgIfaceName: ifaceName,
WgAddr: wgaddr.MustParseWGAddress(resp.PeerConfig.Address),
WgPrivateKey: key,
WgPort: wgPort,
IFaceBlackList: profilemanager.DefaultInterfaceBlacklist,
MTU: iface.DefaultMTU,
}
relayMgr := relayClient.NewManager(ctx, nil, key.PublicKey().String(), iface.DefaultMTU)
+1 -1
View File
@@ -6,6 +6,6 @@ import (
"github.com/netbirdio/netbird/client/internal/stdnet"
)
func (e *Engine) newStdNet() (*stdnet.Net, error) {
func (e *Engine) newStdNet() *stdnet.Net {
return stdnet.NewNet(e.clientCtx, e.config.IFaceBlackList)
}
+1 -1
View File
@@ -2,6 +2,6 @@ package internal
import "github.com/netbirdio/netbird/client/internal/stdnet"
func (e *Engine) newStdNet() (*stdnet.Net, error) {
func (e *Engine) newStdNet() *stdnet.Net {
return stdnet.NewNetWithDiscover(e.clientCtx, e.mobileDep.IFaceDiscover, e.config.IFaceBlackList)
}
+20 -9
View File
@@ -161,7 +161,6 @@ func (m *MockWGIface) GetProxy() wgproxy.Proxy {
return m.GetProxyFunc()
}
func (m *MockWGIface) GetNet() *netstack.Net {
return m.GetNetFunc()
}
@@ -689,10 +688,7 @@ func TestEngine_UpdateNetworkMapWithRoutes(t *testing.T) {
StatusRecorder: peer.NewRecorder("https://mgm"),
}, MobileDependency{})
engine.ctx = ctx
newNet, err := stdnet.NewNet(context.Background(), nil)
if err != nil {
t.Fatal(err)
}
newNet := stdnet.NewNet(context.Background(), profilemanager.DefaultInterfaceBlacklist)
opts := iface.WGIFaceOpts{
IFaceName: wgIfaceName,
@@ -897,10 +893,7 @@ func TestEngine_UpdateNetworkMapWithDNSUpdate(t *testing.T) {
}, MobileDependency{})
engine.ctx = ctx
newNet, err := stdnet.NewNet(context.Background(), nil)
if err != nil {
t.Fatal(err)
}
newNet := stdnet.NewNet(context.Background(), profilemanager.DefaultInterfaceBlacklist)
opts := iface.WGIFaceOpts{
IFaceName: wgIfaceName,
Address: wgaddr.MustParseWGAddress(wgAddr),
@@ -1500,3 +1493,21 @@ func TestOverlayAddrsFromAllowedIPs(t *testing.T) {
})
}
}
func TestEngine_SyncResponsePersistence(t *testing.T) {
e := &Engine{}
_, err := e.GetLatestSyncResponse()
require.Error(t, err, "persistence is disabled by default")
e.SetSyncResponsePersistence(true)
e.persistSyncResponse(&mgmtProto.SyncResponse{NetworkMap: &mgmtProto.NetworkMap{Serial: 7}})
got, err := e.GetLatestSyncResponse()
require.NoError(t, err)
assert.Equal(t, uint64(7), got.GetNetworkMap().GetSerial())
e.SetSyncResponsePersistence(false)
_, err = e.GetLatestSyncResponse()
require.Error(t, err)
}
+1 -4
View File
@@ -39,10 +39,7 @@ func NewAgent(ctx context.Context, iFaceDiscover stdnet.ExternalIFaceDiscover, c
iceFailedTimeout := iceFailedTimeout()
iceRelayAcceptanceMinWait := iceRelayAcceptanceMinWait()
transportNet, err := newStdNet(ctx, iFaceDiscover, config.InterfaceBlackList)
if err != nil {
log.Errorf("failed to create pion's stdnet: %s", err)
}
transportNet := newStdNet(ctx, iFaceDiscover, config.InterfaceBlackList)
fac := logging.NewDefaultLoggerFactory()
+1 -1
View File
@@ -8,6 +8,6 @@ import (
"github.com/netbirdio/netbird/client/internal/stdnet"
)
func newStdNet(ctx context.Context, _ stdnet.ExternalIFaceDiscover, ifaceBlacklist []string) (*stdnet.Net, error) {
func newStdNet(ctx context.Context, _ stdnet.ExternalIFaceDiscover, ifaceBlacklist []string) *stdnet.Net {
return stdnet.NewNet(ctx, ifaceBlacklist)
}
+1 -1
View File
@@ -6,6 +6,6 @@ import (
"github.com/netbirdio/netbird/client/internal/stdnet"
)
func newStdNet(ctx context.Context, iFaceDiscover stdnet.ExternalIFaceDiscover, ifaceBlacklist []string) (*stdnet.Net, error) {
func newStdNet(ctx context.Context, iFaceDiscover stdnet.ExternalIFaceDiscover, ifaceBlacklist []string) *stdnet.Net {
return stdnet.NewNetWithDiscover(ctx, iFaceDiscover, ifaceBlacklist)
}
+20
View File
@@ -205,6 +205,7 @@ type Status struct {
muxRelays sync.RWMutex
peers map[string]State
ipToKey map[string]string
activeRoutePeers map[route.HAUniqueID]string
changeNotify map[string]map[string]*StatusChangeSubscription // map[peerID]map[subscriptionID]*StatusChangeSubscription
signalState bool
signalError error
@@ -268,6 +269,7 @@ func NewRecorder(mgmAddress string) *Status {
return &Status{
peers: make(map[string]State),
ipToKey: make(map[string]string),
activeRoutePeers: make(map[route.HAUniqueID]string),
changeNotify: make(map[string]map[string]*StatusChangeSubscription),
eventStreams: make(map[string]chan *proto.SystemEvent),
eventQueue: NewEventQueue(eventQueueSize),
@@ -492,6 +494,24 @@ func (d *Status) RemovePeerStateRoute(peer string, route string) error {
return nil
}
func (d *Status) AddActiveRoutePeer(haID route.HAUniqueID, peer string) {
d.mux.Lock()
defer d.mux.Unlock()
d.activeRoutePeers[haID] = peer
}
func (d *Status) RemoveActiveRoutePeer(haID route.HAUniqueID) {
d.mux.Lock()
defer d.mux.Unlock()
delete(d.activeRoutePeers, haID)
}
func (d *Status) GetActiveRoutePeers() map[route.HAUniqueID]string {
d.mux.RLock()
defer d.mux.RUnlock()
return maps.Clone(d.activeRoutePeers)
}
// CheckRoutes checks if the source and destination addresses are within the same route
// and returns the resource ID of the route that contains the addresses
func (d *Status) CheckRoutes(ip netip.Addr) ([]byte, bool) {
+23
View File
@@ -9,6 +9,8 @@ import (
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"github.com/netbirdio/netbird/route"
)
func TestAddPeer(t *testing.T) {
@@ -372,3 +374,24 @@ func TestMarkServerStateDoesNotNotifyWhenUnchanged(t *testing.T) {
status.MarkManagementDisconnected(err)
assert.False(t, notified(ch), "redundant disconnect should not notify")
}
func TestActiveRoutePeers(t *testing.T) {
status := NewRecorder("https://mgm")
netA := route.HAUniqueID("net-a-10.0.0.0/24")
netB := route.HAUniqueID("net-b-10.0.0.0/24")
status.AddActiveRoutePeer(netA, "peerA")
status.AddActiveRoutePeer(netB, "peerB")
active := status.GetActiveRoutePeers()
assert.Equal(t, "peerA", active[netA])
assert.Equal(t, "peerB", active[netB])
status.RemoveActiveRoutePeer(netA)
delete(active, netB)
active = status.GetActiveRoutePeers()
_, ok := active[netA]
assert.False(t, ok)
assert.Equal(t, "peerB", active[netB])
}
+2 -10
View File
@@ -201,11 +201,7 @@ func (p *StunTurnProbe) probeSTUN(ctx context.Context, uri *stun.URI) (addr stri
}
}()
net, err := stdnet.NewNet(ctx, nil)
if err != nil {
probeErr = fmt.Errorf("new net: %w", err)
return
}
net := stdnet.NewNet(ctx, nil)
client, err := stun.DialURI(uri, &stun.DialConfig{
Net: net,
@@ -290,11 +286,7 @@ func (p *StunTurnProbe) probeTURN(ctx context.Context, uri *stun.URI) (addr stri
}
}()
net, err := stdnet.NewNet(ctx, nil)
if err != nil {
probeErr = fmt.Errorf("new net: %w", err)
return
}
net := stdnet.NewNet(ctx, nil)
cfg := &turn.ClientConfig{
STUNServerAddr: turnServerAddr,
TURNServerAddr: turnServerAddr,
@@ -294,6 +294,7 @@ func (w *Watcher) addAllowedIPs(route *route.Route) error {
return fmt.Errorf("add allowed IPs for peer %s: %w", route.Peer, err)
}
w.statusRecorder.AddActiveRoutePeer(route.GetHAUniqueID(), route.Peer)
if err := w.statusRecorder.AddPeerStateRoute(route.Peer, w.handler.String(), route.GetResourceID()); err != nil {
log.Warnf("Failed to update peer state: %v", err)
}
@@ -303,6 +304,7 @@ func (w *Watcher) addAllowedIPs(route *route.Route) error {
}
func (w *Watcher) removeAllowedIPs(route *route.Route, rsn reason) error {
w.statusRecorder.RemoveActiveRoutePeer(route.GetHAUniqueID())
if err := w.statusRecorder.RemovePeerStateRoute(route.Peer, w.handler.String()); err != nil {
log.Warnf("Failed to update peer state: %v", err)
}
+2 -4
View File
@@ -8,6 +8,7 @@ import (
"net/netip"
"testing"
"github.com/netbirdio/netbird/client/internal/profilemanager"
"github.com/netbirdio/netbird/client/internal/stdnet"
"golang.zx2c4.com/wireguard/wgctrl/wgtypes"
@@ -406,10 +407,7 @@ func TestManagerUpdateRoutes(t *testing.T) {
for n, testCase := range testCases {
t.Run(testCase.name, func(t *testing.T) {
peerPrivateKey, _ := wgtypes.GeneratePrivateKey()
newNet, err := stdnet.NewNet(context.Background(), nil)
if err != nil {
t.Fatal(err)
}
newNet := stdnet.NewNet(context.Background(), profilemanager.DefaultInterfaceBlacklist)
opts := iface.WGIFaceOpts{
IFaceName: fmt.Sprintf("utun43%d", n),
Address: wgaddr.MustParseWGAddress("100.65.65.2/24"),
@@ -15,6 +15,7 @@ import (
"syscall"
"testing"
"github.com/netbirdio/netbird/client/internal/profilemanager"
"github.com/netbirdio/netbird/client/internal/stdnet"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
@@ -436,8 +437,7 @@ func createWGInterface(t *testing.T, interfaceName, ipAddressCIDR string, listen
peerPrivateKey, err := wgtypes.GeneratePrivateKey()
require.NoError(t, err)
newNet, err := stdnet.NewNet(context.Background(), nil)
require.NoError(t, err)
newNet := stdnet.NewNet(context.Background(), profilemanager.DefaultInterfaceBlacklist)
opts := iface.WGIFaceOpts{
IFaceName: interfaceName,
+47 -38
View File
@@ -45,7 +45,7 @@ type Net struct {
}
// NewNetWithDiscover creates a new StdNet instance.
func NewNetWithDiscover(ctx context.Context, iFaceDiscover ExternalIFaceDiscover, disallowList []string) (*Net, error) {
func NewNetWithDiscover(ctx context.Context, iFaceDiscover ExternalIFaceDiscover, disallowList []string) *Net {
if ctx == nil {
ctx = context.Background()
}
@@ -60,20 +60,19 @@ func NewNetWithDiscover(ctx context.Context, iFaceDiscover ExternalIFaceDiscover
} else {
n.iFaceDiscover = newMobileIFaceDiscover(iFaceDiscover)
}
return n, n.UpdateInterfaces()
return n
}
// NewNet creates a new StdNet instance.
func NewNet(ctx context.Context, disallowList []string) (*Net, error) {
func NewNet(ctx context.Context, disallowList []string) *Net {
if ctx == nil {
ctx = context.Background()
}
n := &Net{
return &Net{
iFaceDiscover: pionDiscover{},
interfaceFilter: InterfaceFilter(disallowList),
ctx: ctx,
}
return n, n.UpdateInterfaces()
}
// resolveAddr performs DNS resolution with context support and timeout.
@@ -122,45 +121,18 @@ func (n *Net) resolveAddr(network, address string) (netip.AddrPort, error) {
return netip.AddrPortFrom(addrs[0], uint16(port)), nil
}
// UpdateInterfaces updates the internal list of network interfaces
// and associated addresses filtering them by name.
// The interfaces are discovered by an external iFaceDiscover function or by a default discoverer if the external one
// wasn't specified.
func (n *Net) UpdateInterfaces() (err error) {
n.mu.Lock()
defer n.mu.Unlock()
return n.updateInterfaces()
}
func (n *Net) updateInterfaces() (err error) {
allIfaces, err := n.iFaceDiscover.iFaces()
if err != nil {
return err
}
n.interfaces = n.filterInterfaces(allIfaces)
n.lastUpdate = time.Now()
return nil
}
// Interfaces returns a slice of interfaces which are available on the
// system
func (n *Net) Interfaces() ([]*transport.Interface, error) {
n.mu.Lock()
defer n.mu.Unlock()
if time.Since(n.lastUpdate) < updateInterval {
return slices.Clone(n.interfaces), nil
iFaces, err := n.freshInterfacesLocked()
if err != nil {
return nil, err
}
if err := n.updateInterfaces(); err != nil {
return nil, fmt.Errorf("update interfaces: %w", err)
}
return slices.Clone(n.interfaces), nil
return slices.Clone(iFaces), nil
}
// InterfaceByIndex returns the interface specified by index.
@@ -171,7 +143,13 @@ func (n *Net) Interfaces() ([]*transport.Interface, error) {
func (n *Net) InterfaceByIndex(index int) (*transport.Interface, error) {
n.mu.Lock()
defer n.mu.Unlock()
for _, ifc := range n.interfaces {
iFaces, err := n.freshInterfacesLocked()
if err != nil {
return nil, err
}
for _, ifc := range iFaces {
if ifc.Index == index {
return ifc, nil
}
@@ -184,7 +162,13 @@ func (n *Net) InterfaceByIndex(index int) (*transport.Interface, error) {
func (n *Net) InterfaceByName(name string) (*transport.Interface, error) {
n.mu.Lock()
defer n.mu.Unlock()
for _, ifc := range n.interfaces {
iFaces, err := n.freshInterfacesLocked()
if err != nil {
return nil, err
}
for _, ifc := range iFaces {
if ifc.Name == name {
return ifc, nil
}
@@ -193,6 +177,31 @@ func (n *Net) InterfaceByName(name string) (*transport.Interface, error) {
return nil, fmt.Errorf("%w: %s", transport.ErrInterfaceNotFound, name)
}
func (n *Net) freshInterfacesLocked() ([]*transport.Interface, error) {
if time.Since(n.lastUpdate) < updateInterval {
return n.interfaces, nil
}
if err := n.updateInterfacesLocked(); err != nil {
return nil, fmt.Errorf("update interfaces: %w", err)
}
return n.interfaces, nil
}
func (n *Net) updateInterfacesLocked() error {
allIFaces, err := n.iFaceDiscover.iFaces()
if err != nil {
return err
}
n.interfaces = n.filterInterfaces(allIFaces)
n.lastUpdate = time.Now()
return nil
}
func (n *Net) filterInterfaces(interfaces []*transport.Interface) []*transport.Interface {
if n.interfaceFilter == nil {
return interfaces
+136
View File
@@ -0,0 +1,136 @@
package stdnet
import (
"context"
"errors"
"net"
"testing"
"github.com/pion/transport/v3"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)
type countingDiscover struct {
calls int
list []*transport.Interface
err error
}
func (d *countingDiscover) iFaces() ([]*transport.Interface, error) {
d.calls++
if d.err != nil {
return nil, d.err
}
return d.list, nil
}
func newTestNet(t *testing.T, d iFaceDiscover) *Net {
t.Helper()
return &Net{
iFaceDiscover: d,
ctx: context.Background(),
}
}
func testIFace(index int, name string) *transport.Interface {
return transport.NewInterface(net.Interface{Index: index, Name: name})
}
func TestNet_InterfacesDiscoversLazilyAndCaches(t *testing.T) {
d := &countingDiscover{list: []*transport.Interface{testIFace(1, "eth0")}}
n := newTestNet(t, d)
require.Zero(t, d.calls, "construction must not discover interfaces")
iFaces, err := n.Interfaces()
require.NoError(t, err)
require.Len(t, iFaces, 1)
assert.Equal(t, 1, d.calls)
_, err = n.Interfaces()
require.NoError(t, err)
assert.Equal(t, 1, d.calls)
}
func TestNewNet_DoesNotDiscoverAtConstruction(t *testing.T) {
n := NewNet(context.Background(), nil)
require.NotNil(t, n)
assert.True(t, n.lastUpdate.IsZero(), "constructor must leave the cache cold")
}
func TestNewNetWithDiscover_DoesNotDiscoverAtConstruction(t *testing.T) {
n := NewNetWithDiscover(context.Background(), nil, nil)
require.NotNil(t, n)
assert.True(t, n.lastUpdate.IsZero(), "constructor must leave the cache cold")
}
func TestNet_InterfacesRetryAfterDiscoveryFailure(t *testing.T) {
discoverErr := errors.New("discover failed")
d := &countingDiscover{err: discoverErr}
n := newTestNet(t, d)
_, err := n.Interfaces()
require.ErrorIs(t, err, discoverErr)
d.err = nil
d.list = []*transport.Interface{testIFace(1, "eth0")}
iFaces, err := n.Interfaces()
require.NoError(t, err)
require.Len(t, iFaces, 1)
assert.Equal(t, 2, d.calls)
}
func TestNet_InterfaceByNameRefreshes(t *testing.T) {
d := &countingDiscover{list: []*transport.Interface{testIFace(3, "eth0")}}
n := newTestNet(t, d)
ifc, err := n.InterfaceByName("eth0")
require.NoError(t, err)
assert.Equal(t, "eth0", ifc.Name)
assert.Equal(t, 1, d.calls)
_, err = n.InterfaceByName("nope")
require.ErrorIs(t, err, transport.ErrInterfaceNotFound)
}
func TestNet_InterfaceByIndexRefreshes(t *testing.T) {
d := &countingDiscover{list: []*transport.Interface{testIFace(3, "eth0")}}
n := newTestNet(t, d)
ifc, err := n.InterfaceByIndex(3)
require.NoError(t, err)
assert.Equal(t, "eth0", ifc.Name)
assert.Equal(t, 1, d.calls)
_, err = n.InterfaceByIndex(99)
require.ErrorIs(t, err, transport.ErrInterfaceNotFound)
}
func TestNet_InterfaceLookupPropagatesDiscoveryError(t *testing.T) {
discoverErr := errors.New("discover failed")
n := newTestNet(t, &countingDiscover{err: discoverErr})
_, err := n.InterfaceByName("eth0")
require.ErrorIs(t, err, discoverErr)
_, err = n.InterfaceByIndex(1)
require.ErrorIs(t, err, discoverErr)
}
func TestNet_InterfacesReturnsCopy(t *testing.T) {
d := &countingDiscover{list: []*transport.Interface{testIFace(1, "eth0")}}
n := newTestNet(t, d)
iFaces, err := n.Interfaces()
require.NoError(t, err)
require.Len(t, iFaces, 1)
iFaces[0] = testIFace(2, "tampered")
iFaces, err = n.Interfaces()
require.NoError(t, err)
require.Len(t, iFaces, 1)
assert.Equal(t, "eth0", iFaces[0].Name)
}
@@ -1,4 +1,5 @@
import { useEffect, useRef } from "react";
import { useSearchParams } from "react-router-dom";
import { Events } from "@wailsio/runtime";
import { useStatus } from "@/contexts/StatusContext.tsx";
@@ -6,13 +7,15 @@ const EVENT_WINDOW_PAINTED = "netbird:window-painted";
export const ReadySignal = () => {
const { isReady } = useStatus();
const sent = useRef(false);
const [params] = useSearchParams();
const generation = params.get("gen") ?? "";
const sent = useRef<string | null>(null);
useEffect(() => {
if (!isReady || sent.current) return;
sent.current = true;
void Events.Emit(EVENT_WINDOW_PAINTED);
}, [isReady]);
if (!isReady || sent.current === generation) return;
sent.current = generation;
void Events.Emit(EVENT_WINDOW_PAINTED, generation);
}, [isReady, generation]);
return null;
};
@@ -1,23 +1,27 @@
import { useLayoutEffect, useRef } from "react";
import { Window } from "@wailsio/runtime";
import { useSearchParams } from "react-router-dom";
import { Events, Window } from "@wailsio/runtime";
import i18next from "@/lib/i18n";
import { isLinux } from "@/lib/platform";
const EVENT_WINDOW_PAINTED = "netbird:window-painted";
// Sizes the current Wails window to the measured content height (keeping `width`),
// then shows it. Re-applies on content resize and language change.
// then reports it as painted so Go shows it. Re-applies on content resize and language change.
export function useAutoSizeWindow<T extends HTMLElement>(width: number, ready: boolean = true) {
const ref = useRef<T | null>(null);
const [params] = useSearchParams();
const generation = params.get("gen") ?? "";
useLayoutEffect(() => {
const el = ref.current;
if (!el) return;
let shown = false;
let painted = false;
let raf1 = 0;
let raf2 = 0;
const showOnce = () => {
if (shown) return;
shown = true;
Window.Show().catch(() => {});
Window.Focus().catch(() => {});
const paintedOnce = () => {
if (painted) return;
painted = true;
Events.Emit(EVENT_WINDOW_PAINTED, generation).catch(() => {});
};
const apply = async () => {
if (!ready) return;
@@ -33,7 +37,7 @@ export function useAutoSizeWindow<T extends HTMLElement>(width: number, ready: b
await Window.SetMaxSize(width, targetH);
}
await Window.SetSize(width, targetH);
showOnce();
paintedOnce();
} catch {
// window gone / not ready — ignore
}
@@ -55,6 +59,6 @@ export function useAutoSizeWindow<T extends HTMLElement>(width: number, ready: b
cancelAnimationFrame(raf2);
i18next.off("languageChanged", scheduleApply);
};
}, [width, ready]);
}, [width, ready, generation]);
return ref;
}
@@ -1,4 +1,4 @@
import { useCallback, useEffect, useRef } from "react";
import { useCallback } from "react";
import { useTranslation } from "react-i18next";
import { useSearchParams } from "react-router-dom";
import { Events } from "@wailsio/runtime";
@@ -21,7 +21,6 @@ export default function LoginWaitingForBrowserDialog() {
const [params] = useSearchParams();
const uri = params.get("uri") ?? "";
const contentRef = useAutoSizeWindow<HTMLDivElement>(WINDOW_WIDTH);
const openedRef = useRef(false);
const reportOpenFailure = useCallback(
(e: unknown) => {
@@ -33,13 +32,6 @@ export default function LoginWaitingForBrowserDialog() {
[t],
);
// Open the browser only after mount, or it lands on top of the still-hidden popup.
useEffect(() => {
if (!uri || openedRef.current) return;
openedRef.current = true;
Connection.OpenURL(uri).catch(reportOpenFailure);
}, [uri, reportOpenFailure]);
const tryAgain = useCallback(() => {
if (!uri) return;
Connection.OpenURL(uri).catch(reportOpenFailure);
+17 -13
View File
@@ -205,19 +205,7 @@ func (s *Connection) Down(ctx context.Context) error {
// window.open, so the SSO verification page can't pop inline. Honors $BROWSER
// before the platform default.
func (s *Connection) OpenURL(url string) error {
if browser := os.Getenv("BROWSER"); browser != "" {
return exec.Command(browser, url).Start()
}
switch runtime.GOOS {
case "windows":
return exec.Command("rundll32", "url.dll,FileProtocolHandler", url).Start()
case "darwin":
return exec.Command("open", url).Start()
case "linux":
return exec.Command("xdg-open", url).Start()
default:
return fmt.Errorf("unsupported platform")
}
return openURL(url)
}
func (s *Connection) Logout(ctx context.Context, p LogoutParams) error {
@@ -288,3 +276,19 @@ func (s *Connection) waitSSOLogin(ctx context.Context, p WaitSSOParams) (string,
func (s *Connection) classifyDaemonError(err error) *ClientError {
return s.classifier.classify(err)
}
func openURL(url string) error {
if browser := os.Getenv("BROWSER"); browser != "" {
return exec.Command(browser, url).Start()
}
switch runtime.GOOS {
case "windows":
return exec.Command("rundll32", "url.dll,FileProtocolHandler", url).Start()
case "darwin":
return exec.Command("open", url).Start()
case "linux":
return exec.Command("xdg-open", url).Start()
default:
return fmt.Errorf("unsupported platform")
}
}
+364 -120
View File
@@ -5,6 +5,7 @@ package services
import (
"net/url"
"strconv"
"strings"
"sync"
"sync/atomic"
"time"
@@ -26,6 +27,16 @@ type windowOp func(w *application.WebviewWindow, created bool)
type windowCloser func(w *application.WebviewWindow)
// hideableWindow is the slice of application.Window the hide/restore bookkeeping needs.
// Narrow enough to fake in tests, which application.Window itself is not: it carries
// unexported methods.
type hideableWindow interface {
Show() application.Window
Hide() application.Window
IsVisible() bool
Name() string
}
// EventTriggerLogin asks the frontend's startLogin() to begin an SSO flow.
const EventTriggerLogin = "trigger-login"
@@ -37,7 +48,10 @@ const EventSettingsOpen = "netbird:settings:open"
const EventWindowPainted = "netbird:window-painted"
const paintedFallback = 2 * time.Second
// generationParam carries the painted-report token in each dialog's start URL.
const generationParam = "gen"
const paintedFallback = 3 * time.Second
const headlessTeardownDelay = 2 * time.Second
@@ -201,6 +215,12 @@ func DialogWindowOptions(name, title, url string, linuxIcon []byte) application.
}
}
// hiddenWindow records a window hidden by owner, the name of the popup that hid it.
type hiddenWindow struct {
win hideableWindow
owner string
}
type WindowManager struct {
app *application.App
mainWindow *application.WebviewWindow
@@ -213,19 +233,35 @@ type WindowManager struct {
installProgress *application.WebviewWindow
welcome *application.WebviewWindow
errorDialog *application.WebviewWindow
// hiddenForLogin holds windows hidden while the BrowserLogin popup is open, restored on close.
hiddenForLogin []application.Window
mu sync.Mutex
newMain func(startURL string) *application.WebviewWindow
creating map[string]bool
pendingOps map[string][]windowOp
pendingClose map[string]windowCloser
restoreGen uint64
ready map[uint]bool
// hiddenWindows holds windows hidden while a popup owns the screen, each tagged with
// the popup that hid it so closing one popup cannot restore what another still hides.
hiddenWindows []hiddenWindow
hiding map[string]bool
// allWindows and raiseMain are the seams the hide/restore tests replace; both are nil
// in production, where the Wails app and the platform helper are used directly.
allWindows func() []hideableWindow
raiseMain func()
mu sync.Mutex
newMain func(startURL string) *application.WebviewWindow
creating map[string]bool
pendingOps map[string][]windowOp
pendingClose map[string]windowCloser
restoreGen map[string]uint64
// painted gates showing a window: set by the frontend's first render, or by the
// fallback timer so a webview that never wakes up still becomes visible.
painted map[uint]bool
// mounted gates emitting to a window: set only by a real frontend report, since an
// event emitted to a frontend that has not subscribed yet is dropped, not queued.
mounted map[uint]bool
showPending map[uint]bool
pendingTab map[uint]string
pendingEmits map[uint][]string
fallbackTimers map[uint]*time.Timer
afterShow map[uint]func()
// generation maps a window name to the token stamped into its current start URL, so a
// painted report from a replaced window can be told apart from the live one's.
generation map[string]uint64
lastGeneration uint64
headlessMain bool
headlessTimer *time.Timer
// recenterOnShow is set only on the minimal-WM/XEmbed path, where the WM neither centers nor
@@ -243,11 +279,16 @@ func NewWindowManager(app *application.App, mainWindow *application.WebviewWindo
creating: map[string]bool{},
pendingOps: map[string][]windowOp{},
pendingClose: map[string]windowCloser{},
ready: map[uint]bool{},
restoreGen: map[string]uint64{},
hiding: map[string]bool{},
painted: map[uint]bool{},
mounted: map[uint]bool{},
showPending: map[uint]bool{},
pendingTab: map[uint]string{},
pendingEmits: map[uint][]string{},
fallbackTimers: map[uint]*time.Timer{},
afterShow: map[uint]func(){},
generation: map[string]uint64{},
}
s.watchPainted()
s.watchTriggerLogin()
@@ -307,13 +348,13 @@ func (s *WindowManager) OpenSettings(tab string) {
s.withWindow(windowSettings, &s.settings, s.newSettingsWindow, func(w *application.WebviewWindow, _ bool) {
s.mu.Lock()
ready := s.ready[w.ID()]
if !ready {
mounted := s.mounted[w.ID()]
if !mounted {
s.pendingTab[w.ID()] = target
}
s.mu.Unlock()
if ready {
if mounted {
s.app.Event.Emit(EventSettingsOpen, target)
}
s.showWhenReady(w)
@@ -327,21 +368,37 @@ func (s *WindowManager) OpenBrowserLogin(uri string) {
startURL = "/#/dialog/browser-login?uri=" + url.QueryEscape(uri)
}
s.withWindow(windowBrowserLogin, &s.browserLogin, func() *application.WebviewWindow {
return s.newBrowserLoginWindow(startURL)
return s.newBrowserLoginWindow(s.stampGeneration(windowBrowserLogin, startURL))
}, func(w *application.WebviewWindow, created bool) {
if created {
s.centerOnCursorScreen(w)
return
}
if uri != "" {
w.SetURL(startURL)
if !created && uri != "" {
w.SetURL(s.stampGeneration(windowBrowserLogin, startURL))
}
s.centerOnCursorScreen(w)
w.Show()
w.Focus()
s.showThenOpenBrowser(w, uri)
})
}
func (s *WindowManager) showThenOpenBrowser(w *application.WebviewWindow, uri string) {
if uri != "" {
s.mu.Lock()
s.afterShow[w.ID()] = func() { s.openBrowser(uri) }
s.mu.Unlock()
}
s.showWhenReady(w)
}
func (s *WindowManager) openBrowser(uri string) {
if uri == "" {
return
}
go func() {
if err := openURL(uri); err != nil {
log.Errorf("open browser for SSO login: %v", err)
s.OpenError(s.title("browserLogin.openFailedTitle"), err.Error(), "")
}
}()
}
func (s *WindowManager) newBrowserLoginWindow(startURL string) *application.WebviewWindow {
s.hideOtherWindows(windowBrowserLogin)
opts := DialogWindowOptions(windowBrowserLogin, s.title("window.title.signIn"), startURL, s.linuxIcon)
@@ -360,12 +417,14 @@ func (s *WindowManager) newBrowserLoginWindow(startURL string) *application.Webv
if userClosed {
s.browserLogin = nil
}
s.forgetWindowLocked(w)
s.mu.Unlock()
if userClosed {
s.restoreHiddenWindows()
s.restoreHiddenWindows(windowBrowserLogin)
s.app.Event.Emit(EventBrowserLoginCancel)
}
})
s.armReady(w)
return w
}
@@ -386,13 +445,11 @@ func (s *WindowManager) InstallProgressWindow() *application.WebviewWindow {
}
func (s *WindowManager) CloseBrowserLogin() {
// The WindowClosing hook no-ops on a programmatic close, so restore here —
// but only if a popup was actually open. The frontend calls this even when no
// popup was ever shown (e.g. resetDialog() after an early RequestExtend failure,
// or connection.ts's catch path), and hiddenForLogin is shared with
// OpenInstallProgress, so an unconditional restore could re-show windows a
// still-running install-progress is hiding.
s.closeWindow(windowBrowserLogin, &s.browserLogin, s.restoreAndClose)
// The WindowClosing hook no-ops on a programmatic close, so the closer restores.
// The frontend calls this even when no popup was ever shown (resetDialog() after an
// early RequestExtend failure, or connection.ts's catch path); closeWindow skips the
// closer then, and an owner-scoped restore cannot touch what install-progress hides.
s.closeWindow(windowBrowserLogin, &s.browserLogin, s.restoringCloser(windowBrowserLogin))
}
// OpenSessionExpiration shows the countdown warning on the cursor's display; seconds seeds
@@ -404,16 +461,13 @@ func (s *WindowManager) OpenSessionExpiration(seconds int, deadlineUnixMilli int
startURL += "&deadline=" + strconv.FormatInt(deadlineUnixMilli, 10)
}
s.withWindow(windowSessionExpiration, &s.sessionExpiration, func() *application.WebviewWindow {
return s.newSessionExpirationWindow(startURL)
return s.newSessionExpirationWindow(s.stampGeneration(windowSessionExpiration, startURL))
}, func(w *application.WebviewWindow, created bool) {
if created {
s.centerOnCursorScreen(w)
return
if !created {
w.SetURL(s.stampGeneration(windowSessionExpiration, startURL))
}
w.SetURL(startURL)
s.centerOnCursorScreen(w)
w.Show()
w.Focus()
s.showWhenReady(w)
})
}
@@ -427,8 +481,10 @@ func (s *WindowManager) newSessionExpirationWindow(startURL string) *application
if s.sessionExpiration == w {
s.sessionExpiration = nil
}
s.forgetWindowLocked(w)
s.mu.Unlock()
})
s.armReady(w)
return w
}
@@ -440,20 +496,20 @@ func (s *WindowManager) CloseSessionExpiration() {
// closes the browser-login popup and the session-expiration window together.
func (s *WindowManager) CloseRenewFlow() {
s.mu.Lock()
bl := s.takeWindowLocked(windowBrowserLogin, &s.browserLogin, s.restoreAndClose)
bl := s.takeWindowLocked(windowBrowserLogin, &s.browserLogin, s.restoringCloser(windowBrowserLogin))
se := s.takeWindowLocked(windowSessionExpiration, &s.sessionExpiration, closeOnly)
if se != nil {
kept := s.hiddenForLogin[:0]
for _, w := range s.hiddenForLogin {
if w != se {
kept = append(kept, w)
kept := s.hiddenWindows[:0]
for _, hidden := range s.hiddenWindows {
if !sameWindow(hidden.win, se) {
kept = append(kept, hidden)
}
}
s.hiddenForLogin = kept
s.hiddenWindows = kept
}
s.mu.Unlock()
s.restoreHiddenWindows()
s.restoreHiddenWindows(windowBrowserLogin)
// Close after unlock so the re-entrant handlers can take s.mu.
if bl != nil {
bl.Close()
@@ -471,14 +527,12 @@ func (s *WindowManager) OpenInstallProgress(version string) {
startURL = "/#/dialog/install-progress?version=" + url.QueryEscape(version)
}
s.withWindow(windowInstallProgress, &s.installProgress, func() *application.WebviewWindow {
return s.newInstallProgressWindow(startURL)
return s.newInstallProgressWindow(s.stampGeneration(windowInstallProgress, startURL))
}, func(w *application.WebviewWindow, created bool) {
if !created {
w.SetURL(startURL)
w.Show()
w.Focus()
w.SetURL(s.stampGeneration(windowInstallProgress, startURL))
}
s.centerWhenReady(w)
s.showWhenReady(w)
})
}
@@ -489,32 +543,33 @@ func (s *WindowManager) newInstallProgressWindow(startURL string) *application.W
)
w.OnWindowEvent(events.Common.WindowClosing, func(_ *application.WindowEvent) {
s.mu.Lock()
if s.installProgress == w {
userClosed := s.installProgress == w
if userClosed {
s.installProgress = nil
}
s.forgetWindowLocked(w)
s.mu.Unlock()
s.restoreHiddenWindows()
if userClosed {
s.restoreHiddenWindows(windowInstallProgress)
}
})
s.armReady(w)
return w
}
func (s *WindowManager) CloseInstallProgress() {
s.closeWindow(windowInstallProgress, &s.installProgress, closeOnly)
s.closeWindow(windowInstallProgress, &s.installProgress, s.restoringCloser(windowInstallProgress))
}
// OpenWelcome shows the first-launch onboarding window. Singleton, destroyed on close.
func (s *WindowManager) OpenWelcome() {
s.withWindow(windowWelcome, &s.welcome, s.newWelcomeWindow, func(w *application.WebviewWindow, created bool) {
if !created {
w.Show()
w.Focus()
}
s.centerWhenReady(w)
s.withWindow(windowWelcome, &s.welcome, s.newWelcomeWindow, func(w *application.WebviewWindow, _ bool) {
s.showWhenReady(w)
})
}
func (s *WindowManager) newWelcomeWindow() *application.WebviewWindow {
opts := DialogWindowOptions(windowWelcome, s.title("window.title.welcome"), "/#/dialog/welcome", s.linuxIcon)
opts := DialogWindowOptions(windowWelcome, s.title("window.title.welcome"), s.stampGeneration(windowWelcome, "/#/dialog/welcome"), s.linuxIcon)
opts.Width = 420
opts.InitialPosition = application.WindowCentered
w := s.app.Window.NewWithOptions(opts)
@@ -523,8 +578,10 @@ func (s *WindowManager) newWelcomeWindow() *application.WebviewWindow {
if s.welcome == w {
s.welcome = nil
}
s.forgetWindowLocked(w)
s.mu.Unlock()
})
s.armReady(w)
return w
}
@@ -542,14 +599,12 @@ func (s *WindowManager) OpenError(title, message, command string) {
}
startURL := errorDialogURL(title, message, command)
s.withWindow(windowError, &s.errorDialog, func() *application.WebviewWindow {
return s.newErrorWindow(startURL)
return s.newErrorWindow(s.stampGeneration(windowError, startURL))
}, func(w *application.WebviewWindow, created bool) {
if !created {
w.SetURL(startURL)
w.Show()
w.Focus()
w.SetURL(s.stampGeneration(windowError, startURL))
}
s.centerWhenReady(w)
s.showWhenReady(w)
})
}
@@ -562,8 +617,10 @@ func (s *WindowManager) newErrorWindow(startURL string) *application.WebviewWind
if s.errorDialog == w {
s.errorDialog = nil
}
s.forgetWindowLocked(w)
s.mu.Unlock()
})
s.armReady(w)
return w
}
@@ -589,14 +646,14 @@ func (s *WindowManager) ShowMainAndEmit(event string) {
s.ensureMain("/", func(w *application.WebviewWindow, _ bool) {
id := w.ID()
s.mu.Lock()
ready := s.ready[id]
if !ready {
mounted := s.mounted[id]
if !mounted {
s.pendingEmits[id] = append(s.pendingEmits[id], event)
}
s.mu.Unlock()
s.showWhenReady(w)
if ready {
if mounted {
s.app.Event.Emit(event)
}
})
@@ -741,31 +798,66 @@ func (s *WindowManager) releaseCreationLocked(name string) {
delete(s.pendingClose, name)
}
func (s *WindowManager) restoreAndClose(w *application.WebviewWindow) {
s.restoreHiddenWindows()
w.Close()
func (s *WindowManager) restoringCloser(owner string) windowCloser {
return func(w *application.WebviewWindow) {
s.restoreHiddenWindows(owner)
w.Close()
}
}
// armReady starts the fallback that shows w even if its frontend never reports a first
// render. The timer starts at creation, because a hidden webview can be suspended before
// it reaches WindowRuntimeReady — the very case this fallback covers. That makes the first
// budget cover webview boot as well, so the runtime-ready hook rearms it to give the
// frontend its own full budget to mount and paint.
func (s *WindowManager) armReady(w *application.WebviewWindow) {
if w == nil {
return
}
s.armPaintedFallback(w)
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()
s.armPaintedFallback(w)
})
}
func (s *WindowManager) armPaintedFallback(w *application.WebviewWindow) {
id := w.ID()
timer := time.AfterFunc(paintedFallback, func() {
s.mu.Lock()
painted := s.painted[id]
s.mu.Unlock()
if painted {
return
}
log.Warnf("window %q never reported a first render, showing it anyway", w.Name())
s.markPainted(w)
})
s.mu.Lock()
if prev := s.fallbackTimers[id]; prev != nil {
prev.Stop()
}
if s.painted[id] {
timer.Stop()
delete(s.fallbackTimers, id)
} else {
s.fallbackTimers[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)
w := s.windowByName(e.Sender)
if w == nil {
return
}
if !s.matchesGeneration(e.Sender, paintedGeneration(e.Data)) {
log.Debugf("ignoring stale painted report for window %q", e.Sender)
return
}
s.markPainted(w)
s.markMounted(w)
})
}
@@ -777,7 +869,7 @@ func (s *WindowManager) watchTriggerLogin() {
s.headlessTimer = nil
}
w := s.mainWindow
ready := w != nil && s.ready[w.ID()]
ready := w != nil && s.mounted[w.ID()]
s.mu.Unlock()
if ready {
return
@@ -788,7 +880,7 @@ func (s *WindowManager) watchTriggerLogin() {
if created {
s.headlessMain = true
}
pending := !s.ready[w.ID()]
pending := !s.mounted[w.ID()]
if pending {
s.pendingEmits[w.ID()] = append(s.pendingEmits[w.ID()], EventTriggerLogin)
}
@@ -850,18 +942,67 @@ func (s *WindowManager) forgetWindowLocked(w *application.WebviewWindow) {
timer.Stop()
}
delete(s.fallbackTimers, id)
delete(s.ready, id)
delete(s.painted, id)
delete(s.mounted, id)
delete(s.showPending, id)
delete(s.pendingTab, id)
delete(s.pendingEmits, id)
delete(s.afterShow, id)
kept := s.hiddenForLogin[:0]
for _, hidden := range s.hiddenForLogin {
if hidden != application.Window(w) {
kept := s.hiddenWindows[:0]
for _, hidden := range s.hiddenWindows {
if !sameWindow(hidden.win, w) {
kept = append(kept, hidden)
}
}
s.hiddenForLogin = kept
s.hiddenWindows = kept
}
func (s *WindowManager) stampGeneration(name, startURL string) string {
s.mu.Lock()
defer s.mu.Unlock()
s.lastGeneration++
s.generation[name] = s.lastGeneration
return appendGeneration(startURL, s.lastGeneration)
}
func (s *WindowManager) matchesGeneration(name string, gen uint64) bool {
s.mu.Lock()
defer s.mu.Unlock()
want, tracked := s.generation[name]
if !tracked {
return true
}
return want == gen
}
func (s *WindowManager) hideableWindows() []hideableWindow {
if s.allWindows != nil {
return s.allWindows()
}
all := s.app.Window.GetAll()
windows := make([]hideableWindow, 0, len(all))
for _, w := range all {
windows = append(windows, w)
}
return windows
}
func (s *WindowManager) isMainWindow(w hideableWindow, mainWindow *application.WebviewWindow) bool {
if s.allWindows != nil {
return w != nil && w.Name() == windowMain
}
return sameWindow(w, mainWindow)
}
func (s *WindowManager) raiseMainWindow(mainWindow *application.WebviewWindow) {
if s.raiseMain != nil {
s.raiseMain()
return
}
if mainWindow != nil {
raiseToForeground(mainWindow)
}
}
func (s *WindowManager) windowByName(name string) *application.WebviewWindow {
@@ -872,24 +1013,50 @@ func (s *WindowManager) windowByName(name string) *application.WebviewWindow {
return s.mainWindow
case windowSettings:
return s.settings
case windowBrowserLogin:
return s.browserLogin
case windowSessionExpiration:
return s.sessionExpiration
case windowInstallProgress:
return s.installProgress
case windowWelcome:
return s.welcome
case windowError:
return s.errorDialog
default:
return nil
}
}
func (s *WindowManager) markReady(w *application.WebviewWindow) {
func (s *WindowManager) markPainted(w *application.WebviewWindow) {
id := w.ID()
s.mu.Lock()
already := s.ready[id]
s.ready[id] = true
already := s.painted[id]
s.painted[id] = true
wanted := s.showPending[id]
tab, hasTab := s.pendingTab[id]
emits := s.pendingEmits[id]
delete(s.showPending, id)
if timer := s.fallbackTimers[id]; timer != nil {
timer.Stop()
delete(s.fallbackTimers, id)
}
delete(s.showPending, id)
s.mu.Unlock()
if already || !wanted {
return
}
s.showNow(w)
}
// markMounted records that the window's frontend is subscribed, and flushes the events
// held back for it. The fallback timer never calls this: showing a blank window is
// recoverable, emitting into a frontend that cannot hear it is not.
func (s *WindowManager) markMounted(w *application.WebviewWindow) {
id := w.ID()
s.mu.Lock()
already := s.mounted[id]
s.mounted[id] = true
tab, hasTab := s.pendingTab[id]
emits := s.pendingEmits[id]
delete(s.pendingTab, id)
delete(s.pendingEmits, id)
s.mu.Unlock()
@@ -902,10 +1069,6 @@ func (s *WindowManager) markReady(w *application.WebviewWindow) {
s.app.Event.Emit(EventSettingsOpen, tab)
}
if wanted {
s.showNow(w)
}
for _, event := range emits {
s.app.Event.Emit(event)
}
@@ -918,18 +1081,19 @@ func (s *WindowManager) showWhenReady(w *application.WebviewWindow) {
id := w.ID()
s.mu.Lock()
ready := s.ready[id]
if !ready {
painted := s.painted[id]
if !painted {
s.showPending[id] = true
}
s.mu.Unlock()
if ready {
if painted {
s.showNow(w)
}
}
func (s *WindowManager) showNow(w *application.WebviewWindow) {
id := w.ID()
s.mu.Lock()
if w == s.mainWindow {
s.headlessMain = false
@@ -938,10 +1102,15 @@ func (s *WindowManager) showNow(w *application.WebviewWindow) {
s.headlessTimer = nil
}
}
after := s.afterShow[id]
delete(s.afterShow, id)
s.mu.Unlock()
w.Show()
w.Focus()
s.centerWhenReady(w)
if after != nil {
after()
}
}
func (s *WindowManager) ShowMainAt(url string) {
@@ -1070,13 +1239,19 @@ func (s *WindowManager) retitleAll() {
}
}
// hideOtherWindows hides every visible window except keepName, recording them against
// keepName so only its own restore brings them back. A window already hidden by an
// earlier popup is skipped, leaving it tagged to the popup that actually hid it. The
// per-owner generation catches a restore for keepName that ran between the snapshot and
// the record, in which case the windows are re-shown rather than stranded.
func (s *WindowManager) hideOtherWindows(keepName string) {
s.mu.Lock()
gen := s.restoreGen
s.hiding[keepName] = true
gen := s.restoreGen[keepName]
s.mu.Unlock()
var hidden []application.Window
for _, w := range s.app.Window.GetAll() {
var hidden []hideableWindow
for _, w := range s.hideableWindows() {
if w == nil || w.Name() == keepName || !w.IsVisible() {
continue
}
@@ -1088,9 +1263,11 @@ func (s *WindowManager) hideOtherWindows(keepName string) {
}
s.mu.Lock()
restored := s.restoreGen != gen
restored := s.restoreGen[keepName] != gen
if !restored {
s.hiddenForLogin = append(s.hiddenForLogin, hidden...)
for _, w := range hidden {
s.hiddenWindows = append(s.hiddenWindows, hiddenWindow{win: w, owner: keepName})
}
}
s.mu.Unlock()
if !restored {
@@ -1101,33 +1278,58 @@ func (s *WindowManager) hideOtherWindows(keepName string) {
}
}
// restoreHiddenWindows re-shows windows hidden by hideOtherWindows. If the main
// window was among them, raiseToForeground lifts it above the SSO browser, which
// still owns the foreground — a plain Show/Focus would be demoted to a taskbar
// flash and leave it stranded behind.
func (s *WindowManager) restoreHiddenWindows() {
// restoreHiddenWindows re-shows the windows owner hid, unless another popup still covers
// them, in which case they are handed to that popup. If the main window was among them,
// raiseToForeground lifts it above the SSO browser, which still owns the foreground — a
// plain Show/Focus would be demoted to a taskbar flash and leave it stranded behind.
func (s *WindowManager) restoreHiddenWindows(owner string) {
s.mu.Lock()
hidden := s.hiddenForLogin
s.hiddenForLogin = nil
s.restoreGen++
mainWindow := s.mainWindow
delete(s.hiding, owner)
var restore []hideableWindow
kept := s.hiddenWindows[:0]
for _, hidden := range s.hiddenWindows {
if hidden.owner != owner {
kept = append(kept, hidden)
continue
}
if coverer, covered := s.coveringPopupLocked(hidden.win); covered {
hidden.owner = coverer
kept = append(kept, hidden)
continue
}
if hidden.win != nil {
restore = append(restore, hidden.win)
}
}
s.hiddenWindows = kept
s.restoreGen[owner]++
s.mu.Unlock()
mainRestored := false
for _, w := range hidden {
if w == nil {
continue
}
for _, w := range restore {
w.Show()
if w == mainWindow {
if s.isMainWindow(w, mainWindow) {
mainRestored = true
}
}
if mainRestored && mainWindow != nil {
raiseToForeground(mainWindow)
if mainRestored {
s.raiseMainWindow(mainWindow)
}
}
func (s *WindowManager) coveringPopupLocked(w hideableWindow) (string, bool) {
if w == nil {
return "", false
}
for name := range s.hiding {
if name != w.Name() {
return name, true
}
}
return "", false
}
// getScreenBasedOnCursorPosition returns the cursor's display, falling back to the
// main-window screen, then nil (OS-default placement).
func (s *WindowManager) getScreenBasedOnCursorPosition() *application.Screen {
@@ -1169,6 +1371,48 @@ func errorDialogURL(title, message, command string) string {
return startURL
}
// appendGeneration adds the painted-report token to a dialog start URL, keeping any
// existing query params intact across the "/#/path?params" hash-router form.
func appendGeneration(startURL string, gen uint64) string {
sep := "?"
if strings.Contains(startURL, "?") {
sep = "&"
}
return startURL + sep + generationParam + "=" + strconv.FormatUint(gen, 10)
}
// paintedGeneration reads the token a painted report carries back, returning 0 when the
// frontend sent none (an older bundle, or the main window, which is never stamped).
func paintedGeneration(data any) uint64 {
switch v := data.(type) {
case string:
gen, err := strconv.ParseUint(v, 10, 64)
if err != nil {
return 0
}
return gen
case float64:
return uint64(v)
case []any:
if len(v) == 0 {
return 0
}
return paintedGeneration(v[0])
default:
return 0
}
}
// sameWindow reports whether a hidden entry refers to w, comparing through the interface
// so a nil entry never matches a live window.
func sameWindow(hidden hideableWindow, w *application.WebviewWindow) bool {
if hidden == nil || w == nil {
return false
}
other, ok := hidden.(*application.WebviewWindow)
return ok && other == w
}
// u32ptr returns a pointer to v, for the optional *uint32 Wails theme fields.
func u32ptr(v uint32) *uint32 { return &v }
+291 -2
View File
@@ -17,9 +17,65 @@ func newTestWindowManager() *WindowManager {
creating: map[string]bool{},
pendingOps: map[string][]windowOp{},
pendingClose: map[string]windowCloser{},
restoreGen: map[string]uint64{},
hiding: map[string]bool{},
generation: map[string]uint64{},
}
}
type fakeWindow struct {
name string
visible bool
shown int
hidden int
}
func newFakeWindow(name string) *fakeWindow {
return &fakeWindow{name: name, visible: true}
}
func (f *fakeWindow) Show() application.Window {
f.visible = true
f.shown++
return nil
}
func (f *fakeWindow) Hide() application.Window {
f.visible = false
f.hidden++
return nil
}
func (f *fakeWindow) IsVisible() bool { return f.visible }
func (f *fakeWindow) Name() string { return f.name }
type fakeDesktop struct {
windows []*fakeWindow
raised int
}
func newFakeDesktop(s *WindowManager, windows ...*fakeWindow) *fakeDesktop {
d := &fakeDesktop{windows: windows}
s.allWindows = func() []hideableWindow {
all := make([]hideableWindow, 0, len(d.windows))
for _, w := range d.windows {
all = append(all, w)
}
return all
}
s.raiseMain = func() { d.raised++ }
return d
}
func ownersOf(hidden []hiddenWindow) []string {
owners := make([]string, 0, len(hidden))
for _, h := range hidden {
owners = append(owners, h.owner)
}
return owners
}
func waitDone(t *testing.T, done <-chan struct{}, msg string) {
t.Helper()
select {
@@ -339,12 +395,245 @@ func TestCloseRenewFlowDuringBrowserLoginCreationRestoresHiddenWindows(t *testin
// Seeded after the call so the deferred closer, not CloseRenewFlow's own
// immediate restore, is what has to drain it. A nil entry is skipped by
// restoreHiddenWindows, so no Wails window is needed.
s.hiddenForLogin = []application.Window{nil}
s.hiddenWindows = []hiddenWindow{{owner: windowBrowserLogin}}
return &application.WebviewWindow{}
}, func(*application.WebviewWindow, bool) {})
require.Nil(t, s.browserLogin)
require.Empty(t, s.hiddenForLogin)
require.Empty(t, s.hiddenWindows)
require.Empty(t, s.creating)
require.Empty(t, s.pendingClose)
}
func TestHideOtherWindowsSkipsKeepNameAndInvisible(t *testing.T) {
main := newFakeWindow(windowMain)
settings := newFakeWindow(windowSettings)
settings.visible = false
popup := newFakeWindow(windowBrowserLogin)
s := newTestWindowManager()
newFakeDesktop(s, main, settings, popup)
s.hideOtherWindows(windowBrowserLogin)
require.False(t, main.visible)
require.Equal(t, 1, main.hidden)
require.Equal(t, 0, settings.hidden, "an already hidden window must not be recorded")
require.Equal(t, 0, popup.hidden, "the popup itself must stay visible")
require.Equal(t, []string{windowBrowserLogin}, ownersOf(s.hiddenWindows))
}
func TestInstallDuringLoginKeepsMainHiddenUntilLoginCloses(t *testing.T) {
main := newFakeWindow(windowMain)
login := newFakeWindow(windowBrowserLogin)
install := newFakeWindow(windowInstallProgress)
install.visible = false
s := newTestWindowManager()
d := newFakeDesktop(s, main, login, install)
s.hideOtherWindows(windowBrowserLogin)
require.False(t, main.visible)
install.visible = true
s.hideOtherWindows(windowInstallProgress)
require.False(t, login.visible, "the install popup hides the login popup")
s.restoreHiddenWindows(windowInstallProgress)
require.True(t, login.visible, "the install popup restores the login popup it hid")
require.False(t, main.visible, "the main window stays hidden for the login popup")
require.Equal(t, 0, d.raised)
require.Equal(t, []string{windowBrowserLogin}, ownersOf(s.hiddenWindows))
s.restoreHiddenWindows(windowBrowserLogin)
require.True(t, main.visible)
require.Equal(t, 1, d.raised, "restoring the main window raises it above the SSO browser")
require.Empty(t, s.hiddenWindows)
}
func TestLoginClosingUnderInstallHandsMainToInstall(t *testing.T) {
main := newFakeWindow(windowMain)
login := newFakeWindow(windowBrowserLogin)
install := newFakeWindow(windowInstallProgress)
install.visible = false
s := newTestWindowManager()
d := newFakeDesktop(s, main, login, install)
s.hideOtherWindows(windowBrowserLogin)
install.visible = true
s.hideOtherWindows(windowInstallProgress)
require.False(t, main.visible)
require.False(t, login.visible, "the install popup hides the login popup")
// The login popup closes while the install popup is still up: the main window it
// hid must not resurface under the install popup, it is handed over instead.
s.restoreHiddenWindows(windowBrowserLogin)
require.False(t, main.visible, "the install popup still covers the main window")
require.Equal(t, 0, d.raised)
require.Equal(t, []string{windowInstallProgress, windowInstallProgress}, ownersOf(s.hiddenWindows))
s.restoreHiddenWindows(windowInstallProgress)
require.True(t, main.visible, "the install popup restores the handed-over main window")
require.Equal(t, 1, d.raised)
require.Empty(t, s.hiddenWindows)
}
func TestInstallClosingUnderLoginHandsMainToLogin(t *testing.T) {
main := newFakeWindow(windowMain)
install := newFakeWindow(windowInstallProgress)
login := newFakeWindow(windowBrowserLogin)
login.visible = false
s := newTestWindowManager()
d := newFakeDesktop(s, main, install, login)
s.hideOtherWindows(windowInstallProgress)
login.visible = true
s.hideOtherWindows(windowBrowserLogin)
require.False(t, install.visible, "the login popup hides the install popup")
s.restoreHiddenWindows(windowInstallProgress)
require.False(t, main.visible, "the login popup still covers the main window")
require.Equal(t, 0, d.raised)
require.Equal(t, []string{windowBrowserLogin, windowBrowserLogin}, ownersOf(s.hiddenWindows))
s.restoreHiddenWindows(windowBrowserLogin)
require.True(t, main.visible)
require.Equal(t, 1, d.raised)
require.Empty(t, s.hiddenWindows)
}
func TestPopupClosingReshowsTheCoveringPopupItself(t *testing.T) {
main := newFakeWindow(windowMain)
install := newFakeWindow(windowInstallProgress)
login := newFakeWindow(windowBrowserLogin)
login.visible = false
s := newTestWindowManager()
d := newFakeDesktop(s, main, install, login)
s.hideOtherWindows(windowInstallProgress)
login.visible = true
s.hideOtherWindows(windowBrowserLogin)
// The login popup hid the install popup itself; closing the login popup must bring
// the install popup back rather than hand it over to its own owner.
s.restoreHiddenWindows(windowBrowserLogin)
require.True(t, install.visible, "a popup is never handed over to itself")
require.False(t, main.visible, "the main window stays with the install popup")
require.Equal(t, 0, d.raised)
require.Equal(t, []string{windowInstallProgress}, ownersOf(s.hiddenWindows))
}
func TestRestoreHiddenWindowsUnknownOwnerKeepsEverything(t *testing.T) {
main := newFakeWindow(windowMain)
s := newTestWindowManager()
d := newFakeDesktop(s, main)
s.hideOtherWindows(windowBrowserLogin)
s.restoreHiddenWindows(windowWelcome)
require.False(t, main.visible)
require.Equal(t, []string{windowBrowserLogin}, ownersOf(s.hiddenWindows))
require.Equal(t, 0, d.raised)
}
func TestRestoreHiddenWindowsWithoutMainDoesNotRaise(t *testing.T) {
settings := newFakeWindow(windowSettings)
s := newTestWindowManager()
d := newFakeDesktop(s, settings)
s.hideOtherWindows(windowBrowserLogin)
s.restoreHiddenWindows(windowBrowserLogin)
require.True(t, settings.visible)
require.Equal(t, 0, d.raised)
}
func TestRestoreHiddenWindowsEmptyIsNoop(t *testing.T) {
s := newTestWindowManager()
require.NotPanics(t, func() { s.restoreHiddenWindows(windowBrowserLogin) })
require.Empty(t, s.hiddenWindows)
}
func TestHideOtherWindowsRacingOwnRestoreReshowsWhatItHid(t *testing.T) {
main := newFakeWindow(windowMain)
s := newTestWindowManager()
d := newFakeDesktop(s, main)
enumerate := s.allWindows
// A restore for the same owner lands between the generation snapshot and the record.
s.allWindows = func() []hideableWindow {
s.restoreHiddenWindows(windowBrowserLogin)
return enumerate()
}
s.hideOtherWindows(windowBrowserLogin)
require.True(t, main.visible)
require.Equal(t, 1, main.hidden)
require.Empty(t, s.hiddenWindows)
require.Equal(t, 0, d.raised)
}
func TestHideOtherWindowsIgnoresRestoreOfAnotherOwner(t *testing.T) {
main := newFakeWindow(windowMain)
s := newTestWindowManager()
newFakeDesktop(s, main)
enumerate := s.allWindows
s.allWindows = func() []hideableWindow {
s.restoreHiddenWindows(windowInstallProgress)
return enumerate()
}
s.hideOtherWindows(windowBrowserLogin)
require.False(t, main.visible)
require.Equal(t, []string{windowBrowserLogin}, ownersOf(s.hiddenWindows))
}
func TestRestoringCloserRestoresOnlyItsOwner(t *testing.T) {
main := newFakeWindow(windowMain)
s := newTestWindowManager()
newFakeDesktop(s, main)
s.hideOtherWindows(windowBrowserLogin)
s.hiddenWindows = append(s.hiddenWindows, hiddenWindow{owner: windowInstallProgress})
s.restoringCloser(windowBrowserLogin)(&application.WebviewWindow{})
require.True(t, main.visible)
require.Equal(t, []string{windowInstallProgress}, ownersOf(s.hiddenWindows))
}
func TestStampGenerationTracksLatestPerWindow(t *testing.T) {
s := newTestWindowManager()
first := s.stampGeneration(windowBrowserLogin, "/#/dialog/browser-login")
require.Equal(t, "/#/dialog/browser-login?gen=1", first)
require.True(t, s.matchesGeneration(windowBrowserLogin, 1))
second := s.stampGeneration(windowBrowserLogin, "/#/dialog/browser-login?uri=x")
require.Equal(t, "/#/dialog/browser-login?uri=x&gen=2", second)
require.False(t, s.matchesGeneration(windowBrowserLogin, 1))
require.True(t, s.matchesGeneration(windowBrowserLogin, 2))
}
func TestMatchesGenerationUntrackedWindowAccepts(t *testing.T) {
s := newTestWindowManager()
require.True(t, s.matchesGeneration(windowMain, 0))
}
func TestPaintedGeneration(t *testing.T) {
tests := []struct {
name string
data any
want uint64
}{
{"string", "7", 7},
{"float", float64(7), 7},
{"slice", []any{"7"}, 7},
{"empty slice", []any{}, 0},
{"unparsable", "abc", 0},
{"nil", nil, 0},
}
for _, tc := range tests {
t.Run(tc.name, func(t *testing.T) {
require.Equal(t, tc.want, paintedGeneration(tc.data))
})
}
}