Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
80 changes: 64 additions & 16 deletions internal/daemon/daemon.go
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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
Expand Down Expand Up @@ -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) }()
Expand All @@ -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)
Expand Down
120 changes: 120 additions & 0 deletions internal/daemon/daemon_test.go
Original file line number Diff line number Diff line change
@@ -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)
}
61 changes: 45 additions & 16 deletions internal/tunnel/tunnel.go
Original file line number Diff line number Diff line change
Expand Up @@ -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
}

Expand Down Expand Up @@ -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)
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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 {
Expand All @@ -243,7 +264,7 @@ func (t *Tunnel) run() {
return
}
}
t.Status = Closed
t.setStatus(Closed)
close(t.Closed)
}

Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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) {
Expand Down
Loading