diff --git a/internal/daemon/daemon.go b/internal/daemon/daemon.go index 6505c95..febf00c 100644 --- a/internal/daemon/daemon.go +++ b/internal/daemon/daemon.go @@ -48,6 +48,9 @@ type daemon struct { // TODO: write proper concurrent map structure for this tunnels map[string]*tunnel.Tunnel + // opening holds the tunnels whose Open is in progress. The channel is + // closed once it finishes, successful or not. + opening map[string]chan struct{} mutex sync.RWMutex once sync.Once @@ -56,8 +59,13 @@ type daemon struct { func newDaemon(parent context.Context, ln net.Listener) (*daemon, context.CancelFunc) { ctx, cancel := context.WithCancel(parent) - tunnels := make(map[string]*tunnel.Tunnel) - d := &daemon{ctx: ctx, cancel: cancel, ln: ln, tunnels: tunnels} + d := &daemon{ + ctx: ctx, + cancel: cancel, + ln: ln, + tunnels: make(map[string]*tunnel.Tunnel), + opening: make(map[string]chan struct{}), + } go func() { // Parent-driven shutdown @@ -157,35 +165,72 @@ func (d *daemon) openTunnel(conn net.Conn, desc *tunnel.Desc) { var err error defer func() { respond(conn, err, nil) }() - d.mutex.RLock() - _, exists := d.tunnels[desc.Name] - d.mutex.RUnlock() - if exists { - err = AlreadyRunning + if err = d.reserve(desc.Name); err != nil { log.Errorf("%v: could not open: %v", desc.Name, err) return } t := tunnel.FromDesc(desc) - if err = t.Open(); err != nil { - log.Errorf("%v: could not open: %v", t.Name, err) - return - } + err = t.Open() d.mutex.Lock() - d.tunnels[t.Name] = t + close(d.opening[desc.Name]) + delete(d.opening, desc.Name) + if err == nil { + d.tunnels[t.Name] = t + } d.mutex.Unlock() + if err != nil { + log.Errorf("%v: could not open: %v", t.Name, err) + return + } + // Register closing logic go func() { <-t.Closed - d.mutex.Lock() - delete(d.tunnels, t.Name) - d.mutex.Unlock() + d.removeTunnel(t) log.Infof("Closed tunnel %s", t.Name) }() } +// reserve marks the tunnel name as being opened. If another client is +// opening the same tunnel, it waits for that attempt to finish first: if it +// succeeded the tunnel is already running, otherwise this client tries +// again itself. +func (d *daemon) reserve(name string) error { + for { + d.mutex.Lock() + if _, ok := d.tunnels[name]; ok { + d.mutex.Unlock() + return AlreadyRunning + } + wait, ok := d.opening[name] + if !ok { + d.opening[name] = make(chan struct{}) + d.mutex.Unlock() + return nil + } + d.mutex.Unlock() + + select { + case <-wait: + case <-d.ctx.Done(): + return d.ctx.Err() + } + } +} + +// removeTunnel forgets about t, unless its name has meanwhile been taken +// by a tunnel that was opened after it. +func (d *daemon) removeTunnel(t *tunnel.Tunnel) { + d.mutex.Lock() + defer d.mutex.Unlock() + if d.tunnels[t.Name] == t { + delete(d.tunnels, t.Name) + } +} + func (d *daemon) closeTunnel(conn net.Conn, q *tunnel.Desc) { var err error defer func() { respond(conn, err, nil) }() @@ -204,13 +249,16 @@ func (d *daemon) closeTunnel(conn net.Conn, q *tunnel.Desc) { return } <-t.Closed + // Also remove it here, so it is gone by the time the client gets the + // response, rather than whenever the closing goroutine gets to it. + d.removeTunnel(t) } func (d *daemon) listTunnels(conn net.Conn) { d.mutex.RLock() ts := make(map[string]tunnel.Desc, len(d.tunnels)) for n, t := range d.tunnels { - ts[n] = *t.Desc + ts[n] = t.Snapshot() } d.mutex.RUnlock() respond(conn, nil, ts) diff --git a/internal/daemon/daemon_test.go b/internal/daemon/daemon_test.go new file mode 100644 index 0000000..b8e7baf --- /dev/null +++ b/internal/daemon/daemon_test.go @@ -0,0 +1,120 @@ +package daemon + +import ( + "context" + "errors" + "testing" + "time" + + "github.com/alebeck/boring/internal/tunnel" +) + +func testDaemon(t *testing.T) *daemon { + ctx, cancel := context.WithCancel(context.Background()) + t.Cleanup(cancel) + return &daemon{ + ctx: ctx, + cancel: cancel, + tunnels: make(map[string]*tunnel.Tunnel), + opening: make(map[string]chan struct{}), + } +} + +func TestRemoveTunnelKeepsNewer(t *testing.T) { + d := testDaemon(t) + old := tunnel.FromDesc(&tunnel.Desc{Name: "a"}) + cur := tunnel.FromDesc(&tunnel.Desc{Name: "a"}) + d.tunnels["a"] = cur + + d.removeTunnel(old) + if d.tunnels["a"] != cur { + t.Fatal("removing an old tunnel dropped the one that replaced it") + } + + d.removeTunnel(cur) + if _, ok := d.tunnels["a"]; ok { + t.Fatal("tunnel not removed") + } +} + +func TestReserveRunning(t *testing.T) { + d := testDaemon(t) + d.tunnels["a"] = tunnel.FromDesc(&tunnel.Desc{Name: "a"}) + if err := d.reserve("a"); !errors.Is(err, AlreadyRunning) { + t.Fatalf("expected %v, got %v", AlreadyRunning, err) + } +} + +// finishOpen mimics the end of openTunnel. +func finishOpen(d *daemon, name string, ok bool) { + d.mutex.Lock() + defer d.mutex.Unlock() + close(d.opening[name]) + delete(d.opening, name) + if ok { + d.tunnels[name] = tunnel.FromDesc(&tunnel.Desc{Name: name}) + } +} + +func reserveAsync(d *daemon, name string) chan error { + res := make(chan error, 1) + go func() { res <- d.reserve(name) }() + return res +} + +func expectBlocked(t *testing.T, res chan error) { + t.Helper() + select { + case err := <-res: + t.Fatalf("reserve returned early: %v", err) + case <-time.After(50 * time.Millisecond): + } +} + +func expectResult(t *testing.T, res chan error, want error) { + t.Helper() + select { + case err := <-res: + if !errors.Is(err, want) { + t.Fatalf("expected %v, got %v", want, err) + } + case <-time.After(time.Second): + t.Fatal("reserve did not return") + } +} + +func TestReserveWaitsForSuccessfulOpen(t *testing.T) { + d := testDaemon(t) + if err := d.reserve("a"); err != nil { + t.Fatal(err) + } + res := reserveAsync(d, "a") + expectBlocked(t, res) + finishOpen(d, "a", true) + expectResult(t, res, AlreadyRunning) +} + +func TestReserveRetriesAfterFailedOpen(t *testing.T) { + d := testDaemon(t) + if err := d.reserve("a"); err != nil { + t.Fatal(err) + } + res := reserveAsync(d, "a") + expectBlocked(t, res) + finishOpen(d, "a", false) + expectResult(t, res, nil) + if _, ok := d.opening["a"]; !ok { + t.Fatal("second reserve did not take over the name") + } +} + +func TestReserveShutdown(t *testing.T) { + d := testDaemon(t) + if err := d.reserve("a"); err != nil { + t.Fatal(err) + } + res := reserveAsync(d, "a") + expectBlocked(t, res) + d.cancel() + expectResult(t, res, context.Canceled) +} diff --git a/internal/tunnel/tunnel.go b/internal/tunnel/tunnel.go index 46c1fc1..89a5a31 100644 --- a/internal/tunnel/tunnel.go +++ b/internal/tunnel/tunnel.go @@ -46,11 +46,15 @@ type Tunnel struct { hops []ssh_config.Hop Closed chan struct{} stop chan struct{} + stopOnce sync.Once listener net.Listener wg sync.WaitGroup client *ssh.Client localAddr *address remoteAddr *address + // mu guards Status and LastConn, which the tunnel's own goroutines + // update while the daemon may be reading them for a listing. + mu sync.Mutex *Desc } @@ -85,14 +89,31 @@ func (t *Tunnel) Open() (err error) { t.Closed = make(chan struct{}) } + t.mu.Lock() + t.Status = Open + t.LastConn = time.Now() + t.mu.Unlock() + go t.run() log.Infof("%v: opened tunnel", t.Name) - t.Status = Open - t.LastConn = time.Now() return } +// Snapshot returns a copy of the tunnel's description that is safe to take +// while the tunnel is running. +func (t *Tunnel) Snapshot() Desc { + t.mu.Lock() + defer t.mu.Unlock() + return *t.Desc +} + +func (t *Tunnel) setStatus(s Status) { + t.mu.Lock() + t.Status = s + t.mu.Unlock() +} + func (t *Tunnel) prepare() error { // We need to pass the user as it's needed for matching Match blocks sc, err := ssh_config.ParseSSHConfig(t.Host, t.User) @@ -175,7 +196,7 @@ func (t *Tunnel) makeClient() error { } // Wait for all wrapped clients to close in case of tunnel closing or reconnection - go t.waitFor(func() { wg.Wait() }) + t.goWait(wg.Wait) t.client = c return nil @@ -222,8 +243,8 @@ func (t *Tunnel) run() { close(disconn) }() - go t.waitFor(func() { t.keepAlive(disconn) }) - go t.waitFor(func() { t.handleConns() }) + t.goWait(func() { t.keepAlive(disconn) }) + t.goWait(t.handleConns) stopped := false select { @@ -243,7 +264,7 @@ func (t *Tunnel) run() { return } } - t.Status = Closed + t.setStatus(Closed) close(t.Closed) } @@ -290,7 +311,7 @@ func (t *Tunnel) handleForward() { log.Errorf("%v: could not accept: %v", t.Name, err) return } - go t.waitFor(func() { + t.goWait(func() { addr := t.remoteAddr if t.Mode == Remote || t.Mode == RemoteSocks { addr = t.localAddr @@ -335,12 +356,12 @@ func (t *Tunnel) handleSocks() { log.Errorf("%v: could not accept: %v", t.Name, err) return } - go t.waitFor(func() { serv.ServeConn(conn) }) + t.goWait(func() { serv.ServeConn(conn) }) } } func (t *Tunnel) reconnectLoop() error { - t.Status = Reconn + t.setStatus(Reconn) timeout := time.After(reconnectTimeout) wait := time.NewTimer(2 * time.Millisecond) // First time try (essent.) immediately waitTime := initReconnectWait @@ -368,20 +389,28 @@ func (t *Tunnel) reconnectLoop() error { } } +// Close signals the tunnel to stop. It is safe to call concurrently and +// more than once; wait on Closed for the tunnel to actually shut down. func (t *Tunnel) Close() error { - if t.Status == Closed { + t.mu.Lock() + closed := t.Status == Closed + t.mu.Unlock() + if closed { return fmt.Errorf("trying to close a closed tunnel") } - close(t.stop) + t.stopOnce.Do(func() { close(t.stop) }) return nil } -// Logic registered with waitFor will be waited for upon tunnel closing -// and reconnecting. -func (t *Tunnel) waitFor(f func()) { +// goWait runs f in a new goroutine that will be waited for upon tunnel +// closing and reconnecting. The wait group is incremented before the +// goroutine starts, so a concurrent Wait cannot miss it. +func (t *Tunnel) goWait(f func()) { t.wg.Add(1) - defer t.wg.Done() - f() + go func() { + defer t.wg.Done() + f() + }() } func parseAddr(addr string, allowShort bool) (*address, error) { diff --git a/internal/tunnel/tunnel_test.go b/internal/tunnel/tunnel_test.go new file mode 100644 index 0000000..5eac639 --- /dev/null +++ b/internal/tunnel/tunnel_test.go @@ -0,0 +1,73 @@ +package tunnel + +import ( + "sync" + "testing" +) + +// openTunnel returns a tunnel in the state Open leaves it in, without +// connecting anywhere. +func openTunnel() *Tunnel { + return &Tunnel{ + Desc: &Desc{Name: "test", Status: Open}, + stop: make(chan struct{}), + Closed: make(chan struct{}), + } +} + +func TestCloseTwice(t *testing.T) { + tun := openTunnel() + if err := tun.Close(); err != nil { + t.Fatalf("first close: %v", err) + } + // The tunnel stays open until its run loop notices the stop signal, + // so a second close in that window must not panic. + if err := tun.Close(); err != nil { + t.Fatalf("second close: %v", err) + } + select { + case <-tun.stop: + default: + t.Fatal("stop channel not closed") + } +} + +func TestCloseConcurrent(t *testing.T) { + tun := openTunnel() + var wg sync.WaitGroup + for range 10 { + wg.Go(func() { + if err := tun.Close(); err != nil { + t.Errorf("close: %v", err) + } + }) + } + wg.Wait() +} + +func TestCloseClosed(t *testing.T) { + tun := openTunnel() + tun.Status = Closed + if err := tun.Close(); err == nil { + t.Fatal("expected error when closing a closed tunnel") + } +} + +// Run with -race: status updates from the tunnel's goroutines must not race +// with the daemon taking snapshots for a listing. +func TestSnapshotConcurrent(t *testing.T) { + tun := openTunnel() + var wg sync.WaitGroup + wg.Go(func() { + for range 1000 { + tun.setStatus(Reconn) + tun.setStatus(Open) + } + }) + for range 1000 { + if s := tun.Snapshot().Status; s != Open && s != Reconn { + t.Fatalf("unexpected status %v", s) + } + } + wg.Wait() +} diff --git a/test/e2e/concurrency_test.go b/test/e2e/concurrency_test.go new file mode 100644 index 0000000..3f5c8c5 --- /dev/null +++ b/test/e2e/concurrency_test.go @@ -0,0 +1,133 @@ +package e2e + +import ( + "strings" + "testing" + + "golang.org/x/sync/errgroup" +) + +// runConcurrently runs the same CLI command n times in parallel and fails +// the test if any of them exits with a non-zero code. +func runConcurrently(t *testing.T, env []string, n int, args ...string) { + t.Helper() + runConcurrentlyAllow(t, env, n, nil, args...) +} + +// runConcurrentlyAllow is like runConcurrently, but also accepts a non-zero +// exit code if the output contains one of the allowed messages. +func runConcurrentlyAllow(t *testing.T, env []string, n int, allowed []string, args ...string) { + t.Helper() + var g errgroup.Group + outs := make([]string, n) + codes := make([]int, n) + for i := range n { + g.Go(func() error { + var err error + codes[i], outs[i], err = cliCommand(env, args...) + return err + }) + } + if err := g.Wait(); err != nil { + t.Fatalf("failed to run CLI command: %v", err) + } + for i := range n { + if codes[i] != 0 && !containsAny(outs[i], allowed) { + t.Fatalf("%v: exit code %d: %s", args, codes[i], outs[i]) + } + } +} + +func containsAny(s string, subs []string) bool { + for _, sub := range subs { + if strings.Contains(s, sub) { + return true + } + } + return false +} + +func listStatus(t *testing.T, env []string) string { + t.Helper() + c, out, err := cliCommand(env, "list") + if err != nil { + t.Fatalf("failed to run CLI command: %v", err) + } + if c != 0 { + t.Fatalf("list: exit code %d: %s", c, out) + } + lines := strings.Split(strings.TrimSpace(stripANSI(out)), "\n") + return strings.Fields(lines[1])[0] +} + +// Closing the same tunnel from several clients at once used to crash the +// daemon with "close of closed channel". The daemon is started with +// BORING_NO_SPAWN, so a crash makes the following commands fail. +// +// Clients that lose the race legitimately find the tunnel no longer +// running, either when listing running tunnels or when sending the close +// command, and exit with an error. That is expected here. +// Messages of clients that lost the race, depending on how far the winner +// got: the tunnel is gone from the list, gone from the daemon, or already +// shut down but not yet removed. +var notRunning = []string{ + "No running tunnels match", + "tunnel not running", + "trying to close a closed tunnel", +} + +func TestCloseConcurrent(t *testing.T) { + env, cancel, err := makeDefaultEnvWithDaemon(t) + if err != nil { + t.Fatalf("%v", err.Error()) + } + defer cancel() + + for range 5 { + runConcurrently(t, env, 1, "open", "test") + runConcurrentlyAllow(t, env, 4, notRunning, "close", "test") + if s := listStatus(t, env); s != "closed" { + t.Fatalf("expected tunnel to be closed, got %q", s) + } + } +} + +// Opening the same tunnel from several clients at once must open it +// exactly once. The others should report it as already running instead +// of failing to bind the local port. +func TestOpenConcurrent(t *testing.T) { + env, cancel, err := makeDefaultEnvWithDaemon(t) + if err != nil { + t.Fatalf("%v", err.Error()) + } + defer cancel() + + for range 5 { + runConcurrently(t, env, 4, "open", "test") + if s := listStatus(t, env); s == "closed" { + t.Fatal("expected tunnel to be open") + } + testTunnel(t, "localhost:49711", "localhost:49712") + runConcurrently(t, env, 1, "close", "test") + } +} + +// A tunnel that is reopened right after being closed must stay tracked by +// the daemon, i.e. it shows up in the list and can be closed again. +func TestCloseReopen(t *testing.T) { + env, cancel, err := makeDefaultEnvWithDaemon(t) + if err != nil { + t.Fatalf("%v", err.Error()) + } + defer cancel() + + runConcurrently(t, env, 1, "open", "test") + for range 10 { + runConcurrently(t, env, 1, "close", "test") + runConcurrently(t, env, 1, "open", "test") + if s := listStatus(t, env); s == "closed" { + t.Fatal("reopened tunnel not tracked by the daemon") + } + } + runConcurrently(t, env, 1, "close", "test") +}