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

Skip to content

Commit 156c8b6

Browse files
authored
fix: prevent coderd crashes from unsupported Tailnet RPCs (#28429)
Workspace agents could call RefreshResumeToken and WorkspaceUpdates, even though agent connections do not configure the providers those RPCs require. The resulting nil dereference panicked a DRPC handler and crashed coderd. Limit the agent Tailnet surface to its supported RPCs, return `Unimplemented` when optional providers are absent, and recover panics in Coder-owned DRPC servers so they return an internal error instead of crashing the process. Refs: https://linear.app/codercom/issue/PLAT-464
1 parent 68ca033 commit 156c8b6

14 files changed

Lines changed: 335 additions & 11 deletions

File tree

agent/agentsocket/server.go

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

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

agent/agenttest/client.go

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

coderd/agentapi/api.go

Lines changed: 28 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -59,7 +59,7 @@ type API struct {
5959
*SubAgentAPI
6060
*BoundaryLogsAPI
6161
*ContextAPI
62-
*tailnet.DRPCService
62+
tailnetService *tailnet.DRPCService
6363

6464
cachedWorkspaceFields *CachedWorkspaceFields
6565

@@ -68,6 +68,28 @@ type API struct {
6868

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

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

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

282-
err = tailnetproto.DRPCRegisterTailnet(mux, a)
304+
err = tailnetproto.DRPCRegisterTailnet(mux, &agentTailnetService{
305+
service: a.tailnetService,
306+
})
283307
if err != nil {
284308
return nil, xerrors.Errorf("register tailnet API protocol in DRPC mux: %w", err)
285309
}
286310

287-
return drpcserver.NewWithOptions(&tracing.DRPCHandler{Handler: mux},
311+
return drpcsdk.NewServer(a.opts.Log, &tracing.DRPCHandler{Handler: mux},
288312
drpcserver.Options{
289313
Manager: drpcsdk.DefaultDRPCOptions(nil),
290314
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 := aibridgedserver.Register(mux, srv); err != nil {
8585
return nil, 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
@@ -2587,7 +2587,7 @@ func (api *API) CreateInMemoryTaggedProvisionerDaemon(dialCtx context.Context, n
25872587
if err != nil {
25882588
return nil, err
25892589
}
2590-
server := drpcserver.NewWithOptions(&tracing.DRPCHandler{Handler: mux},
2590+
server := drpcsdk.NewServer(logger, &tracing.DRPCHandler{Handler: mux},
25912591
drpcserver.Options{
25922592
Manager: drpcsdk.DefaultDRPCOptions(nil),
25932593
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"
@@ -20,6 +21,8 @@ import (
2021
"github.com/coder/coder/v2/codersdk"
2122
"github.com/coder/coder/v2/codersdk/agentsdk"
2223
"github.com/coder/coder/v2/provisionersdk/proto"
24+
"github.com/coder/coder/v2/tailnet"
25+
tailnetproto "github.com/coder/coder/v2/tailnet/proto"
2326
"github.com/coder/coder/v2/testutil"
2427
)
2528

@@ -112,6 +115,47 @@ func TestWorkspaceAgentReportStats(t *testing.T) {
112115
}
113116
}
114117

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

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/aibridgeserve.go

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -159,7 +159,7 @@ func (api *API) aiGatewayServe(rw http.ResponseWriter, r *http.Request) {
159159
return
160160
}
161161

162-
server := drpcserver.NewWithOptions(&tracing.DRPCHandler{Handler: mux},
162+
server := drpcsdk.NewServer(logger, &tracing.DRPCHandler{Handler: mux},
163163
drpcserver.Options{
164164
Manager: drpcsdk.DefaultDRPCOptions(nil),
165165
Log: func(err error) {

0 commit comments

Comments
 (0)