diff --git a/tailcat.go b/tailcat.go index a0eba13c4..9a4837ee6 100644 --- a/tailcat.go +++ b/tailcat.go @@ -232,6 +232,7 @@ type locoBackend struct { nm *netmap.NetworkMap allowedClients map[key.NodePublic]bool // or nil map for all eps []netip.AddrPort // our current local UDP endpoints, sorted + closeOnce sync.Once } func (b *locoBackend) derpRegionID() int { @@ -245,12 +246,23 @@ func (b *locoBackend) derpRegionID() int { } func (b *locoBackend) Close() error { - if e, ok := b.sys.Engine.GetOK(); ok { - e.Close() - } - if m, ok := b.sys.NetMon.GetOK(); ok { - m.Close() - } + b.closeOnce.Do(func() { + if b.ns != nil { + b.ns.Close() + } + if e, ok := b.sys.Engine.GetOK(); ok { + e.Close() + } + if m, ok := b.sys.NetMon.GetOK(); ok { + m.Close() + } + if d, ok := b.sys.Dialer.GetOK(); ok { + d.Close() + } + if bus, ok := b.sys.Bus.GetOK(); ok { + bus.Close() + } + }) return nil } diff --git a/tailcat_test.go b/tailcat_test.go index f6089a855..39fd0551c 100644 --- a/tailcat_test.go +++ b/tailcat_test.go @@ -216,6 +216,71 @@ func TestHalfClose(t *testing.T) { } } +func TestServerCloseClosesActiveConnections(t *testing.T) { + t.Parallel() + + dm := integration.RunDERPAndSTUN(t, mkLogger(t, "derpstun"), "127.0.0.1") + reg := dm.Regions[1] + if reg == nil { + t.Fatal("no region 1 in derpmap") + } + + clientKey := key.NewNode() + accepted := make(chan net.Conn, 1) + handlerDone := make(chan struct{}) + s := &Server{ + Logf: mkLogger(t, "server"), + Region: reg, + AllowedClients: []key.NodePublic{clientKey.Public()}, + ServedTCPPorts: []filter.PortRange{{First: 80, Last: 80}}, + OnTCP: func(port uint16) func(net.Conn) { + if port != 80 { + return nil + } + return func(conn net.Conn) { + accepted <- conn + defer close(handlerDone) + var buf [1]byte + _, _ = conn.Read(buf[:]) + } + }, + } + if err := s.Start(); err != nil { + t.Fatalf("server Start: %v", err) + } + t.Cleanup(func() { s.Close() }) + + c := &Client{Server: s.ConnBlob(), Key: clientKey, Logf: mkLogger(t, "client")} + t.Cleanup(func() { c.Close() }) + PingForTest(t, s, c) + + ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second) + defer cancel() + clientConn, err := c.DialTCPPort(ctx, 80) + if err != nil { + t.Fatalf("DialTCPPort: %v", err) + } + t.Cleanup(func() { clientConn.Close() }) + + var serverConn net.Conn + select { + case serverConn = <-accepted: + t.Cleanup(func() { serverConn.Close() }) + case <-ctx.Done(): + t.Fatalf("waiting for accepted connection: %v", ctx.Err()) + } + + if err := s.Close(); err != nil { + t.Fatalf("server Close: %v", err) + } + + select { + case <-handlerDone: + case <-time.After(5 * time.Second): + t.Fatal("server Close left an active netstack connection open") + } +} + func TestConnBlob(t *testing.T) { akey := func(a [32]byte) NodePublic { return NodePublic{key.NodePublicFromRaw32(mem.B(a[:]))}