Skip to content
Merged
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
16 changes: 16 additions & 0 deletions enterprise/tailnet/connio.go
Original file line number Diff line number Diff line change
Expand Up @@ -3,6 +3,7 @@ package tailnet
import (
"context"
"fmt"
"runtime/debug"
"slices"
"sync"

Expand Down Expand Up @@ -100,6 +101,17 @@ func (c *connIO) recvLoop() {
}
}()
defer c.Close()
// This must be deferred after Close so it runs first: Enqueue needs the
// response channel open to deliver CloseErrInternal to the peer.
defer func() {
if recovered := recover(); recovered != nil {
c.logger.Error(c.peerCtx, "panic handling peer request (recovered)",
slog.F("panic", recovered),
slog.F("stack", string(debug.Stack())),
)
_ = c.Enqueue(&proto.CoordinateResponse{Error: agpl.CloseErrInternal})
}
}()
for {
select {
case <-c.coordCtx.Done():
Expand Down Expand Up @@ -128,6 +140,10 @@ var errDisconnect = xerrors.New("graceful disconnect")

func (c *connIO) handleRequest(req *proto.CoordinateRequest) error {
c.logger.Debug(c.peerCtx, "got request")
if err := agpl.ValidateCoordinateRequest(req); err != nil {
c.logger.Warn(c.peerCtx, "invalid coordinate request", slog.Error(err))
return err
}
err := c.auth.Authorize(c.peerCtx, req)
if err != nil {
c.logger.Warn(c.peerCtx, "unauthorized request", slog.Error(err))
Expand Down
89 changes: 89 additions & 0 deletions enterprise/tailnet/connio_internal_test.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,89 @@
package tailnet

import (
"context"
"strings"
"testing"

"github.com/google/uuid"
"github.com/stretchr/testify/require"

"cdr.dev/slog/v3/sloggers/sloghuman"
"cdr.dev/slog/v3/sloggers/slogtest"
agpl "github.com/coder/coder/v2/tailnet"
"github.com/coder/coder/v2/tailnet/proto"
agpltest "github.com/coder/coder/v2/tailnet/test"
"github.com/coder/coder/v2/testutil"
)

func TestConnIOHandleRequestRejectsBeforeMutation(t *testing.T) {
t.Parallel()
ctx := testutil.Context(t, testutil.WaitShort)
bindings := make(chan binding, 1)
var logbuf strings.Builder
c := &connIO{
id: uuid.New(),
coordCtx: ctx,
peerCtx: ctx,
logger: slogtest.Make(t, &slogtest.Options{IgnoreErrors: true}).AppendSinks(sloghuman.Sink(&logbuf)),
bindings: bindings,
auth: agpl.SingleTailnetCoordinateeAuth{},
}

err := c.handleRequest(&proto.CoordinateRequest{
UpdateSelf: &proto.CoordinateRequest_UpdateSelf{Node: &proto.Node{
PreferredDerp: 31002,
}},
ReadyForHandshake: []*proto.CoordinateRequest_ReadyForHandshake{nil},
})
require.EqualError(t, err, "ready_for_handshake entries must not be nil")
require.Contains(t, logbuf.String(), "invalid coordinate request")
select {
case binding := <-bindings:
t.Fatalf("unexpected binding: %+v", binding)
default:
}
}

func TestConnIORecvLoopPanicAfterAuthorization(t *testing.T) {
t.Parallel()
ctx := testutil.Context(t, testutil.WaitShort)
logger := slogtest.Make(t, &slogtest.Options{IgnoreErrors: true})
bindings := make(chan binding, 1)
requests := make(chan *proto.CoordinateRequest, 1)
responses := make(chan *proto.CoordinateResponse, 1)
rfhs := make(chan readyForHandshake)
close(rfhs)
authorized := make(chan struct{}, 1)
peerID := uuid.New()
destinationID := uuid.New()
c := newConnIO(ctx, ctx, logger, bindings, make(chan tunnel, 1), rfhs,
requests, responses, peerID, t.Name(), agpltest.CoordinateeAuthFunc(func(context.Context, *proto.CoordinateRequest) error {
authorized <- struct{}{}
return nil
}))
t.Cleanup(func() { require.NoError(t, c.Close()) })
c.setLatestMapping([]mapping{{peer: destinationID}})

testutil.RequireSend(ctx, t, requests, &proto.CoordinateRequest{
ReadyForHandshake: []*proto.CoordinateRequest_ReadyForHandshake{{Id: agpl.UUIDToByteSlice(destinationID)}},
})
testutil.RequireReceive(ctx, t, authorized)
response := testutil.RequireReceive(ctx, t, responses)
require.Equal(t, agpl.CloseErrInternal, response.Error)
require.NotContains(t, response.Error, "send on closed channel")
select {
case _, ok := <-responses:
require.False(t, ok)
case <-ctx.Done():
t.Fatal("response channel did not close")
}
select {
case <-c.Done():
case <-ctx.Done():
t.Fatal("connection did not close")
}
withdrawn := testutil.RequireReceive(ctx, t, bindings)
require.Equal(t, bKey(peerID), withdrawn.bKey)
require.Equal(t, proto.CoordinateResponse_PeerUpdate_LOST, withdrawn.kind)
}
55 changes: 55 additions & 0 deletions enterprise/tailnet/mapper_internal_test.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,55 @@
package tailnet

import (
"testing"

"github.com/google/uuid"
"github.com/stretchr/testify/require"

"cdr.dev/slog/v3/sloggers/slogtest"
agpl "github.com/coder/coder/v2/tailnet"
"github.com/coder/coder/v2/tailnet/proto"
"github.com/coder/coder/v2/testutil"
)

func TestMapperPanic(t *testing.T) {
t.Parallel()
ctx := testutil.Context(t, testutil.WaitLong)
logger := slogtest.Make(t, &slogtest.Options{IgnoreErrors: true})
bindings := make(chan binding, 1)
responses := make(chan *proto.CoordinateResponse, 2)
c := newConnIO(ctx, ctx, logger, bindings, make(chan tunnel, 1), nil,
make(chan *proto.CoordinateRequest), responses, uuid.New(), t.Name(), agpl.SingleTailnetCoordinateeAuth{})
t.Cleanup(func() { require.NoError(t, c.Close()) })
coordinatorID := uuid.New()
m := &mapper{
ctx: c.peerCtx,
logger: logger,
c: c,
mappings: make(chan []mapping),
heartbeats: &heartbeats{self: coordinatorID},
// A nil sent map forces a panic after the mapping has been selected.
sent: nil,
}
go m.run()
require.NoError(t, agpl.SendCtx(ctx, m.mappings, []mapping{{
peer: uuid.New(), coordinator: coordinatorID,
node: &proto.Node{}, kind: proto.CoordinateResponse_PeerUpdate_NODE,
}}))
response := testutil.RequireReceive(ctx, t, responses)
require.Equal(t, agpl.CloseErrInternal, response.Error)
select {
case _, ok := <-responses:
require.False(t, ok)
case <-ctx.Done():
t.Fatal("response channel did not close")
}
select {
case <-c.Done():
case <-ctx.Done():
t.Fatal("connection did not close")
}
withdrawn := testutil.RequireReceive(ctx, t, bindings)
require.Equal(t, bKey(c.id), withdrawn.bKey)
require.Equal(t, proto.CoordinateResponse_PeerUpdate_LOST, withdrawn.kind)
}
11 changes: 11 additions & 0 deletions enterprise/tailnet/pgcoord.go
Original file line number Diff line number Diff line change
Expand Up @@ -4,6 +4,7 @@ import (
"context"
"database/sql"
"math"
"runtime/debug"
"slices"
"strings"
"sync"
Expand Down Expand Up @@ -742,6 +743,16 @@ func newMapper(c *connIO, logger slog.Logger, h *heartbeats) *mapper {
}

func (m *mapper) run() {
defer func() {
if recovered := recover(); recovered != nil {
m.logger.Error(m.ctx, "panic processing peer mappings (recovered)",
slog.F("panic", recovered),
slog.F("stack", string(debug.Stack())),
)
_ = m.c.Enqueue(&proto.CoordinateResponse{Error: agpl.CloseErrInternal})
_ = m.c.Close()
}
}()
for {
var best map[uuid.UUID]mapping
select {
Expand Down
33 changes: 33 additions & 0 deletions enterprise/tailnet/requests_test.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,33 @@
package tailnet_test

import (
"testing"

"github.com/stretchr/testify/require"

"cdr.dev/slog/v3/sloggers/slogtest"
"github.com/coder/coder/v2/coderd/database/dbtestutil"
"github.com/coder/coder/v2/enterprise/tailnet"
"github.com/coder/coder/v2/tailnet/test"
"github.com/coder/coder/v2/testutil"
)

func TestPGCoordinator_InvalidRequests(t *testing.T) {
t.Parallel()
ctx := testutil.Context(t, testutil.WaitLong)
store, ps := dbtestutil.NewDB(t)
coordinator, err := tailnet.NewPGCoord(ctx, slogtest.Make(t, &slogtest.Options{IgnoreErrors: true}), ps, store)
require.NoError(t, err)
t.Cleanup(func() { _ = coordinator.Close() })
test.InvalidCoordinateRequestTest(ctx, t, coordinator)
}

func TestPGCoordinator_RequestPanic(t *testing.T) {
t.Parallel()
ctx := testutil.Context(t, testutil.WaitLong)
store, ps := dbtestutil.NewDB(t)
coordinator, err := tailnet.NewPGCoord(ctx, slogtest.Make(t, &slogtest.Options{IgnoreErrors: true}), ps, store)
require.NoError(t, err)
t.Cleanup(func() { _ = coordinator.Close() })
test.CoordinateRequestPanicTest(ctx, t, coordinator)
}
29 changes: 29 additions & 0 deletions tailnet/coordinator.go
Original file line number Diff line number Diff line change
Expand Up @@ -7,6 +7,7 @@ import (
"io"
"net/http"
"net/netip"
runtimedebug "runtime/debug"
"sync"
"time"

Expand All @@ -27,6 +28,7 @@ const (
RequestBufferSize = 32
CloseErrOverwritten = "peer ID overwritten by new connection"
CloseErrCoordinatorClose = "coordinator closed"
CloseErrInternal = "internal coordinator error"
ReadyForHandshakeError = "ready for handshake error"
)

Expand Down Expand Up @@ -176,6 +178,15 @@ func (c *coordinator) Coordinate(
c.wg.Add(1)
go func() {
defer c.wg.Done()
// Peer cleanup can also panic, outside reqLoop's recovery boundary.
defer func() {
if recovered := recover(); recovered != nil {
logger.Error(ctx, "panic coordinating peer (recovered)",
slog.F("panic", recovered),
slog.F("stack", string(runtimedebug.Stack())),
)
}
}()
loopErr := p.reqLoop(ctx, logger, c.core.handleRequest)
closeErrStr := ""
if loopErr != nil {
Expand Down Expand Up @@ -235,6 +246,9 @@ func (c *core) node(id uuid.UUID) *Node {
}

func (c *core) handleRequest(ctx context.Context, p *peer, req *proto.CoordinateRequest) error {
if err := ValidateCoordinateRequest(req); err != nil {
return err
}
c.mutex.Lock()
defer c.mutex.Unlock()
if c.closed {
Expand Down Expand Up @@ -290,6 +304,10 @@ func (c *core) handleRequest(ctx context.Context, p *peer, req *proto.Coordinate
}
if req.Disconnect != nil {
c.removePeerLocked(p.id, proto.CoordinateResponse_PeerUpdate_DISCONNECTED, "graceful disconnect", "")
// The peer and its response channel are gone, so nothing else in this
// request can be delivered. Matches the enterprise coordinator, which
// stops processing a request once it sees Disconnect.
return nil
}
if rfhs := req.ReadyForHandshake; rfhs != nil {
err := c.handleReadyForHandshakeLocked(pr, rfhs)
Expand Down Expand Up @@ -349,6 +367,10 @@ func (c *core) nodeUpdateLocked(p *peer, node *proto.Node) (err error) {

p.node = node
c.updateTunnelPeersLocked(p.id, node, proto.CoordinateResponse_PeerUpdate_NODE, "node update")
// Fan-out can recursively remove this peer if its response buffer is full.
if c.peers[p.id] != p {
return ErrAlreadyRemoved
}
return nil
}

Expand Down Expand Up @@ -480,6 +502,13 @@ func (c *core) removePeerLocked(id uuid.UUID, kind proto.CoordinateResponse_Peer
return
}
c.updateTunnelPeersLocked(id, nil, kind, reason)
if _, ok := c.peers[id]; !ok {
// A tunnel peer with a full response buffer gets removed during the
// fan-out above, and removing it fans out to us in turn. When that
// nested call already closed and deleted this peer, p is stale and
// closing p.resps again would panic. Nothing is left to do.
return
}
c.tunnels.removeAll(id)
if closeErr != "" {
select {
Expand Down
Loading
Loading