fix: prevent coderd crashes from unsupported Tailnet RPCs by geokat · Pull Request #28429 · coder/coder · GitHub
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
2 changes: 1 addition & 1 deletion agent/agentsocket/server.go
1 change: 1 addition & 0 deletions agent/agenttest/client.go
Original file line number Diff line number Diff line change
Expand Up @@ -78,6 +78,7 @@ func NewClientWithSecrets(t testing.TB,
fakeAAPI := NewFakeAgentAPI(t, logger, mp, statsChan)
err = agentproto.DRPCRegisterAgent(mux, fakeAAPI)
require.NoError(t, err)
// Keep panics unrecovered in this test server so they fail tests loudly.
server := drpcserver.NewWithOptions(mux, drpcserver.Options{
Manager: drpcsdk.DefaultDRPCOptions(nil),
Log: func(err error) {
Expand Down
32 changes: 28 additions & 4 deletions coderd/agentapi/api.go
Original file line number Diff line number Diff line change
Expand Up @@ -59,7 +59,7 @@ type API struct {
*SubAgentAPI
*BoundaryLogsAPI
*ContextAPI
*tailnet.DRPCService
tailnetService *tailnet.DRPCService

cachedWorkspaceFields *CachedWorkspaceFields

Expand All @@ -68,6 +68,28 @@ type API struct {

var _ agentproto.DRPCAgentServer = &API{}

// agentTailnetService exposes only Tailnet RPCs intended for workspace agents.
// Other current and future RPCs remain unavailable until explicitly forwarded.
type agentTailnetService struct {
tailnetproto.DRPCTailnetUnimplementedServer

service *tailnet.DRPCService
}

func (s *agentTailnetService) PostTelemetry(ctx context.Context, req *tailnetproto.TelemetryRequest) (*tailnetproto.TelemetryResponse, error) {
return s.service.PostTelemetry(ctx, req)
}

func (s *agentTailnetService) StreamDERPMaps(req *tailnetproto.StreamDERPMapsRequest, stream tailnetproto.DRPCTailnet_StreamDERPMapsStream) error {
return s.service.StreamDERPMaps(req, stream)
}

func (s *agentTailnetService) Coordinate(stream tailnetproto.DRPCTailnet_CoordinateStream) error {
return s.service.Coordinate(stream)
}

var _ tailnetproto.DRPCTailnetServer = (*agentTailnetService)(nil)

type Options struct {
AgentID uuid.UUID
OwnerID uuid.UUID
Expand Down Expand Up @@ -226,7 +248,7 @@ func New(opts Options, workspace database.Workspace, agent database.WorkspaceAge
Log: opts.Log,
}

api.DRPCService = &tailnet.DRPCService{
api.tailnetService = &tailnet.DRPCService{
CoordPtr: opts.TailnetCoordinator,
Logger: opts.Log,
DerpMapUpdateFrequency: opts.DerpMapUpdateFrequency,
Expand Down Expand Up @@ -279,12 +301,14 @@ func (a *API) Server(ctx context.Context) (*drpcserver.Server, error) {
return nil, xerrors.Errorf("register agent API protocol in DRPC mux: %w", err)
}

err = tailnetproto.DRPCRegisterTailnet(mux, a)
err = tailnetproto.DRPCRegisterTailnet(mux, &agentTailnetService{
service: a.tailnetService,
})
if err != nil {
return nil, xerrors.Errorf("register tailnet API protocol in DRPC mux: %w", err)
}

return drpcserver.NewWithOptions(&tracing.DRPCHandler{Handler: mux},
return drpcsdk.NewServer(a.opts.Log, &tracing.DRPCHandler{Handler: mux},
drpcserver.Options{
Manager: drpcsdk.DefaultDRPCOptions(nil),
Log: func(err error) {
Expand Down
2 changes: 1 addition & 1 deletion coderd/aibridged.go
Original file line number Diff line number Diff line change
Expand Up @@ -84,7 +84,7 @@ func (api *API) CreateInMemoryAIBridgeServer(dialCtx context.Context) (client ai
if err := aibridgedserver.Register(mux, srv); err != nil {
return nil, err
}
server := drpcserver.NewWithOptions(&tracing.DRPCHandler{Handler: mux},
server := drpcsdk.NewServer(api.Logger, &tracing.DRPCHandler{Handler: mux},
drpcserver.Options{
Manager: drpcsdk.DefaultDRPCOptions(nil),
Log: func(err error) {
Expand Down
2 changes: 1 addition & 1 deletion coderd/coderd.go
Original file line number Diff line number Diff line change
Expand Up @@ -2587,7 +2587,7 @@ func (api *API) CreateInMemoryTaggedProvisionerDaemon(dialCtx context.Context, n
if err != nil {
return nil, err
}
server := drpcserver.NewWithOptions(&tracing.DRPCHandler{Handler: mux},
server := drpcsdk.NewServer(logger, &tracing.DRPCHandler{Handler: mux},
drpcserver.Options{
Manager: drpcsdk.DefaultDRPCOptions(nil),
Log: func(err error) {
Expand Down
44 changes: 44 additions & 0 deletions coderd/workspaceagentsrpc_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -7,6 +7,7 @@ import (

"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"storj.io/drpc/drpcerr"

agentproto "github.com/coder/coder/v2/agent/proto"
"github.com/coder/coder/v2/coderd/coderdtest"
Expand All @@ -20,6 +21,8 @@ import (
"github.com/coder/coder/v2/codersdk"
"github.com/coder/coder/v2/codersdk/agentsdk"
"github.com/coder/coder/v2/provisionersdk/proto"
"github.com/coder/coder/v2/tailnet"
tailnetproto "github.com/coder/coder/v2/tailnet/proto"
"github.com/coder/coder/v2/testutil"
)

Expand Down Expand Up @@ -112,6 +115,47 @@ func TestWorkspaceAgentReportStats(t *testing.T) {
}
}

func TestWorkspaceAgentRPC_TailnetMethods(t *testing.T) {
t.Parallel()

ctx := testutil.Context(t, testutil.WaitLong)
client, db := coderdtest.NewWithDatabase(t, nil)
user := coderdtest.CreateFirstUser(t, client)
workspace := dbfake.WorkspaceBuild(t, db, database.WorkspaceTable{
OrganizationID: user.OrganizationID,
OwnerID: user.UserID,
}).WithAgent().Do()

agentClient := agentsdk.New(client.URL, agentsdk.WithFixedToken(workspace.AgentToken))
conn, err := agentClient.ConnectRPC(ctx)
require.NoError(t, err)
t.Cleanup(func() {
_ = conn.Close()
})

tailnetClient := tailnetproto.NewDRPCTailnetClient(conn)
_, err = tailnetClient.RefreshResumeToken(ctx, &tailnetproto.RefreshResumeTokenRequest{})
require.Error(t, err)
require.EqualValues(t, drpcerr.Unimplemented, drpcerr.Code(err))

updates, err := tailnetClient.WorkspaceUpdates(ctx, &tailnetproto.WorkspaceUpdatesRequest{
WorkspaceOwnerId: tailnet.UUIDToByteSlice(user.UserID),
})
if err == nil {
_, err = updates.Recv()
}
require.Error(t, err)
require.EqualValues(t, drpcerr.Unimplemented, drpcerr.Code(err))

telemetry, err := tailnetClient.PostTelemetry(ctx, &tailnetproto.TelemetryRequest{})
require.NoError(t, err)
require.NotNil(t, telemetry)

agentAPI := agentproto.NewDRPCAgentClient(conn)
_, err = agentAPI.GetManifest(ctx, &agentproto.GetManifestRequest{})
require.NoError(t, err)
}

func TestAgentAPI_LargeManifest(t *testing.T) {
t.Parallel()

Expand Down
39 changes: 39 additions & 0 deletions codersdk/drpcsdk/server.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,39 @@
package drpcsdk

import (
"runtime/debug"

"storj.io/drpc"
"storj.io/drpc/drpcserver"

"cdr.dev/slog/v3"
)

// NewServer constructs a dRPC server that recovers panics from RPC handlers.
func NewServer(logger slog.Logger, handler drpc.Handler, options drpcserver.Options) *drpcserver.Server {
return drpcserver.NewWithOptions(&recoverHandler{
logger: logger,
handler: handler,
}, options)
}

type recoverHandler struct {
logger slog.Logger
handler drpc.Handler
}

func (h *recoverHandler) HandleRPC(stream drpc.Stream, rpc string) (err error) {
defer func() {
if r := recover(); r != nil {
h.logger.Error(stream.Context(),
"panic serving dRPC request (recovered)",
slog.F("rpc", rpc),
slog.F("panic", r),
slog.F("stack", string(debug.Stack())),
)
err = drpc.InternalError.New("panic serving dRPC request")
}
}()

return h.handler.HandleRPC(stream, rpc)
}
85 changes: 85 additions & 0 deletions codersdk/drpcsdk/server_internal_test.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,85 @@
package drpcsdk

import (
"context"
"testing"

"github.com/stretchr/testify/require"
"golang.org/x/xerrors"
"storj.io/drpc"

"cdr.dev/slog/v3"
"github.com/coder/coder/v2/testutil"
)

func TestRecoverHandler(t *testing.T) {
t.Parallel()

t.Run("Panic", func(t *testing.T) {
t.Parallel()

const panicValue = "sensitive panic details"
sink := testutil.NewFakeSink(t)
handler := &recoverHandler{
logger: sink.Logger(),
handler: handlerFunc(func(drpc.Stream, string) error {
panic(panicValue)
}),
}

err := handler.HandleRPC(contextStream{ctx: t.Context()}, "/test.Service/Panic")
require.Error(t, err)
require.True(t, drpc.InternalError.Has(err))
require.NotContains(t, err.Error(), panicValue)

entries := sink.Entries()
require.Len(t, entries, 1)
require.Equal(t, slog.LevelError, entries[0].Level)
require.Equal(t, "panic serving dRPC request (recovered)", entries[0].Message)
require.Equal(t, "/test.Service/Panic", fieldValue(entries[0].Fields, "rpc"))
require.Equal(t, panicValue, fieldValue(entries[0].Fields, "panic"))
stackValue := fieldValue(entries[0].Fields, "stack")
stack, ok := stackValue.(string)
require.True(t, ok, "stack field must be a string, got %T", stackValue)
require.Contains(t, stack, "goroutine ")
})

t.Run("Error", func(t *testing.T) {
t.Parallel()

expected := xerrors.New("handler error")
handler := &recoverHandler{
handler: handlerFunc(func(drpc.Stream, string) error {
return expected
}),
}

err := handler.HandleRPC(contextStream{ctx: t.Context()}, "/test.Service/Error")
require.ErrorIs(t, err, expected)
})
}

type handlerFunc func(drpc.Stream, string) error

func (f handlerFunc) HandleRPC(stream drpc.Stream, rpc string) error {
return f(stream, rpc)
}

type contextStream struct {
ctx context.Context
}

func (s contextStream) Context() context.Context { return s.ctx }
func (contextStream) MsgSend(drpc.Message, drpc.Encoding) error { return nil }
func (contextStream) MsgRecv(drpc.Message, drpc.Encoding) error { return nil }
func (contextStream) CloseSend() error { return nil }
func (contextStream) Close() error { return nil }

func fieldValue(fields slog.Map, name string) any {
for _, field := range fields {
if field.Name == name {
return field.Value
}
}
return nil
}
94 changes: 94 additions & 0 deletions codersdk/drpcsdk/server_test.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,94 @@
package drpcsdk_test

import (
"context"
"testing"

"github.com/stretchr/testify/require"
"golang.org/x/xerrors"
"storj.io/drpc"
"storj.io/drpc/drpcserver"

"github.com/coder/coder/v2/codersdk/drpcsdk"
"github.com/coder/coder/v2/testutil"
)

func TestNewServerRecoversPanics(t *testing.T) {
t.Parallel()

const (
panicRPC = "/test.Service/Panic"
echoRPC = "/test.Service/Echo"
panicValue = "sensitive panic details"
)

ctx := testutil.Context(t, testutil.WaitShort)
serverCtx, cancel := context.WithCancel(ctx)
defer cancel()

client, listener := drpcsdk.MemTransportPipe()
defer func() {
_ = client.Close()
_ = listener.Close()
}()

handler := testHandlerFunc(func(stream drpc.Stream, rpc string) error {
switch rpc {
case panicRPC:
panic(panicValue)
case echoRPC:
var message string
if err := stream.MsgRecv(&message, stringEncoding{}); err != nil {
return err
}
return stream.MsgSend(&message, stringEncoding{})
default:
return xerrors.Errorf("unexpected RPC %q", rpc)
}
})
server := drpcsdk.NewServer(testutil.NewFakeSink(t).Logger(), handler, drpcserver.Options{
Manager: drpcsdk.DefaultDRPCOptions(nil),
})
serverDone := make(chan error, 1)
go func() {
serverDone <- server.Serve(serverCtx, listener)
}()

request, response := "request", ""
err := client.Invoke(ctx, panicRPC, stringEncoding{}, &request, &response)
require.EqualError(t, err, "internal error: panic serving dRPC request")
require.NotContains(t, err.Error(), panicValue)

request, response = "healthy", ""
err = client.Invoke(ctx, echoRPC, stringEncoding{}, &request, &response)
require.NoError(t, err)
require.Equal(t, request, response)

cancel()
require.NoError(t, testutil.RequireReceive(ctx, t, serverDone))
}

type testHandlerFunc func(drpc.Stream, string) error

func (f testHandlerFunc) HandleRPC(stream drpc.Stream, rpc string) error {
return f(stream, rpc)
}

type stringEncoding struct{}

func (stringEncoding) Marshal(message drpc.Message) ([]byte, error) {
value, ok := message.(*string)
if !ok {
return nil, xerrors.Errorf("marshal %T: expected *string", message)
}
return []byte(*value), nil
}

func (stringEncoding) Unmarshal(data []byte, message drpc.Message) error {
value, ok := message.(*string)
if !ok {
return xerrors.Errorf("unmarshal %T: expected *string", message)
}
*value = string(data)
return nil
}
2 changes: 1 addition & 1 deletion enterprise/coderd/aibridgeserve.go
Loading
Loading