diff --git a/agent/agentsocket/server.go b/agent/agentsocket/server.go index 380b792da1d..d8cad44347b 100644 --- a/agent/agentsocket/server.go +++ b/agent/agentsocket/server.go @@ -55,7 +55,7 @@ func NewServer(logger slog.Logger, opts ...Option) (*Server, error) { return nil, xerrors.Errorf("failed to register drpc service: %w", err) } - server.drpcServer = drpcserver.NewWithOptions(mux, drpcserver.Options{ + server.drpcServer = drpcsdk.NewServer(logger, mux, drpcserver.Options{ Manager: drpcsdk.DefaultDRPCOptions(nil), Log: func(err error) { if errors.Is(err, context.Canceled) || diff --git a/agent/agenttest/client.go b/agent/agenttest/client.go index 474469d7ff0..24f84050a81 100644 --- a/agent/agenttest/client.go +++ b/agent/agenttest/client.go @@ -77,6 +77,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) { diff --git a/coderd/agentapi/api.go b/coderd/agentapi/api.go index 32d65adee29..d00b67b0c3e 100644 --- a/coderd/agentapi/api.go +++ b/coderd/agentapi/api.go @@ -58,7 +58,7 @@ type API struct { *ConnLogAPI *SubAgentAPI *BoundaryLogsAPI - *tailnet.DRPCService + tailnetService *tailnet.DRPCService cachedWorkspaceFields *CachedWorkspaceFields @@ -67,6 +67,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 @@ -217,7 +239,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, @@ -258,12 +280,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) { diff --git a/coderd/aibridged.go b/coderd/aibridged.go index f448be39d07..5abb08fb699 100644 --- a/coderd/aibridged.go +++ b/coderd/aibridged.go @@ -84,7 +84,7 @@ func (api *API) CreateInMemoryAIBridgeServer(dialCtx context.Context) (client ai if err != nil { return nil, xerrors.Errorf("register key validator service: %w", 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) { diff --git a/coderd/coderd.go b/coderd/coderd.go index de6cbbdacf6..558a9981a87 100644 --- a/coderd/coderd.go +++ b/coderd/coderd.go @@ -2492,7 +2492,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) { diff --git a/coderd/coderd_test.go b/coderd/coderd_test.go index ccf9c8de8fd..2a7f3ad9bf3 100644 --- a/coderd/coderd_test.go +++ b/coderd/coderd_test.go @@ -523,3 +523,51 @@ func TestRateLimitByUser(t *testing.T) { "member should not be able to bypass rate limit") }) } + +// TestRateLimitPathNormalization is a regression test for CDM-02-003 +// (Cure53): a client could bypass a rate limit by inserting redundant +// slashes into the request path. Coder's router still routes the +// respelled path to the same handler as the canonical path, but the rate +// limiter previously keyed its bucket on the raw, un-normalized path, so +// the respelled request landed in a fresh bucket instead of the one +// already exhausted by the canonical path. +func TestRateLimitPathNormalization(t *testing.T) { + t.Parallel() + + const rateLimit = 2 + + client := coderdtest.New(t, &coderdtest.Options{ + LoginRateLimit: rateLimit, + }) + + ctx := testutil.Context(t, testutil.WaitLong) + + post := func(path string) int { + req, err := http.NewRequestWithContext(ctx, http.MethodPost, + client.URL.String()+path, strings.NewReader(`{"password":"hunter2"}`)) + require.NoError(t, err) + req.Header.Set("Content-Type", "application/json") + + resp, err := client.HTTPClient.Do(req) + require.NoError(t, err) + defer resp.Body.Close() + return resp.StatusCode + } + + // Exhaust the limit against the canonical path. + for i := range rateLimit { + require.Equal(t, http.StatusOK, post("/api/v2/users/validate-password"), + "request %d against the canonical path should succeed", i+1) + } + + // The canonical path is now rate limited. + require.Equal(t, http.StatusTooManyRequests, post("/api/v2/users/validate-password"), + "canonical path should be rate limited after exhausting the limit") + + // Respelling the same endpoint with redundant slashes must not grant a + // fresh bucket: it's the same handler, so it must still be limited. + require.Equal(t, http.StatusTooManyRequests, post("/api/v2/users//validate-password"), + "double-slash variant must share the canonical path's rate-limit bucket") + require.Equal(t, http.StatusTooManyRequests, post("/api/v2/users///validate-password"), + "triple-slash variant must share the canonical path's rate-limit bucket") +} diff --git a/coderd/httpmw/ratelimit.go b/coderd/httpmw/ratelimit.go index e89a280530e..17af4be2421 100644 --- a/coderd/httpmw/ratelimit.go +++ b/coderd/httpmw/ratelimit.go @@ -3,6 +3,7 @@ package httpmw import ( "fmt" "net/http" + "path" "strconv" "sync/atomic" "time" @@ -85,7 +86,7 @@ func RateLimit(count int, window time.Duration) func(http.Handler) http.Handler "%q provided but user is not %v", codersdk.BypassRatelimitHeader, rbac.RoleOwner(), ) - }, httprate.KeyByEndpoint), + }, keyByNormalizedEndpoint), httprate.WithLimitHandler(func(w http.ResponseWriter, r *http.Request) { httpapi.Write(r.Context(), w, http.StatusTooManyRequests, codersdk.Response{ Message: fmt.Sprintf("You've been rate limited for sending more than %v requests in %v.", count, window), @@ -94,6 +95,21 @@ func RateLimit(count int, window time.Duration) func(http.Handler) http.Handler ) } +// keyByNormalizedEndpoint mirrors httprate.KeyByEndpoint, but cleans the +// request path first. chi's router tolerates redundant slashes (see +// singleSlashMW in coderd.go) and routes them to the same handler as the +// canonical path, but only normalizes its internal route-matching path, +// not r.URL.Path. Without normalizing here too, a client can respell a +// path, for example inserting an extra slash, to get a fresh rate-limit +// bucket for an endpoint it's already been throttled on. +func keyByNormalizedEndpoint(r *http.Request) (string, error) { + p := r.URL.Path + if p == "" { + p = "/" + } + return path.Clean(p), nil +} + // RateLimitByAuthToken returns a handler that limits requests based on the // authentication token in the request. // diff --git a/coderd/httpmw/ratelimit_test.go b/coderd/httpmw/ratelimit_test.go index 49e46ccf467..88441326c74 100644 --- a/coderd/httpmw/ratelimit_test.go +++ b/coderd/httpmw/ratelimit_test.go @@ -49,6 +49,36 @@ func TestRateLimit(t *testing.T) { } }) + t.Run("PathNormalizationBypass", func(t *testing.T) { + t.Parallel() + rtr := chi.NewRouter() + rtr.Use(httpmw.RateLimit(1, time.Second)) + // A wildcard route so that requests for both the canonical path and + // its redundant-slash variants reach the same handler, mirroring + // how chi's router resolves /api/v2/users//validate-password to the + // same handler as /api/v2/users/validate-password in production. + rtr.Post("/*", func(rw http.ResponseWriter, r *http.Request) { + rw.WriteHeader(http.StatusOK) + }) + + remoteAddr := randRemoteAddr() + paths := []string{ + "/api/v2/users/validate-password", + "/api/v2/users//validate-password", + "/api/v2/users///validate-password", + "/api/v2/users/validate-password", + } + for i, p := range paths { + req := httptest.NewRequest("POST", p, nil) + req.RemoteAddr = remoteAddr + rec := httptest.NewRecorder() + rtr.ServeHTTP(rec, req) + resp := rec.Result() + _ = resp.Body.Close() + require.Equal(t, i != 0, resp.StatusCode == http.StatusTooManyRequests, "request %d (%s)", i, p) + } + }) + t.Run("RandomIPs", func(t *testing.T) { t.Parallel() rtr := chi.NewRouter() diff --git a/coderd/workspaceagentsrpc_test.go b/coderd/workspaceagentsrpc_test.go index 1595462d191..5d083c131f9 100644 --- a/coderd/workspaceagentsrpc_test.go +++ b/coderd/workspaceagentsrpc_test.go @@ -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" @@ -17,6 +18,8 @@ import ( "github.com/coder/coder/v2/coderd/rbac" "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" ) @@ -109,6 +112,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() diff --git a/codersdk/drpcsdk/server.go b/codersdk/drpcsdk/server.go new file mode 100644 index 00000000000..6fc9a505d93 --- /dev/null +++ b/codersdk/drpcsdk/server.go @@ -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) +} diff --git a/codersdk/drpcsdk/server_internal_test.go b/codersdk/drpcsdk/server_internal_test.go new file mode 100644 index 00000000000..66298af68f5 --- /dev/null +++ b/codersdk/drpcsdk/server_internal_test.go @@ -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 +} diff --git a/codersdk/drpcsdk/server_test.go b/codersdk/drpcsdk/server_test.go new file mode 100644 index 00000000000..e906daa82e0 --- /dev/null +++ b/codersdk/drpcsdk/server_test.go @@ -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 +} diff --git a/enterprise/coderd/provisionerdaemons.go b/enterprise/coderd/provisionerdaemons.go index 17a00d22421..f0ace6c7c57 100644 --- a/enterprise/coderd/provisionerdaemons.go +++ b/enterprise/coderd/provisionerdaemons.go @@ -376,7 +376,7 @@ func (api *API) provisionerDaemonServe(rw http.ResponseWriter, r *http.Request) _ = conn.Close(websocket.StatusInternalError, httpapi.WebsocketCloseSprintf("drpc register provisioner daemon: %s", err)) return } - server := drpcserver.NewWithOptions(mux, drpcserver.Options{ + server := drpcsdk.NewServer(logger, mux, drpcserver.Options{ Manager: drpcsdk.DefaultDRPCOptions(nil), Log: func(err error) { if xerrors.Is(err, io.EOF) { diff --git a/provisionersdk/serve.go b/provisionersdk/serve.go index 4afcee96269..013626b22a4 100644 --- a/provisionersdk/serve.go +++ b/provisionersdk/serve.go @@ -92,7 +92,7 @@ func Serve(ctx context.Context, server Server, options *ServeOptions) error { if err != nil { return xerrors.Errorf("register provisioner: %w", err) } - srv := drpcserver.NewWithOptions(&tracing.DRPCHandler{Handler: mux}, drpcserver.Options{ + srv := drpcsdk.NewServer(options.Logger, &tracing.DRPCHandler{Handler: mux}, drpcserver.Options{ Manager: drpcsdk.DefaultDRPCOptions(nil), }) diff --git a/site/src/components/SyntaxHighlighter/SyntaxHighlighter.stories.tsx b/site/src/components/SyntaxHighlighter/SyntaxHighlighter.stories.tsx new file mode 100644 index 00000000000..8740b3af492 --- /dev/null +++ b/site/src/components/SyntaxHighlighter/SyntaxHighlighter.stories.tsx @@ -0,0 +1,111 @@ +import type { Meta, StoryObj } from "@storybook/react-vite"; +import type * as Monaco from "monaco-editor"; +import * as monaco from "monaco-editor"; +import { useState } from "react"; +import { expect, userEvent, waitFor } from "storybook/test"; +import { withDashboardProvider } from "#/testHelpers/storybook"; +import { SyntaxHighlighter } from "./SyntaxHighlighter"; + +const original = `resource "coder_agent" "main" { + os = "linux" + arch = "amd64" +} +`; + +const modified = `resource "coder_agent" "main" { + os = "linux" + arch = "arm64" +} +`; + +// The diff editor's gutter menu and occurrence highlighter register delayed +// disposables whose teardown throws when editors unmount in Storybook tests. +// They are irrelevant to model disposal, so we turn them off in stories to keep +// the test runner clean without changing production behavior. +const stableTeardownOptions: Monaco.editor.IStandaloneDiffEditorConstructionOptions = + { + minimap: { enabled: false }, + renderSideBySide: true, + readOnly: true, + renderGutterMenu: false, + occurrencesHighlight: "off", + }; + +const meta: Meta = { + title: "components/SyntaxHighlighter", + component: SyntaxHighlighter, + decorators: [withDashboardProvider], + args: { + language: "hcl", + editorProps: { options: stableTeardownOptions }, + }, +}; + +export default meta; +type Story = StoryObj; + +export const Plain: Story = { + args: { + value: original, + }, +}; + +export const Diff: Story = { + args: { + value: modified, + compareWith: original, + }, +}; + +// Reproduces the leak from DEVEX-736: a single SyntaxHighlighter instance that +// stays mounted while a file switches between diff and plain across template +// versions. Each diff editor owns two Monaco models, and they must be disposed +// when the diff goes away. Before the fix the models were only disposed on full +// unmount, so toggling diff -> plain -> diff leaked two models per cycle. +const DiffToggle = () => { + const [showDiff, setShowDiff] = useState(true); + return ( +
+ + +
+ ); +}; + +export const DisposesModelsOnDiffToggle: Story = { + render: () => , + play: async ({ canvas }) => { + const toggle = canvas.getByRole("button", { name: "Toggle diff" }); + + // Wait for the diff editor to mount its original + modified models, then + // record the total as a baseline. Every full toggle cycle must return to + // this number; growth would mean abandoned models are being retained. + let baseline = 0; + await waitFor(() => { + baseline = monaco.editor.getModels().length; + expect(baseline).toBeGreaterThanOrEqual(2); + }); + + for (let cycle = 0; cycle < 3; cycle++) { + // Switch to plain: the diff editor unmounts and must dispose its models. + await userEvent.click(toggle); + await waitFor(() => + expect(monaco.editor.getModels().length).toBeLessThan(baseline), + ); + + // Switch back to diff: a new diff editor mounts and the total must land + // back on the baseline rather than climbing. + await userEvent.click(toggle); + await waitFor(() => + expect(monaco.editor.getModels().length).toBe(baseline), + ); + } + }, +}; diff --git a/site/src/components/SyntaxHighlighter/SyntaxHighlighter.tsx b/site/src/components/SyntaxHighlighter/SyntaxHighlighter.tsx index d66d7973951..60e24b69ec4 100644 --- a/site/src/components/SyntaxHighlighter/SyntaxHighlighter.tsx +++ b/site/src/components/SyntaxHighlighter/SyntaxHighlighter.tsx @@ -2,7 +2,13 @@ import { useTheme } from "@emotion/react"; import Editor, { DiffEditor, loader } from "@monaco-editor/react"; import type * as Monaco from "monaco-editor"; import * as monaco from "monaco-editor"; -import { type ComponentProps, type FC, useCallback } from "react"; +import { + type ComponentProps, + type FC, + useCallback, + useEffect, + useRef, +} from "react"; import { useCoderTheme } from "./coderTheme"; loader.config({ monaco }); @@ -38,40 +44,6 @@ export const SyntaxHighlighter: FC = ({ const theme = useTheme(); const coderTheme = useCoderTheme(); - // Auto-scroll to first diff when the diff editor mounts and diffs are computed. - const handleDiffEditorMount = useCallback( - ( - editor: Monaco.editor.IStandaloneDiffEditor, - monacoInstance: typeof Monaco, - ) => { - // Call any existing onMount handler from editorProps. - editorProps?.onMount?.(editor, monacoInstance); - - // Diffs may already be computed by the time onMount fires, - // so check immediately first. If not ready yet, fall back - // to waiting for the onDidUpdateDiff event. - const scrollToFirstDiff = () => { - editor.goToDiff("next"); - }; - - const changes = editor.getLineChanges(); - if (changes && changes.length > 0) { - scrollToFirstDiff(); - return; - } - - const disposable = editor.onDidUpdateDiff(() => { - const updatedChanges = editor.getLineChanges(); - if (!updatedChanges || updatedChanges.length === 0) { - return; - } - disposable.dispose(); - scrollToFirstDiff(); - }); - }, - [editorProps], - ); - const commonProps = { language, theme: coderTheme.name, @@ -99,20 +71,102 @@ export const SyntaxHighlighter: FC = ({ }} > {hasDiff ? ( - + ) : ( )} ); }; + +type DiffFileProps = CommonEditorProps & { + original: string; + modified: string; +}; + +// Renders the diff editor and owns its model cleanup. Scoping this to its own +// component means the cleanup effect runs whenever the diff editor unmounts, +// including when SyntaxHighlighter stays mounted but switches diff -> plain for +// a file that stopped changing between versions. +// +// keepCurrent{Original,Modified}Model stops @monaco-editor/react from disposing +// the models mid-teardown (which throws), so we dispose them ourselves after +// React has torn the editor down. Without this the models accumulate unbounded +// as users open template versions until the tab runs out of memory. +const DiffFile: FC = ({ + original, + modified, + onMount, + ...editorProps +}) => { + const diffModelsRef = useRef<{ + original: Monaco.editor.ITextModel; + modified: Monaco.editor.ITextModel; + } | null>(null); + + const handleMount = useCallback( + ( + editor: Monaco.editor.IStandaloneDiffEditor, + monacoInstance: typeof Monaco, + ) => { + onMount?.(editor, monacoInstance); + + const diffModel = editor.getModel(); + diffModelsRef.current = diffModel + ? { original: diffModel.original, modified: diffModel.modified } + : null; + + // Auto-scroll to the first diff. Diffs may already be computed by the + // time onMount fires, so check immediately and otherwise wait for the + // onDidUpdateDiff event. + const scrollToFirstDiff = () => { + editor.goToDiff("next"); + }; + + const changes = editor.getLineChanges(); + if (changes && changes.length > 0) { + scrollToFirstDiff(); + return; + } + + const disposable = editor.onDidUpdateDiff(() => { + const updatedChanges = editor.getLineChanges(); + if (!updatedChanges || updatedChanges.length === 0) { + return; + } + disposable.dispose(); + scrollToFirstDiff(); + }); + }, + [onMount], + ); + + useEffect(() => { + return () => { + const models = diffModelsRef.current; + if (!models) { + return; + } + diffModelsRef.current = null; + // Defer disposal until after React's commit finishes. @monaco-editor/ + // react disposes the diff widget in its own unmount cleanup; freeing + // the models in the same synchronous teardown makes the widget throw + // "TextModel got disposed before DiffEditorWidget model got reset". + queueMicrotask(() => { + models.original.dispose(); + models.modified.dispose(); + }); + }; + }, []); + + return ( + + ); +}; diff --git a/site/src/modules/apps/apps.test.ts b/site/src/modules/apps/apps.test.ts index 146964af78b..6a9ccd0dcef 100644 --- a/site/src/modules/apps/apps.test.ts +++ b/site/src/modules/apps/apps.test.ts @@ -85,6 +85,24 @@ describe("getAppHref", () => { expect(href).toBe("vscode://example.com?token=user-session-token"); }); + it("replaces the session token for Antigravity IDE URLs", () => { + const externalApp = { + ...MockWorkspaceApp, + external: true, + url: `antigravity-ide://coder.coder-remote/open?token=${SESSION_TOKEN_PLACEHOLDER}`, + }; + const href = getAppHref(externalApp, { + host: "*.apps-host.tld", + path: "/path-base", + agent: MockWorkspaceAgent, + workspace: MockWorkspace, + token: "user-session-token", + }); + expect(href).toBe( + "antigravity-ide://coder.coder-remote/open?token=user-session-token", + ); + }); + it("doesn't return the URL with the session token replaced when using the HTTP protocol", () => { const externalApp = { ...MockWorkspaceApp, diff --git a/site/src/modules/apps/apps.ts b/site/src/modules/apps/apps.ts index d3f5b7c8e28..179358cbd17 100644 --- a/site/src/modules/apps/apps.ts +++ b/site/src/modules/apps/apps.ts @@ -26,6 +26,7 @@ const ALLOWED_EXTERNAL_APP_PROTOCOLS = [ "kiro:", "positron:", "antigravity:", + "antigravity-ide:", ]; type GetVSCodeHrefParams = { diff --git a/tailnet/service.go b/tailnet/service.go index 6ae02876a4d..9773530c76f 100644 --- a/tailnet/service.go +++ b/tailnet/service.go @@ -11,6 +11,7 @@ import ( "github.com/google/uuid" "github.com/hashicorp/yamux" "golang.org/x/xerrors" + "storj.io/drpc/drpcerr" "storj.io/drpc/drpcmux" "storj.io/drpc/drpcserver" "tailscale.com/tailcfg" @@ -92,7 +93,7 @@ func NewClientService(options ClientServiceOptions) ( if err != nil { return nil, xerrors.Errorf("register DRPC service: %w", err) } - server := drpcserver.NewWithOptions(mux, drpcserver.Options{ + server := drpcsdk.NewServer(options.Logger, mux, drpcserver.Options{ Manager: drpcsdk.DefaultDRPCOptions(nil), Log: func(err error) { if xerrors.Is(err, io.EOF) || @@ -185,6 +186,13 @@ func (s *DRPCService) StreamDERPMaps(_ *proto.StreamDERPMapsRequest, stream prot } func (s *DRPCService) RefreshResumeToken(ctx context.Context, _ *proto.RefreshResumeTokenRequest) (*proto.RefreshResumeTokenResponse, error) { + if s.ResumeTokenProvider == nil { + return nil, drpcerr.WithCode( + xerrors.New("resume tokens not supported on this connection"), + drpcerr.Unimplemented, + ) + } + streamID, ok := ctx.Value(streamIDContextKey{}).(StreamID) if !ok { return nil, xerrors.New("no Stream ID") @@ -219,6 +227,13 @@ func (s *DRPCService) Coordinate(stream proto.DRPCTailnet_CoordinateStream) erro } func (s *DRPCService) WorkspaceUpdates(req *proto.WorkspaceUpdatesRequest, stream proto.DRPCTailnet_WorkspaceUpdatesStream) error { + if s.WorkspaceUpdatesProvider == nil { + return drpcerr.WithCode( + xerrors.New("workspace updates not supported on this connection"), + drpcerr.Unimplemented, + ) + } + defer stream.Close() ctx := stream.Context() diff --git a/tailnet/service_test.go b/tailnet/service_test.go index 0c268b05edb..21a290ef358 100644 --- a/tailnet/service_test.go +++ b/tailnet/service_test.go @@ -13,6 +13,7 @@ import ( "github.com/stretchr/testify/require" "go.uber.org/mock/gomock" "golang.org/x/xerrors" + "storj.io/drpc/drpcerr" "tailscale.com/tailcfg" "github.com/coder/coder/v2/tailnet" @@ -178,6 +179,27 @@ func TestClientService_ServeClient_V1(t *testing.T) { require.ErrorIs(t, err, tailnet.ErrUnsupportedVersion) } +func TestClientService_UnsupportedProviders(t *testing.T) { + t.Parallel() + + ctx := testutil.Context(t, testutil.WaitShort) + clientID := uuid.New() + _, client := createUpdateService(t, ctx, clientID, nil) + + _, err := client.RefreshResumeToken(ctx, &proto.RefreshResumeTokenRequest{}) + require.ErrorContains(t, err, "resume tokens not supported on this connection") + require.EqualValues(t, drpcerr.Unimplemented, drpcerr.Code(err)) + + updates, err := client.WorkspaceUpdates(ctx, &proto.WorkspaceUpdatesRequest{ + WorkspaceOwnerId: tailnet.UUIDToByteSlice(clientID), + }) + if err == nil { + _, err = updates.Recv() + } + require.ErrorContains(t, err, "workspace updates not supported on this connection") + require.EqualValues(t, drpcerr.Unimplemented, drpcerr.Code(err)) +} + func TestNetworkTelemetryBatcher(t *testing.T) { t.Parallel()