Thanks to visit codestin.com
Credit goes to github.com

Skip to content

Commit a6bab62

Browse files
authored
fix: prevent coderd crashes from unsupported Tailnet RPCs (#28429) (#28965)
Backport of #28429 Original PR: #28429 — fix: prevent coderd crashes from unsupported Tailnet RPCs Merge commit: 156c8b6 PR created manually because the backport CI job is timing out. (Resolved) conflicts: - coderd/agentapi/api.go - enterprise/coderd/aibridgeserve.go
1 parent ac2323f commit a6bab62

13 files changed

Lines changed: 334 additions & 10 deletions

File tree

agent/agentsocket/server.go

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -55,7 +55,7 @@ func NewServer(logger slog.Logger, opts ...Option) (*Server, error) {
5555
return nil, xerrors.Errorf("failed to register drpc service: %w", err)
5656
}
5757

58-
server.drpcServer = drpcserver.NewWithOptions(mux, drpcserver.Options{
58+
server.drpcServer = drpcsdk.NewServer(logger, mux, drpcserver.Options{
5959
Manager: drpcsdk.DefaultDRPCOptions(nil),
6060
Log: func(err error) {
6161
if errors.Is(err, context.Canceled) ||

agent/agenttest/client.go

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -77,6 +77,7 @@ func NewClientWithSecrets(t testing.TB,
7777
fakeAAPI := NewFakeAgentAPI(t, logger, mp, statsChan)
7878
err = agentproto.DRPCRegisterAgent(mux, fakeAAPI)
7979
require.NoError(t, err)
80+
// Keep panics unrecovered in this test server so they fail tests loudly.
8081
server := drpcserver.NewWithOptions(mux, drpcserver.Options{
8182
Manager: drpcsdk.DefaultDRPCOptions(nil),
8283
Log: func(err error) {

coderd/agentapi/api.go

Lines changed: 28 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -58,7 +58,7 @@ type API struct {
5858
*ConnLogAPI
5959
*SubAgentAPI
6060
*BoundaryLogsAPI
61-
*tailnet.DRPCService
61+
tailnetService *tailnet.DRPCService
6262

6363
cachedWorkspaceFields *CachedWorkspaceFields
6464

@@ -67,6 +67,28 @@ type API struct {
6767

6868
var _ agentproto.DRPCAgentServer = &API{}
6969

70+
// agentTailnetService exposes only Tailnet RPCs intended for workspace agents.
71+
// Other current and future RPCs remain unavailable until explicitly forwarded.
72+
type agentTailnetService struct {
73+
tailnetproto.DRPCTailnetUnimplementedServer
74+
75+
service *tailnet.DRPCService
76+
}
77+
78+
func (s *agentTailnetService) PostTelemetry(ctx context.Context, req *tailnetproto.TelemetryRequest) (*tailnetproto.TelemetryResponse, error) {
79+
return s.service.PostTelemetry(ctx, req)
80+
}
81+
82+
func (s *agentTailnetService) StreamDERPMaps(req *tailnetproto.StreamDERPMapsRequest, stream tailnetproto.DRPCTailnet_StreamDERPMapsStream) error {
83+
return s.service.StreamDERPMaps(req, stream)
84+
}
85+
86+
func (s *agentTailnetService) Coordinate(stream tailnetproto.DRPCTailnet_CoordinateStream) error {
87+
return s.service.Coordinate(stream)
88+
}
89+
90+
var _ tailnetproto.DRPCTailnetServer = (*agentTailnetService)(nil)
91+
7092
type Options struct {
7193
AgentID uuid.UUID
7294
OwnerID uuid.UUID
@@ -217,7 +239,7 @@ func New(opts Options, workspace database.Workspace, agent database.WorkspaceAge
217239
Log: opts.Log,
218240
}
219241

220-
api.DRPCService = &tailnet.DRPCService{
242+
api.tailnetService = &tailnet.DRPCService{
221243
CoordPtr: opts.TailnetCoordinator,
222244
Logger: opts.Log,
223245
DerpMapUpdateFrequency: opts.DerpMapUpdateFrequency,
@@ -258,12 +280,14 @@ func (a *API) Server(ctx context.Context) (*drpcserver.Server, error) {
258280
return nil, xerrors.Errorf("register agent API protocol in DRPC mux: %w", err)
259281
}
260282

261-
err = tailnetproto.DRPCRegisterTailnet(mux, a)
283+
err = tailnetproto.DRPCRegisterTailnet(mux, &agentTailnetService{
284+
service: a.tailnetService,
285+
})
262286
if err != nil {
263287
return nil, xerrors.Errorf("register tailnet API protocol in DRPC mux: %w", err)
264288
}
265289

266-
return drpcserver.NewWithOptions(&tracing.DRPCHandler{Handler: mux},
290+
return drpcsdk.NewServer(a.opts.Log, &tracing.DRPCHandler{Handler: mux},
267291
drpcserver.Options{
268292
Manager: drpcsdk.DefaultDRPCOptions(nil),
269293
Log: func(err error) {

coderd/aibridged.go

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -84,7 +84,7 @@ func (api *API) CreateInMemoryAIBridgeServer(dialCtx context.Context) (client ai
8484
if err != nil {
8585
return nil, xerrors.Errorf("register key validator service: %w", err)
8686
}
87-
server := drpcserver.NewWithOptions(&tracing.DRPCHandler{Handler: mux},
87+
server := drpcsdk.NewServer(api.Logger, &tracing.DRPCHandler{Handler: mux},
8888
drpcserver.Options{
8989
Manager: drpcsdk.DefaultDRPCOptions(nil),
9090
Log: func(err error) {

coderd/coderd.go

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -2492,7 +2492,7 @@ func (api *API) CreateInMemoryTaggedProvisionerDaemon(dialCtx context.Context, n
24922492
if err != nil {
24932493
return nil, err
24942494
}
2495-
server := drpcserver.NewWithOptions(&tracing.DRPCHandler{Handler: mux},
2495+
server := drpcsdk.NewServer(logger, &tracing.DRPCHandler{Handler: mux},
24962496
drpcserver.Options{
24972497
Manager: drpcsdk.DefaultDRPCOptions(nil),
24982498
Log: func(err error) {

coderd/workspaceagentsrpc_test.go

Lines changed: 44 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -7,6 +7,7 @@ import (
77

88
"github.com/stretchr/testify/assert"
99
"github.com/stretchr/testify/require"
10+
"storj.io/drpc/drpcerr"
1011

1112
agentproto "github.com/coder/coder/v2/agent/proto"
1213
"github.com/coder/coder/v2/coderd/coderdtest"
@@ -17,6 +18,8 @@ import (
1718
"github.com/coder/coder/v2/coderd/rbac"
1819
"github.com/coder/coder/v2/codersdk/agentsdk"
1920
"github.com/coder/coder/v2/provisionersdk/proto"
21+
"github.com/coder/coder/v2/tailnet"
22+
tailnetproto "github.com/coder/coder/v2/tailnet/proto"
2023
"github.com/coder/coder/v2/testutil"
2124
)
2225

@@ -109,6 +112,47 @@ func TestWorkspaceAgentReportStats(t *testing.T) {
109112
}
110113
}
111114

115+
func TestWorkspaceAgentRPC_TailnetMethods(t *testing.T) {
116+
t.Parallel()
117+
118+
ctx := testutil.Context(t, testutil.WaitLong)
119+
client, db := coderdtest.NewWithDatabase(t, nil)
120+
user := coderdtest.CreateFirstUser(t, client)
121+
workspace := dbfake.WorkspaceBuild(t, db, database.WorkspaceTable{
122+
OrganizationID: user.OrganizationID,
123+
OwnerID: user.UserID,
124+
}).WithAgent().Do()
125+
126+
agentClient := agentsdk.New(client.URL, agentsdk.WithFixedToken(workspace.AgentToken))
127+
conn, err := agentClient.ConnectRPC(ctx)
128+
require.NoError(t, err)
129+
t.Cleanup(func() {
130+
_ = conn.Close()
131+
})
132+
133+
tailnetClient := tailnetproto.NewDRPCTailnetClient(conn)
134+
_, err = tailnetClient.RefreshResumeToken(ctx, &tailnetproto.RefreshResumeTokenRequest{})
135+
require.Error(t, err)
136+
require.EqualValues(t, drpcerr.Unimplemented, drpcerr.Code(err))
137+
138+
updates, err := tailnetClient.WorkspaceUpdates(ctx, &tailnetproto.WorkspaceUpdatesRequest{
139+
WorkspaceOwnerId: tailnet.UUIDToByteSlice(user.UserID),
140+
})
141+
if err == nil {
142+
_, err = updates.Recv()
143+
}
144+
require.Error(t, err)
145+
require.EqualValues(t, drpcerr.Unimplemented, drpcerr.Code(err))
146+
147+
telemetry, err := tailnetClient.PostTelemetry(ctx, &tailnetproto.TelemetryRequest{})
148+
require.NoError(t, err)
149+
require.NotNil(t, telemetry)
150+
151+
agentAPI := agentproto.NewDRPCAgentClient(conn)
152+
_, err = agentAPI.GetManifest(ctx, &agentproto.GetManifestRequest{})
153+
require.NoError(t, err)
154+
}
155+
112156
func TestAgentAPI_LargeManifest(t *testing.T) {
113157
t.Parallel()
114158

codersdk/drpcsdk/server.go

Lines changed: 39 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,39 @@
1+
package drpcsdk
2+
3+
import (
4+
"runtime/debug"
5+
6+
"storj.io/drpc"
7+
"storj.io/drpc/drpcserver"
8+
9+
"cdr.dev/slog/v3"
10+
)
11+
12+
// NewServer constructs a dRPC server that recovers panics from RPC handlers.
13+
func NewServer(logger slog.Logger, handler drpc.Handler, options drpcserver.Options) *drpcserver.Server {
14+
return drpcserver.NewWithOptions(&recoverHandler{
15+
logger: logger,
16+
handler: handler,
17+
}, options)
18+
}
19+
20+
type recoverHandler struct {
21+
logger slog.Logger
22+
handler drpc.Handler
23+
}
24+
25+
func (h *recoverHandler) HandleRPC(stream drpc.Stream, rpc string) (err error) {
26+
defer func() {
27+
if r := recover(); r != nil {
28+
h.logger.Error(stream.Context(),
29+
"panic serving dRPC request (recovered)",
30+
slog.F("rpc", rpc),
31+
slog.F("panic", r),
32+
slog.F("stack", string(debug.Stack())),
33+
)
34+
err = drpc.InternalError.New("panic serving dRPC request")
35+
}
36+
}()
37+
38+
return h.handler.HandleRPC(stream, rpc)
39+
}
Lines changed: 85 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,85 @@
1+
package drpcsdk
2+
3+
import (
4+
"context"
5+
"testing"
6+
7+
"github.com/stretchr/testify/require"
8+
"golang.org/x/xerrors"
9+
"storj.io/drpc"
10+
11+
"cdr.dev/slog/v3"
12+
"github.com/coder/coder/v2/testutil"
13+
)
14+
15+
func TestRecoverHandler(t *testing.T) {
16+
t.Parallel()
17+
18+
t.Run("Panic", func(t *testing.T) {
19+
t.Parallel()
20+
21+
const panicValue = "sensitive panic details"
22+
sink := testutil.NewFakeSink(t)
23+
handler := &recoverHandler{
24+
logger: sink.Logger(),
25+
handler: handlerFunc(func(drpc.Stream, string) error {
26+
panic(panicValue)
27+
}),
28+
}
29+
30+
err := handler.HandleRPC(contextStream{ctx: t.Context()}, "/test.Service/Panic")
31+
require.Error(t, err)
32+
require.True(t, drpc.InternalError.Has(err))
33+
require.NotContains(t, err.Error(), panicValue)
34+
35+
entries := sink.Entries()
36+
require.Len(t, entries, 1)
37+
require.Equal(t, slog.LevelError, entries[0].Level)
38+
require.Equal(t, "panic serving dRPC request (recovered)", entries[0].Message)
39+
require.Equal(t, "/test.Service/Panic", fieldValue(entries[0].Fields, "rpc"))
40+
require.Equal(t, panicValue, fieldValue(entries[0].Fields, "panic"))
41+
stackValue := fieldValue(entries[0].Fields, "stack")
42+
stack, ok := stackValue.(string)
43+
require.True(t, ok, "stack field must be a string, got %T", stackValue)
44+
require.Contains(t, stack, "goroutine ")
45+
})
46+
47+
t.Run("Error", func(t *testing.T) {
48+
t.Parallel()
49+
50+
expected := xerrors.New("handler error")
51+
handler := &recoverHandler{
52+
handler: handlerFunc(func(drpc.Stream, string) error {
53+
return expected
54+
}),
55+
}
56+
57+
err := handler.HandleRPC(contextStream{ctx: t.Context()}, "/test.Service/Error")
58+
require.ErrorIs(t, err, expected)
59+
})
60+
}
61+
62+
type handlerFunc func(drpc.Stream, string) error
63+
64+
func (f handlerFunc) HandleRPC(stream drpc.Stream, rpc string) error {
65+
return f(stream, rpc)
66+
}
67+
68+
type contextStream struct {
69+
ctx context.Context
70+
}
71+
72+
func (s contextStream) Context() context.Context { return s.ctx }
73+
func (contextStream) MsgSend(drpc.Message, drpc.Encoding) error { return nil }
74+
func (contextStream) MsgRecv(drpc.Message, drpc.Encoding) error { return nil }
75+
func (contextStream) CloseSend() error { return nil }
76+
func (contextStream) Close() error { return nil }
77+
78+
func fieldValue(fields slog.Map, name string) any {
79+
for _, field := range fields {
80+
if field.Name == name {
81+
return field.Value
82+
}
83+
}
84+
return nil
85+
}

codersdk/drpcsdk/server_test.go

Lines changed: 94 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,94 @@
1+
package drpcsdk_test
2+
3+
import (
4+
"context"
5+
"testing"
6+
7+
"github.com/stretchr/testify/require"
8+
"golang.org/x/xerrors"
9+
"storj.io/drpc"
10+
"storj.io/drpc/drpcserver"
11+
12+
"github.com/coder/coder/v2/codersdk/drpcsdk"
13+
"github.com/coder/coder/v2/testutil"
14+
)
15+
16+
func TestNewServerRecoversPanics(t *testing.T) {
17+
t.Parallel()
18+
19+
const (
20+
panicRPC = "/test.Service/Panic"
21+
echoRPC = "/test.Service/Echo"
22+
panicValue = "sensitive panic details"
23+
)
24+
25+
ctx := testutil.Context(t, testutil.WaitShort)
26+
serverCtx, cancel := context.WithCancel(ctx)
27+
defer cancel()
28+
29+
client, listener := drpcsdk.MemTransportPipe()
30+
defer func() {
31+
_ = client.Close()
32+
_ = listener.Close()
33+
}()
34+
35+
handler := testHandlerFunc(func(stream drpc.Stream, rpc string) error {
36+
switch rpc {
37+
case panicRPC:
38+
panic(panicValue)
39+
case echoRPC:
40+
var message string
41+
if err := stream.MsgRecv(&message, stringEncoding{}); err != nil {
42+
return err
43+
}
44+
return stream.MsgSend(&message, stringEncoding{})
45+
default:
46+
return xerrors.Errorf("unexpected RPC %q", rpc)
47+
}
48+
})
49+
server := drpcsdk.NewServer(testutil.NewFakeSink(t).Logger(), handler, drpcserver.Options{
50+
Manager: drpcsdk.DefaultDRPCOptions(nil),
51+
})
52+
serverDone := make(chan error, 1)
53+
go func() {
54+
serverDone <- server.Serve(serverCtx, listener)
55+
}()
56+
57+
request, response := "request", ""
58+
err := client.Invoke(ctx, panicRPC, stringEncoding{}, &request, &response)
59+
require.EqualError(t, err, "internal error: panic serving dRPC request")
60+
require.NotContains(t, err.Error(), panicValue)
61+
62+
request, response = "healthy", ""
63+
err = client.Invoke(ctx, echoRPC, stringEncoding{}, &request, &response)
64+
require.NoError(t, err)
65+
require.Equal(t, request, response)
66+
67+
cancel()
68+
require.NoError(t, testutil.RequireReceive(ctx, t, serverDone))
69+
}
70+
71+
type testHandlerFunc func(drpc.Stream, string) error
72+
73+
func (f testHandlerFunc) HandleRPC(stream drpc.Stream, rpc string) error {
74+
return f(stream, rpc)
75+
}
76+
77+
type stringEncoding struct{}
78+
79+
func (stringEncoding) Marshal(message drpc.Message) ([]byte, error) {
80+
value, ok := message.(*string)
81+
if !ok {
82+
return nil, xerrors.Errorf("marshal %T: expected *string", message)
83+
}
84+
return []byte(*value), nil
85+
}
86+
87+
func (stringEncoding) Unmarshal(data []byte, message drpc.Message) error {
88+
value, ok := message.(*string)
89+
if !ok {
90+
return xerrors.Errorf("unmarshal %T: expected *string", message)
91+
}
92+
*value = string(data)
93+
return nil
94+
}

enterprise/coderd/provisionerdaemons.go

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -376,7 +376,7 @@ func (api *API) provisionerDaemonServe(rw http.ResponseWriter, r *http.Request)
376376
_ = conn.Close(websocket.StatusInternalError, httpapi.WebsocketCloseSprintf("drpc register provisioner daemon: %s", err))
377377
return
378378
}
379-
server := drpcserver.NewWithOptions(mux, drpcserver.Options{
379+
server := drpcsdk.NewServer(logger, mux, drpcserver.Options{
380380
Manager: drpcsdk.DefaultDRPCOptions(nil),
381381
Log: func(err error) {
382382
if xerrors.Is(err, io.EOF) {

0 commit comments

Comments
 (0)