mirror of
https://github.com/netbirdio/netbird.git
synced 2026-10-07 05:59:06 +02:00
Merge remote-tracking branch 'origin/main' into feat-post_quantum_ml_kem
This commit is contained in:
@@ -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
@@ -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,
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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")
|
||||
}
|
||||
|
||||
@@ -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{
|
||||
|
||||
@@ -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()
|
||||
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
@@ -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()
|
||||
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
@@ -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) {
|
||||
|
||||
@@ -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])
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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);
|
||||
|
||||
@@ -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
@@ -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 }
|
||||
|
||||
|
||||
@@ -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))
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user