Skip to content
Merged
Show file tree
Hide file tree
Changes from 2 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 entry is required")
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 mapping peer responses (recovered)",

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Would "panic processing peer mappings (recovered)" be more clear?

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Yeah, that reads better. Updated.

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)
}
15 changes: 15 additions & 0 deletions tailnet/coordinator.go
Original file line number Diff line number Diff line change
Expand Up @@ -27,6 +27,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 @@ -235,6 +236,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 +294,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 @@ -480,6 +488,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