From 13bb844fdfe76faed35a67a675c495d4badf943a Mon Sep 17 00:00:00 2001 From: Jon Ayers Date: Tue, 23 Jun 2026 17:44:51 +0000 Subject: [PATCH 1/5] fix(coderd): scope provisioner module file downloads to the daemon's org The DownloadFile gRPC handler authorized module-file downloads using only file.CreatedBy == uuid.Nil and Mimetype == application/x-tar, which every organization's cached Terraform module archive satisfies. A provisioner key scoped to one organization could supply another organization's file UUID and receive its full Terraform module source. Add HasTemplateVersionsUsingCachedModuleFileInOrg to verify the requested module file is referenced by a template version in the calling daemon's organization, and gate the download on it. Cross-org requests return the same error as the metadata check so the handler does not confirm file existence in other organizations. Ref: ANT-2026-22440 --- coderd/database/dbauthz/dbauthz.go | 10 + coderd/database/dbauthz/dbauthz_test.go | 5 + coderd/database/dbmetrics/querymetrics.go | 8 + coderd/database/dbmock/dbmock.go | 15 ++ coderd/database/querier.go | 5 + coderd/database/queries.sql.go | 27 +++ .../templateversionterraformvalues.sql | 14 ++ .../provisionerdserver/download_file_test.go | 187 ++++++++++++++++++ .../provisionerdserver/provisionerdserver.go | 16 ++ 9 files changed, 287 insertions(+) create mode 100644 coderd/provisionerdserver/download_file_test.go diff --git a/coderd/database/dbauthz/dbauthz.go b/coderd/database/dbauthz/dbauthz.go index 5fe1cd7e995..4d4d6fb281c 100644 --- a/coderd/database/dbauthz/dbauthz.go +++ b/coderd/database/dbauthz/dbauthz.go @@ -5697,6 +5697,16 @@ func (q *querier) GetWorkspacesForWorkspaceMetrics(ctx context.Context) ([]datab return q.db.GetWorkspacesForWorkspaceMetrics(ctx) } +func (q *querier) HasTemplateVersionsUsingCachedModuleFileInOrg(ctx context.Context, arg database.HasTemplateVersionsUsingCachedModuleFileInOrgParams) (bool, error) { + // This is an internal cross-table guard used to authorize provisioner + // module-file downloads. Tenant isolation comes from the organization_id + // filter in the query itself; only system-level actors may call it. + if err := q.authorizeContext(ctx, policy.ActionRead, rbac.ResourceSystem); err != nil { + return false, err + } + return q.db.HasTemplateVersionsUsingCachedModuleFileInOrg(ctx, arg) +} + func (q *querier) HydrateAgentChatsContext(ctx context.Context, arg database.HydrateAgentChatsContextParams) error { // System-level operation: an agent context push fans hydration out // across every not-yet-pinned chat for the agent, so it authorizes at diff --git a/coderd/database/dbauthz/dbauthz_test.go b/coderd/database/dbauthz/dbauthz_test.go index 0cc603d4c5c..50cf7524ace 100644 --- a/coderd/database/dbauthz/dbauthz_test.go +++ b/coderd/database/dbauthz/dbauthz_test.go @@ -2674,6 +2674,11 @@ func (s *MethodTestSuite) TestTemplate() { dbm.EXPECT().GetTemplateVersionTerraformValues(gomock.Any(), tv.ID).Return(val, nil).AnyTimes() check.Args(tv.ID).Asserts(t, policy.ActionRead) })) + s.Run("HasTemplateVersionsUsingCachedModuleFileInOrg", s.Mocked(func(dbm *dbmock.MockStore, _ *gofakeit.Faker, check *expects) { + arg := database.HasTemplateVersionsUsingCachedModuleFileInOrgParams{FileID: uuid.New(), OrganizationID: uuid.New()} + dbm.EXPECT().HasTemplateVersionsUsingCachedModuleFileInOrg(gomock.Any(), arg).Return(true, nil).AnyTimes() + check.Args(arg).Asserts(rbac.ResourceSystem, policy.ActionRead).Returns(true) + })) s.Run("GetTemplateVersionVariables", s.Mocked(func(dbm *dbmock.MockStore, faker *gofakeit.Faker, check *expects) { t1 := testutil.Fake(s.T(), faker, database.Template{}) tv := testutil.Fake(s.T(), faker, database.TemplateVersion{TemplateID: uuid.NullUUID{UUID: t1.ID, Valid: true}}) diff --git a/coderd/database/dbmetrics/querymetrics.go b/coderd/database/dbmetrics/querymetrics.go index 40820d9cfb8..8fc504de35b 100644 --- a/coderd/database/dbmetrics/querymetrics.go +++ b/coderd/database/dbmetrics/querymetrics.go @@ -3874,6 +3874,14 @@ func (m queryMetricsStore) GetWorkspacesForWorkspaceMetrics(ctx context.Context) return r0, r1 } +func (m queryMetricsStore) HasTemplateVersionsUsingCachedModuleFileInOrg(ctx context.Context, arg database.HasTemplateVersionsUsingCachedModuleFileInOrgParams) (bool, error) { + start := time.Now() + r0, r1 := m.s.HasTemplateVersionsUsingCachedModuleFileInOrg(ctx, arg) + m.queryLatencies.WithLabelValues("HasTemplateVersionsUsingCachedModuleFileInOrg").Observe(time.Since(start).Seconds()) + m.queryCounts.WithLabelValues(httpmw.ExtractHTTPRoute(ctx), httpmw.ExtractHTTPMethod(ctx), "HasTemplateVersionsUsingCachedModuleFileInOrg").Inc() + return r0, r1 +} + func (m queryMetricsStore) HydrateAgentChatsContext(ctx context.Context, arg database.HydrateAgentChatsContextParams) error { start := time.Now() r0 := m.s.HydrateAgentChatsContext(ctx, arg) diff --git a/coderd/database/dbmock/dbmock.go b/coderd/database/dbmock/dbmock.go index e693e675ee1..2051eed4ec6 100644 --- a/coderd/database/dbmock/dbmock.go +++ b/coderd/database/dbmock/dbmock.go @@ -7242,6 +7242,21 @@ func (mr *MockStoreMockRecorder) GetWorkspacesForWorkspaceMetrics(ctx any) *gomo return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetWorkspacesForWorkspaceMetrics", reflect.TypeOf((*MockStore)(nil).GetWorkspacesForWorkspaceMetrics), ctx) } +// HasTemplateVersionsUsingCachedModuleFileInOrg mocks base method. +func (m *MockStore) HasTemplateVersionsUsingCachedModuleFileInOrg(ctx context.Context, arg database.HasTemplateVersionsUsingCachedModuleFileInOrgParams) (bool, error) { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "HasTemplateVersionsUsingCachedModuleFileInOrg", ctx, arg) + ret0, _ := ret[0].(bool) + ret1, _ := ret[1].(error) + return ret0, ret1 +} + +// HasTemplateVersionsUsingCachedModuleFileInOrg indicates an expected call of HasTemplateVersionsUsingCachedModuleFileInOrg. +func (mr *MockStoreMockRecorder) HasTemplateVersionsUsingCachedModuleFileInOrg(ctx, arg any) *gomock.Call { + mr.mock.ctrl.T.Helper() + return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "HasTemplateVersionsUsingCachedModuleFileInOrg", reflect.TypeOf((*MockStore)(nil).HasTemplateVersionsUsingCachedModuleFileInOrg), ctx, arg) +} + // HydrateAgentChatsContext mocks base method. func (m *MockStore) HydrateAgentChatsContext(ctx context.Context, arg database.HydrateAgentChatsContextParams) error { m.ctrl.T.Helper() diff --git a/coderd/database/querier.go b/coderd/database/querier.go index 2b1b93b4ac4..27ec63a4e1b 100644 --- a/coderd/database/querier.go +++ b/coderd/database/querier.go @@ -975,6 +975,11 @@ type sqlcQuerier interface { GetWorkspacesByTemplateID(ctx context.Context, templateID uuid.UUID) ([]WorkspaceTable, error) GetWorkspacesEligibleForTransition(ctx context.Context, now time.Time) ([]GetWorkspacesEligibleForTransitionRow, error) GetWorkspacesForWorkspaceMetrics(ctx context.Context) ([]GetWorkspacesForWorkspaceMetricsRow, error) + // Reports whether the given file is referenced as cached module files by any + // template version in the given organization. Used to authorize provisioner + // module-file downloads so a daemon cannot read another organization's cached + // Terraform module source. + HasTemplateVersionsUsingCachedModuleFileInOrg(ctx context.Context, arg HasTemplateVersionsUsingCachedModuleFileInOrgParams) (bool, error) // Stamps the pinned hash and error on every not-yet-hydrated chat for // an agent (context_aggregate_hash IS NULL) and copies the agent's // current context resources onto those chats in the same statement, so diff --git a/coderd/database/queries.sql.go b/coderd/database/queries.sql.go index 02e49dcc5ef..6e289e4dd76 100644 --- a/coderd/database/queries.sql.go +++ b/coderd/database/queries.sql.go @@ -27695,6 +27695,33 @@ func (q *sqlQuerier) GetTemplateVersionTerraformValues(ctx context.Context, temp return i, err } +const hasTemplateVersionsUsingCachedModuleFileInOrg = `-- name: HasTemplateVersionsUsingCachedModuleFileInOrg :one +SELECT EXISTS ( + SELECT 1 + FROM template_version_terraform_values tvtv + JOIN template_versions tv + ON tv.id = tvtv.template_version_id + WHERE tvtv.cached_module_files = $1::uuid + AND tv.organization_id = $2::uuid +) +` + +type HasTemplateVersionsUsingCachedModuleFileInOrgParams struct { + FileID uuid.UUID `db:"file_id" json:"file_id"` + OrganizationID uuid.UUID `db:"organization_id" json:"organization_id"` +} + +// Reports whether the given file is referenced as cached module files by any +// template version in the given organization. Used to authorize provisioner +// module-file downloads so a daemon cannot read another organization's cached +// Terraform module source. +func (q *sqlQuerier) HasTemplateVersionsUsingCachedModuleFileInOrg(ctx context.Context, arg HasTemplateVersionsUsingCachedModuleFileInOrgParams) (bool, error) { + row := q.db.QueryRowContext(ctx, hasTemplateVersionsUsingCachedModuleFileInOrg, arg.FileID, arg.OrganizationID) + var exists bool + err := row.Scan(&exists) + return exists, err +} + const insertTemplateVersionTerraformValuesByJobID = `-- name: InsertTemplateVersionTerraformValuesByJobID :exec INSERT INTO template_version_terraform_values ( diff --git a/coderd/database/queries/templateversionterraformvalues.sql b/coderd/database/queries/templateversionterraformvalues.sql index 2ded4a26753..97b20a5103c 100644 --- a/coderd/database/queries/templateversionterraformvalues.sql +++ b/coderd/database/queries/templateversionterraformvalues.sql @@ -23,3 +23,17 @@ VALUES @updated_at, @provisionerd_version ); + +-- name: HasTemplateVersionsUsingCachedModuleFileInOrg :one +-- Reports whether the given file is referenced as cached module files by any +-- template version in the given organization. Used to authorize provisioner +-- module-file downloads so a daemon cannot read another organization's cached +-- Terraform module source. +SELECT EXISTS ( + SELECT 1 + FROM template_version_terraform_values tvtv + JOIN template_versions tv + ON tv.id = tvtv.template_version_id + WHERE tvtv.cached_module_files = @file_id::uuid + AND tv.organization_id = @organization_id::uuid +); diff --git a/coderd/provisionerdserver/download_file_test.go b/coderd/provisionerdserver/download_file_test.go new file mode 100644 index 00000000000..a109514dda1 --- /dev/null +++ b/coderd/provisionerdserver/download_file_test.go @@ -0,0 +1,187 @@ +package provisionerdserver_test + +import ( + "context" + crand "crypto/rand" + "testing" + + "github.com/google/uuid" + "github.com/stretchr/testify/require" + "storj.io/drpc" + + "github.com/coder/coder/v2/coderd/database" + "github.com/coder/coder/v2/coderd/database/dbgen" + "github.com/coder/coder/v2/coderd/database/dbtime" + "github.com/coder/coder/v2/coderd/externalauth" + proto "github.com/coder/coder/v2/provisionerd/proto" + sdkproto "github.com/coder/coder/v2/provisionersdk/proto" + "github.com/coder/coder/v2/testutil" +) + +// mockDownloadStream is an in-memory implementation of +// proto.DRPCProvisionerDaemon_DownloadFileStream that records every message +// sent by the handler so tests can assert on them. +type mockDownloadStream struct { + ctx context.Context + messages []*sdkproto.FileUpload +} + +func (m *mockDownloadStream) Send(f *sdkproto.FileUpload) error { + m.messages = append(m.messages, f) + return nil +} +func (m *mockDownloadStream) Context() context.Context { return m.ctx } +func (*mockDownloadStream) CloseSend() error { return nil } +func (*mockDownloadStream) Close() error { return nil } +func (*mockDownloadStream) MsgSend(drpc.Message, drpc.Encoding) error { return nil } +func (*mockDownloadStream) MsgRecv(drpc.Message, drpc.Encoding) error { return nil } + +// failure returns the error message streamed by the handler, if any. +func (m *mockDownloadStream) failure() (string, bool) { + for _, msg := range m.messages { + if e, ok := msg.Type.(*sdkproto.FileUpload_Error); ok { + return e.Error.Error, true + } + } + return "", false +} + +// downloadedBytes reassembles the file streamed by the handler from its +// DataUpload + ChunkPiece messages. +func (m *mockDownloadStream) downloadedBytes(t *testing.T) []byte { + t.Helper() + var upload *sdkproto.DataUpload + chunks := map[int32][]byte{} + for _, msg := range m.messages { + switch v := msg.Type.(type) { + case *sdkproto.FileUpload_DataUpload: + require.Nil(t, upload, "received more than one DataUpload") + upload = v.DataUpload + case *sdkproto.FileUpload_ChunkPiece: + chunks[v.ChunkPiece.PieceIndex] = v.ChunkPiece.Data + case *sdkproto.FileUpload_Error: + t.Fatalf("unexpected error message in stream: %s", v.Error.Error) + } + } + require.NotNil(t, upload, "no DataUpload message was streamed") + var data []byte + for i := int32(0); i < upload.Chunks; i++ { + data = append(data, chunks[i]...) + } + return data +} + +// insertModuleFile inserts a system-created (CreatedBy=uuid.Nil) tar file and +// links it as the cached module files of a template version in the given +// organization, returning the file. +func insertModuleFile(t *testing.T, db database.Store, orgID uuid.UUID, data []byte) database.File { + t.Helper() + ctx := testutil.Context(t, testutil.WaitShort) + + user := dbgen.User(t, db, database.User{}) + template := dbgen.Template(t, db, database.Template{ + OrganizationID: orgID, + CreatedBy: user.ID, + }) + jobID := uuid.New() + version := dbgen.TemplateVersion(t, db, database.TemplateVersion{ + OrganizationID: orgID, + CreatedBy: user.ID, + TemplateID: uuid.NullUUID{UUID: template.ID, Valid: true}, + JobID: jobID, + }) + // Insert the file directly rather than via dbgen.File: the helper treats a + // zero CreatedBy as "unset" and replaces it with a random UUID, but module + // files must be system-created (CreatedBy=uuid.Nil) to match the handler's + // metadata check. + file, err := db.InsertFile(ctx, database.InsertFileParams{ + ID: uuid.New(), + Hash: uuid.NewString(), + CreatedAt: dbtime.Now(), + CreatedBy: uuid.Nil, + Mimetype: "application/x-tar", + Data: data, + }) + require.NoError(t, err) + err = db.InsertTemplateVersionTerraformValuesByJobID(ctx, database.InsertTemplateVersionTerraformValuesByJobIDParams{ + JobID: version.JobID, + CachedPlan: []byte("{}"), + CachedModuleFiles: uuid.NullUUID{UUID: file.ID, Valid: true}, + UpdatedAt: dbtime.Now(), + }) + require.NoError(t, err) + return file +} + +// TestDownloadFileModuleFileTenantIsolation verifies that a provisioner daemon +// cannot download cached module archives belonging to other organizations +// (ANT-2026-22440), while still being able to download module files from its +// own organization. +func TestDownloadFileModuleFileTenantIsolation(t *testing.T) { + t.Parallel() + + t.Run("RejectsOtherOrgModuleFile", func(t *testing.T) { + t.Parallel() + + // The server is scoped to the default organization (org A). + server, db, _, daemon := setup(t, false, &overrides{ + externalAuthConfigs: []*externalauth.Config{{}}, + }) + ctx := testutil.Context(t, testutil.WaitMedium) + + // Create a module file belonging to a different organization (org B). + otherOrg := dbgen.Organization(t, db, database.Organization{}) + require.NotEqual(t, daemon.OrganizationID, otherOrg.ID) + + moduleData := make([]byte, sdkproto.ChunkSize*2) + _, err := crand.Read(moduleData) + require.NoError(t, err) + file := insertModuleFile(t, db, otherOrg.ID, moduleData) + + stream := &mockDownloadStream{ctx: ctx} + err = server.DownloadFile(&proto.FileRequest{ + FileId: file.ID.String(), + UploadType: sdkproto.DataUploadType_UPLOAD_TYPE_MODULE_FILES, + }, stream) + require.Error(t, err) + require.ErrorContains(t, err, "is not a modules file") + + // The handler must not have streamed any of the file's contents. + msg, ok := stream.failure() + require.True(t, ok, "expected an error message on the stream") + require.Contains(t, msg, "is not a modules file") + for _, m := range stream.messages { + switch m.Type.(type) { + case *sdkproto.FileUpload_DataUpload, *sdkproto.FileUpload_ChunkPiece: + t.Fatal("handler leaked file contents for another org's module file") + } + } + }) + + t.Run("AllowsSameOrgModuleFile", func(t *testing.T) { + t.Parallel() + + // The server is scoped to the default organization (org A). + server, db, _, daemon := setup(t, false, &overrides{ + externalAuthConfigs: []*externalauth.Config{{}}, + }) + ctx := testutil.Context(t, testutil.WaitMedium) + + moduleData := make([]byte, sdkproto.ChunkSize*2+512) + _, err := crand.Read(moduleData) + require.NoError(t, err) + file := insertModuleFile(t, db, daemon.OrganizationID, moduleData) + + stream := &mockDownloadStream{ctx: ctx} + err = server.DownloadFile(&proto.FileRequest{ + FileId: file.ID.String(), + UploadType: sdkproto.DataUploadType_UPLOAD_TYPE_MODULE_FILES, + }, stream) + require.NoError(t, err) + + if msg, ok := stream.failure(); ok { + t.Fatalf("unexpected error on stream: %s", msg) + } + require.Equal(t, moduleData, stream.downloadedBytes(t)) + }) +} diff --git a/coderd/provisionerdserver/provisionerdserver.go b/coderd/provisionerdserver/provisionerdserver.go index d233cb41dd9..a4ab4956429 100644 --- a/coderd/provisionerdserver/provisionerdserver.go +++ b/coderd/provisionerdserver/provisionerdserver.go @@ -1584,6 +1584,22 @@ func (s *server) DownloadFile(request *proto.FileRequest, stream proto.DRPCProvi if file.CreatedBy != uuid.Nil || file.Mimetype != tarMimeType { return fail(xerrors.Errorf("file %s is not a modules file", fid)) } + // Ensure the requested module file belongs to a template version in + // this provisioner daemon's organization. Without this, any + // authenticated provisioner could download cached module archives + // (Terraform source) belonging to other organizations (ANT-2026-22440). + ok, err := s.Database.HasTemplateVersionsUsingCachedModuleFileInOrg(ctx, database.HasTemplateVersionsUsingCachedModuleFileInOrgParams{ + FileID: fid, + OrganizationID: s.OrganizationID, + }) + if err != nil { + return fail(xerrors.Errorf("authorize module file: %w", err)) + } + if !ok { + // Use the same error as the metadata check above so the handler + // does not confirm the existence of files in other organizations. + return fail(xerrors.Errorf("file %s is not a modules file", fid)) + } default: return fail(xerrors.Errorf("unsupported file upload type: %s", request.UploadType)) } From a64f355f5bb2cfd35a620cea8abd1c0cdb3e1b71 Mon Sep 17 00:00:00 2001 From: Jon Ayers Date: Tue, 23 Jun 2026 22:32:32 +0000 Subject: [PATCH 2/5] fix(coderd/provisionerdserver): address review on module file download fix Add a debug-level log before the org-isolation rejection so operators have a breadcrumb without leaking file existence to the calling provisioner. Move the DownloadFile regression test into provisionerdserver_test.go and rename it to TestDownloadFile to match the existing per-RPC test naming. --- .../provisionerdserver/download_file_test.go | 187 ------------------ .../provisionerdserver/provisionerdserver.go | 4 + .../provisionerdserver_test.go | 168 ++++++++++++++++ 3 files changed, 172 insertions(+), 187 deletions(-) delete mode 100644 coderd/provisionerdserver/download_file_test.go diff --git a/coderd/provisionerdserver/download_file_test.go b/coderd/provisionerdserver/download_file_test.go deleted file mode 100644 index a109514dda1..00000000000 --- a/coderd/provisionerdserver/download_file_test.go +++ /dev/null @@ -1,187 +0,0 @@ -package provisionerdserver_test - -import ( - "context" - crand "crypto/rand" - "testing" - - "github.com/google/uuid" - "github.com/stretchr/testify/require" - "storj.io/drpc" - - "github.com/coder/coder/v2/coderd/database" - "github.com/coder/coder/v2/coderd/database/dbgen" - "github.com/coder/coder/v2/coderd/database/dbtime" - "github.com/coder/coder/v2/coderd/externalauth" - proto "github.com/coder/coder/v2/provisionerd/proto" - sdkproto "github.com/coder/coder/v2/provisionersdk/proto" - "github.com/coder/coder/v2/testutil" -) - -// mockDownloadStream is an in-memory implementation of -// proto.DRPCProvisionerDaemon_DownloadFileStream that records every message -// sent by the handler so tests can assert on them. -type mockDownloadStream struct { - ctx context.Context - messages []*sdkproto.FileUpload -} - -func (m *mockDownloadStream) Send(f *sdkproto.FileUpload) error { - m.messages = append(m.messages, f) - return nil -} -func (m *mockDownloadStream) Context() context.Context { return m.ctx } -func (*mockDownloadStream) CloseSend() error { return nil } -func (*mockDownloadStream) Close() error { return nil } -func (*mockDownloadStream) MsgSend(drpc.Message, drpc.Encoding) error { return nil } -func (*mockDownloadStream) MsgRecv(drpc.Message, drpc.Encoding) error { return nil } - -// failure returns the error message streamed by the handler, if any. -func (m *mockDownloadStream) failure() (string, bool) { - for _, msg := range m.messages { - if e, ok := msg.Type.(*sdkproto.FileUpload_Error); ok { - return e.Error.Error, true - } - } - return "", false -} - -// downloadedBytes reassembles the file streamed by the handler from its -// DataUpload + ChunkPiece messages. -func (m *mockDownloadStream) downloadedBytes(t *testing.T) []byte { - t.Helper() - var upload *sdkproto.DataUpload - chunks := map[int32][]byte{} - for _, msg := range m.messages { - switch v := msg.Type.(type) { - case *sdkproto.FileUpload_DataUpload: - require.Nil(t, upload, "received more than one DataUpload") - upload = v.DataUpload - case *sdkproto.FileUpload_ChunkPiece: - chunks[v.ChunkPiece.PieceIndex] = v.ChunkPiece.Data - case *sdkproto.FileUpload_Error: - t.Fatalf("unexpected error message in stream: %s", v.Error.Error) - } - } - require.NotNil(t, upload, "no DataUpload message was streamed") - var data []byte - for i := int32(0); i < upload.Chunks; i++ { - data = append(data, chunks[i]...) - } - return data -} - -// insertModuleFile inserts a system-created (CreatedBy=uuid.Nil) tar file and -// links it as the cached module files of a template version in the given -// organization, returning the file. -func insertModuleFile(t *testing.T, db database.Store, orgID uuid.UUID, data []byte) database.File { - t.Helper() - ctx := testutil.Context(t, testutil.WaitShort) - - user := dbgen.User(t, db, database.User{}) - template := dbgen.Template(t, db, database.Template{ - OrganizationID: orgID, - CreatedBy: user.ID, - }) - jobID := uuid.New() - version := dbgen.TemplateVersion(t, db, database.TemplateVersion{ - OrganizationID: orgID, - CreatedBy: user.ID, - TemplateID: uuid.NullUUID{UUID: template.ID, Valid: true}, - JobID: jobID, - }) - // Insert the file directly rather than via dbgen.File: the helper treats a - // zero CreatedBy as "unset" and replaces it with a random UUID, but module - // files must be system-created (CreatedBy=uuid.Nil) to match the handler's - // metadata check. - file, err := db.InsertFile(ctx, database.InsertFileParams{ - ID: uuid.New(), - Hash: uuid.NewString(), - CreatedAt: dbtime.Now(), - CreatedBy: uuid.Nil, - Mimetype: "application/x-tar", - Data: data, - }) - require.NoError(t, err) - err = db.InsertTemplateVersionTerraformValuesByJobID(ctx, database.InsertTemplateVersionTerraformValuesByJobIDParams{ - JobID: version.JobID, - CachedPlan: []byte("{}"), - CachedModuleFiles: uuid.NullUUID{UUID: file.ID, Valid: true}, - UpdatedAt: dbtime.Now(), - }) - require.NoError(t, err) - return file -} - -// TestDownloadFileModuleFileTenantIsolation verifies that a provisioner daemon -// cannot download cached module archives belonging to other organizations -// (ANT-2026-22440), while still being able to download module files from its -// own organization. -func TestDownloadFileModuleFileTenantIsolation(t *testing.T) { - t.Parallel() - - t.Run("RejectsOtherOrgModuleFile", func(t *testing.T) { - t.Parallel() - - // The server is scoped to the default organization (org A). - server, db, _, daemon := setup(t, false, &overrides{ - externalAuthConfigs: []*externalauth.Config{{}}, - }) - ctx := testutil.Context(t, testutil.WaitMedium) - - // Create a module file belonging to a different organization (org B). - otherOrg := dbgen.Organization(t, db, database.Organization{}) - require.NotEqual(t, daemon.OrganizationID, otherOrg.ID) - - moduleData := make([]byte, sdkproto.ChunkSize*2) - _, err := crand.Read(moduleData) - require.NoError(t, err) - file := insertModuleFile(t, db, otherOrg.ID, moduleData) - - stream := &mockDownloadStream{ctx: ctx} - err = server.DownloadFile(&proto.FileRequest{ - FileId: file.ID.String(), - UploadType: sdkproto.DataUploadType_UPLOAD_TYPE_MODULE_FILES, - }, stream) - require.Error(t, err) - require.ErrorContains(t, err, "is not a modules file") - - // The handler must not have streamed any of the file's contents. - msg, ok := stream.failure() - require.True(t, ok, "expected an error message on the stream") - require.Contains(t, msg, "is not a modules file") - for _, m := range stream.messages { - switch m.Type.(type) { - case *sdkproto.FileUpload_DataUpload, *sdkproto.FileUpload_ChunkPiece: - t.Fatal("handler leaked file contents for another org's module file") - } - } - }) - - t.Run("AllowsSameOrgModuleFile", func(t *testing.T) { - t.Parallel() - - // The server is scoped to the default organization (org A). - server, db, _, daemon := setup(t, false, &overrides{ - externalAuthConfigs: []*externalauth.Config{{}}, - }) - ctx := testutil.Context(t, testutil.WaitMedium) - - moduleData := make([]byte, sdkproto.ChunkSize*2+512) - _, err := crand.Read(moduleData) - require.NoError(t, err) - file := insertModuleFile(t, db, daemon.OrganizationID, moduleData) - - stream := &mockDownloadStream{ctx: ctx} - err = server.DownloadFile(&proto.FileRequest{ - FileId: file.ID.String(), - UploadType: sdkproto.DataUploadType_UPLOAD_TYPE_MODULE_FILES, - }, stream) - require.NoError(t, err) - - if msg, ok := stream.failure(); ok { - t.Fatalf("unexpected error on stream: %s", msg) - } - require.Equal(t, moduleData, stream.downloadedBytes(t)) - }) -} diff --git a/coderd/provisionerdserver/provisionerdserver.go b/coderd/provisionerdserver/provisionerdserver.go index a4ab4956429..09ed0523585 100644 --- a/coderd/provisionerdserver/provisionerdserver.go +++ b/coderd/provisionerdserver/provisionerdserver.go @@ -1596,6 +1596,10 @@ func (s *server) DownloadFile(request *proto.FileRequest, stream proto.DRPCProvi return fail(xerrors.Errorf("authorize module file: %w", err)) } if !ok { + s.Logger.Debug(ctx, "module file download rejected: file not referenced by any template version in daemon org", + slog.F("file_id", fid), + slog.F("organization_id", s.OrganizationID), + ) // Use the same error as the metadata check above so the handler // does not confirm the existence of files in other organizations. return fail(xerrors.Errorf("file %s is not a modules file", fid)) diff --git a/coderd/provisionerdserver/provisionerdserver_test.go b/coderd/provisionerdserver/provisionerdserver_test.go index 6c1bec15706..4ac065b9a2d 100644 --- a/coderd/provisionerdserver/provisionerdserver_test.go +++ b/coderd/provisionerdserver/provisionerdserver_test.go @@ -2,6 +2,7 @@ package provisionerdserver_test import ( "context" + crand "crypto/rand" "database/sql" "encoding/json" "io" @@ -5420,3 +5421,170 @@ func newFakeUsageInserter() (*coderdtest.UsageInserter, *atomic.Pointer[usage.In poitr.Store(&inserter) return fake, poitr } + +// mockDownloadStream is an in-memory implementation of +// proto.DRPCProvisionerDaemon_DownloadFileStream that records every message +// sent by the handler so tests can assert on them. +type mockDownloadStream struct { + ctx context.Context + messages []*sdkproto.FileUpload +} + +func (m *mockDownloadStream) Send(f *sdkproto.FileUpload) error { + m.messages = append(m.messages, f) + return nil +} +func (m *mockDownloadStream) Context() context.Context { return m.ctx } +func (*mockDownloadStream) CloseSend() error { return nil } +func (*mockDownloadStream) Close() error { return nil } +func (*mockDownloadStream) MsgSend(drpc.Message, drpc.Encoding) error { return nil } +func (*mockDownloadStream) MsgRecv(drpc.Message, drpc.Encoding) error { return nil } + +// failure returns the error message streamed by the handler, if any. +func (m *mockDownloadStream) failure() (string, bool) { + for _, msg := range m.messages { + if e, ok := msg.Type.(*sdkproto.FileUpload_Error); ok { + return e.Error.Error, true + } + } + return "", false +} + +// downloadedBytes reassembles the file streamed by the handler from its +// DataUpload + ChunkPiece messages. +func (m *mockDownloadStream) downloadedBytes(t *testing.T) []byte { + t.Helper() + var upload *sdkproto.DataUpload + chunks := map[int32][]byte{} + for _, msg := range m.messages { + switch v := msg.Type.(type) { + case *sdkproto.FileUpload_DataUpload: + require.Nil(t, upload, "received more than one DataUpload") + upload = v.DataUpload + case *sdkproto.FileUpload_ChunkPiece: + chunks[v.ChunkPiece.PieceIndex] = v.ChunkPiece.Data + case *sdkproto.FileUpload_Error: + t.Fatalf("unexpected error message in stream: %s", v.Error.Error) + } + } + require.NotNil(t, upload, "no DataUpload message was streamed") + var data []byte + for i := int32(0); i < upload.Chunks; i++ { + data = append(data, chunks[i]...) + } + return data +} + +// insertModuleFile inserts a system-created (CreatedBy=uuid.Nil) tar file and +// links it as the cached module files of a template version in the given +// organization, returning the file. +func insertModuleFile(t *testing.T, db database.Store, orgID uuid.UUID, data []byte) database.File { + t.Helper() + ctx := testutil.Context(t, testutil.WaitShort) + + user := dbgen.User(t, db, database.User{}) + template := dbgen.Template(t, db, database.Template{ + OrganizationID: orgID, + CreatedBy: user.ID, + }) + jobID := uuid.New() + version := dbgen.TemplateVersion(t, db, database.TemplateVersion{ + OrganizationID: orgID, + CreatedBy: user.ID, + TemplateID: uuid.NullUUID{UUID: template.ID, Valid: true}, + JobID: jobID, + }) + // Insert the file directly rather than via dbgen.File: the helper treats a + // zero CreatedBy as "unset" and replaces it with a random UUID, but module + // files must be system-created (CreatedBy=uuid.Nil) to match the handler's + // metadata check. + file, err := db.InsertFile(ctx, database.InsertFileParams{ + ID: uuid.New(), + Hash: uuid.NewString(), + CreatedAt: dbtime.Now(), + CreatedBy: uuid.Nil, + Mimetype: "application/x-tar", + Data: data, + }) + require.NoError(t, err) + err = db.InsertTemplateVersionTerraformValuesByJobID(ctx, database.InsertTemplateVersionTerraformValuesByJobIDParams{ + JobID: version.JobID, + CachedPlan: []byte("{}"), + CachedModuleFiles: uuid.NullUUID{UUID: file.ID, Valid: true}, + UpdatedAt: dbtime.Now(), + }) + require.NoError(t, err) + return file +} + +// TestDownloadFile verifies that a provisioner daemon cannot download cached +// module archives belonging to other organizations (ANT-2026-22440), while +// still being able to download module files from its own organization. +func TestDownloadFile(t *testing.T) { + t.Parallel() + + t.Run("RejectsOtherOrgModuleFile", func(t *testing.T) { + t.Parallel() + + // The server is scoped to the default organization (org A). + server, db, _, daemon := setup(t, false, &overrides{ + externalAuthConfigs: []*externalauth.Config{{}}, + }) + ctx := testutil.Context(t, testutil.WaitMedium) + + // Create a module file belonging to a different organization (org B). + otherOrg := dbgen.Organization(t, db, database.Organization{}) + require.NotEqual(t, daemon.OrganizationID, otherOrg.ID) + + moduleData := make([]byte, sdkproto.ChunkSize*2) + _, err := crand.Read(moduleData) + require.NoError(t, err) + file := insertModuleFile(t, db, otherOrg.ID, moduleData) + + stream := &mockDownloadStream{ctx: ctx} + err = server.DownloadFile(&proto.FileRequest{ + FileId: file.ID.String(), + UploadType: sdkproto.DataUploadType_UPLOAD_TYPE_MODULE_FILES, + }, stream) + require.Error(t, err) + require.ErrorContains(t, err, "is not a modules file") + + // The handler must not have streamed any of the file's contents. + msg, ok := stream.failure() + require.True(t, ok, "expected an error message on the stream") + require.Contains(t, msg, "is not a modules file") + for _, m := range stream.messages { + switch m.Type.(type) { + case *sdkproto.FileUpload_DataUpload, *sdkproto.FileUpload_ChunkPiece: + t.Fatal("handler leaked file contents for another org's module file") + } + } + }) + + t.Run("AllowsSameOrgModuleFile", func(t *testing.T) { + t.Parallel() + + // The server is scoped to the default organization (org A). + server, db, _, daemon := setup(t, false, &overrides{ + externalAuthConfigs: []*externalauth.Config{{}}, + }) + ctx := testutil.Context(t, testutil.WaitMedium) + + moduleData := make([]byte, sdkproto.ChunkSize*2+512) + _, err := crand.Read(moduleData) + require.NoError(t, err) + file := insertModuleFile(t, db, daemon.OrganizationID, moduleData) + + stream := &mockDownloadStream{ctx: ctx} + err = server.DownloadFile(&proto.FileRequest{ + FileId: file.ID.String(), + UploadType: sdkproto.DataUploadType_UPLOAD_TYPE_MODULE_FILES, + }, stream) + require.NoError(t, err) + + if msg, ok := stream.failure(); ok { + t.Fatalf("unexpected error on stream: %s", msg) + } + require.Equal(t, moduleData, stream.downloadedBytes(t)) + }) +} From 59438844c08a3c2b9dc8c9176b1ad873030ffbd8 Mon Sep 17 00:00:00 2001 From: Jon Ayers Date: Tue, 23 Jun 2026 22:49:31 +0000 Subject: [PATCH 3/5] fix(coderd/provisionerdserver): log module file download rejection at warn --- coderd/provisionerdserver/provisionerdserver.go | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/coderd/provisionerdserver/provisionerdserver.go b/coderd/provisionerdserver/provisionerdserver.go index 09ed0523585..9fb25812600 100644 --- a/coderd/provisionerdserver/provisionerdserver.go +++ b/coderd/provisionerdserver/provisionerdserver.go @@ -1596,7 +1596,7 @@ func (s *server) DownloadFile(request *proto.FileRequest, stream proto.DRPCProvi return fail(xerrors.Errorf("authorize module file: %w", err)) } if !ok { - s.Logger.Debug(ctx, "module file download rejected: file not referenced by any template version in daemon org", + s.Logger.Warn(ctx, "module file download rejected: file not referenced by any template version in daemon org", slog.F("file_id", fid), slog.F("organization_id", s.OrganizationID), ) From b4e6d8f928908bb1ce2aa4e31f400873fbd4c7f2 Mon Sep 17 00:00:00 2001 From: Jon Ayers Date: Tue, 23 Jun 2026 23:13:42 +0000 Subject: [PATCH 4/5] test(coderd/provisionerdserver): drop dead crypto/rand error check --- coderd/provisionerdserver/provisionerdserver_test.go | 12 ++++++------ 1 file changed, 6 insertions(+), 6 deletions(-) diff --git a/coderd/provisionerdserver/provisionerdserver_test.go b/coderd/provisionerdserver/provisionerdserver_test.go index 4ac065b9a2d..235f00aefeb 100644 --- a/coderd/provisionerdserver/provisionerdserver_test.go +++ b/coderd/provisionerdserver/provisionerdserver_test.go @@ -5537,12 +5537,12 @@ func TestDownloadFile(t *testing.T) { require.NotEqual(t, daemon.OrganizationID, otherOrg.ID) moduleData := make([]byte, sdkproto.ChunkSize*2) - _, err := crand.Read(moduleData) - require.NoError(t, err) + // crand.Read never returns an error as of Go 1.24. + _, _ = crand.Read(moduleData) file := insertModuleFile(t, db, otherOrg.ID, moduleData) stream := &mockDownloadStream{ctx: ctx} - err = server.DownloadFile(&proto.FileRequest{ + err := server.DownloadFile(&proto.FileRequest{ FileId: file.ID.String(), UploadType: sdkproto.DataUploadType_UPLOAD_TYPE_MODULE_FILES, }, stream) @@ -5571,12 +5571,12 @@ func TestDownloadFile(t *testing.T) { ctx := testutil.Context(t, testutil.WaitMedium) moduleData := make([]byte, sdkproto.ChunkSize*2+512) - _, err := crand.Read(moduleData) - require.NoError(t, err) + // crand.Read never returns an error as of Go 1.24. + _, _ = crand.Read(moduleData) file := insertModuleFile(t, db, daemon.OrganizationID, moduleData) stream := &mockDownloadStream{ctx: ctx} - err = server.DownloadFile(&proto.FileRequest{ + err := server.DownloadFile(&proto.FileRequest{ FileId: file.ID.String(), UploadType: sdkproto.DataUploadType_UPLOAD_TYPE_MODULE_FILES, }, stream) From 74fe9ca2484750d2ef9a48dff112eb1f3a3980d2 Mon Sep 17 00:00:00 2001 From: Jon Ayers Date: Wed, 24 Jun 2026 16:47:35 +0000 Subject: [PATCH 5/5] refactor(coderd): authorize module file download via ResourceFile.InOrg and use real DRPC stream in test Switch the HasTemplateVersionsUsingCachedModuleFileInOrg authorization from ResourceSystem to ResourceFile.InOrg, a least-privilege check that the caller can read files in the target organization. Replace the hand-rolled DownloadFile mock stream with a real in-memory DRPC client/server pipe and the production HandleReceivingDataUpload reassembly path. --- coderd/database/dbauthz/dbauthz.go | 8 +- coderd/database/dbauthz/dbauthz_test.go | 2 +- .../provisionerdserver_test.go | 120 +++++++----------- 3 files changed, 52 insertions(+), 78 deletions(-) diff --git a/coderd/database/dbauthz/dbauthz.go b/coderd/database/dbauthz/dbauthz.go index 4d4d6fb281c..3ad04ab92cf 100644 --- a/coderd/database/dbauthz/dbauthz.go +++ b/coderd/database/dbauthz/dbauthz.go @@ -5698,10 +5698,10 @@ func (q *querier) GetWorkspacesForWorkspaceMetrics(ctx context.Context) ([]datab } func (q *querier) HasTemplateVersionsUsingCachedModuleFileInOrg(ctx context.Context, arg database.HasTemplateVersionsUsingCachedModuleFileInOrgParams) (bool, error) { - // This is an internal cross-table guard used to authorize provisioner - // module-file downloads. Tenant isolation comes from the organization_id - // filter in the query itself; only system-level actors may call it. - if err := q.authorizeContext(ctx, policy.ActionRead, rbac.ResourceSystem); err != nil { + // This query authorizes provisioner module-file downloads. The caller + // must be able to read files in the target organization; the actual + // tenant isolation comes from the organization_id filter in the query. + if err := q.authorizeContext(ctx, policy.ActionRead, rbac.ResourceFile.InOrg(arg.OrganizationID)); err != nil { return false, err } return q.db.HasTemplateVersionsUsingCachedModuleFileInOrg(ctx, arg) diff --git a/coderd/database/dbauthz/dbauthz_test.go b/coderd/database/dbauthz/dbauthz_test.go index 50cf7524ace..8e77219654a 100644 --- a/coderd/database/dbauthz/dbauthz_test.go +++ b/coderd/database/dbauthz/dbauthz_test.go @@ -2677,7 +2677,7 @@ func (s *MethodTestSuite) TestTemplate() { s.Run("HasTemplateVersionsUsingCachedModuleFileInOrg", s.Mocked(func(dbm *dbmock.MockStore, _ *gofakeit.Faker, check *expects) { arg := database.HasTemplateVersionsUsingCachedModuleFileInOrgParams{FileID: uuid.New(), OrganizationID: uuid.New()} dbm.EXPECT().HasTemplateVersionsUsingCachedModuleFileInOrg(gomock.Any(), arg).Return(true, nil).AnyTimes() - check.Args(arg).Asserts(rbac.ResourceSystem, policy.ActionRead).Returns(true) + check.Args(arg).Asserts(rbac.ResourceFile.InOrg(arg.OrganizationID), policy.ActionRead).Returns(true) })) s.Run("GetTemplateVersionVariables", s.Mocked(func(dbm *dbmock.MockStore, faker *gofakeit.Faker, check *expects) { t1 := testutil.Fake(s.T(), faker, database.Template{}) diff --git a/coderd/provisionerdserver/provisionerdserver_test.go b/coderd/provisionerdserver/provisionerdserver_test.go index 235f00aefeb..8f0112decf7 100644 --- a/coderd/provisionerdserver/provisionerdserver_test.go +++ b/coderd/provisionerdserver/provisionerdserver_test.go @@ -26,6 +26,8 @@ import ( "golang.org/x/xerrors" "google.golang.org/protobuf/types/known/timestamppb" "storj.io/drpc" + "storj.io/drpc/drpcmux" + "storj.io/drpc/drpcserver" "cdr.dev/slog/v3" "cdr.dev/slog/v3/sloggers/slogtest" @@ -53,6 +55,7 @@ import ( "github.com/coder/coder/v2/coderd/usage/usagetypes" "github.com/coder/coder/v2/coderd/wspubsub" "github.com/coder/coder/v2/codersdk" + "github.com/coder/coder/v2/codersdk/drpcsdk" "github.com/coder/coder/v2/provisionerd/proto" "github.com/coder/coder/v2/provisionersdk" sdkproto "github.com/coder/coder/v2/provisionersdk/proto" @@ -5422,57 +5425,33 @@ func newFakeUsageInserter() (*coderdtest.UsageInserter, *atomic.Pointer[usage.In return fake, poitr } -// mockDownloadStream is an in-memory implementation of -// proto.DRPCProvisionerDaemon_DownloadFileStream that records every message -// sent by the handler so tests can assert on them. -type mockDownloadStream struct { - ctx context.Context - messages []*sdkproto.FileUpload -} - -func (m *mockDownloadStream) Send(f *sdkproto.FileUpload) error { - m.messages = append(m.messages, f) - return nil -} -func (m *mockDownloadStream) Context() context.Context { return m.ctx } -func (*mockDownloadStream) CloseSend() error { return nil } -func (*mockDownloadStream) Close() error { return nil } -func (*mockDownloadStream) MsgSend(drpc.Message, drpc.Encoding) error { return nil } -func (*mockDownloadStream) MsgRecv(drpc.Message, drpc.Encoding) error { return nil } - -// failure returns the error message streamed by the handler, if any. -func (m *mockDownloadStream) failure() (string, bool) { - for _, msg := range m.messages { - if e, ok := msg.Type.(*sdkproto.FileUpload_Error); ok { - return e.Error.Error, true - } - } - return "", false -} - -// downloadedBytes reassembles the file streamed by the handler from its -// DataUpload + ChunkPiece messages. -func (m *mockDownloadStream) downloadedBytes(t *testing.T) []byte { +// serveProvisionerDaemon serves the provisioner daemon server over an +// in-memory pipe and returns a connected client, mirroring how coderd serves +// in-memory provisioner daemons. This exercises the real DRPC streaming path +// instead of a hand-rolled mock stream. +func serveProvisionerDaemon(t *testing.T, srv proto.DRPCProvisionerDaemonServer) proto.DRPCProvisionerDaemonClient { t.Helper() - var upload *sdkproto.DataUpload - chunks := map[int32][]byte{} - for _, msg := range m.messages { - switch v := msg.Type.(type) { - case *sdkproto.FileUpload_DataUpload: - require.Nil(t, upload, "received more than one DataUpload") - upload = v.DataUpload - case *sdkproto.FileUpload_ChunkPiece: - chunks[v.ChunkPiece.PieceIndex] = v.ChunkPiece.Data - case *sdkproto.FileUpload_Error: - t.Fatalf("unexpected error message in stream: %s", v.Error.Error) - } - } - require.NotNil(t, upload, "no DataUpload message was streamed") - var data []byte - for i := int32(0); i < upload.Chunks; i++ { - data = append(data, chunks[i]...) - } - return data + clientPipe, serverPipe := drpcsdk.MemTransportPipe() + t.Cleanup(func() { + _ = clientPipe.Close() + _ = serverPipe.Close() + }) + mux := drpcmux.New() + require.NoError(t, proto.DRPCRegisterProvisionerDaemon(mux, srv)) + server := drpcserver.NewWithOptions(mux, drpcserver.Options{ + Manager: drpcsdk.DefaultDRPCOptions(nil), + }) + ctx, cancel := context.WithCancel(context.Background()) + closed := make(chan struct{}) + go func() { + defer close(closed) + _ = server.Serve(ctx, serverPipe) + }() + t.Cleanup(func() { + cancel() + <-closed + }) + return proto.NewDRPCProvisionerDaemonClient(clientPipe) } // insertModuleFile inserts a system-created (CreatedBy=uuid.Nil) tar file and @@ -5527,10 +5506,11 @@ func TestDownloadFile(t *testing.T) { t.Parallel() // The server is scoped to the default organization (org A). - server, db, _, daemon := setup(t, false, &overrides{ + srv, db, _, daemon := setup(t, false, &overrides{ externalAuthConfigs: []*externalauth.Config{{}}, }) ctx := testutil.Context(t, testutil.WaitMedium) + client := serveProvisionerDaemon(t, srv) // Create a module file belonging to a different organization (org B). otherOrg := dbgen.Organization(t, db, database.Organization{}) @@ -5541,50 +5521,44 @@ func TestDownloadFile(t *testing.T) { _, _ = crand.Read(moduleData) file := insertModuleFile(t, db, otherOrg.ID, moduleData) - stream := &mockDownloadStream{ctx: ctx} - err := server.DownloadFile(&proto.FileRequest{ + stream, err := client.DownloadFile(ctx, &proto.FileRequest{ FileId: file.ID.String(), UploadType: sdkproto.DataUploadType_UPLOAD_TYPE_MODULE_FILES, - }, stream) + }) + require.NoError(t, err) + + // The handler must reject the cross-org download with an error rather + // than streaming the file's contents. + _, err = provisionersdk.HandleReceivingDataUpload(stream) require.Error(t, err) require.ErrorContains(t, err, "is not a modules file") - - // The handler must not have streamed any of the file's contents. - msg, ok := stream.failure() - require.True(t, ok, "expected an error message on the stream") - require.Contains(t, msg, "is not a modules file") - for _, m := range stream.messages { - switch m.Type.(type) { - case *sdkproto.FileUpload_DataUpload, *sdkproto.FileUpload_ChunkPiece: - t.Fatal("handler leaked file contents for another org's module file") - } - } }) t.Run("AllowsSameOrgModuleFile", func(t *testing.T) { t.Parallel() // The server is scoped to the default organization (org A). - server, db, _, daemon := setup(t, false, &overrides{ + srv, db, _, daemon := setup(t, false, &overrides{ externalAuthConfigs: []*externalauth.Config{{}}, }) ctx := testutil.Context(t, testutil.WaitMedium) + client := serveProvisionerDaemon(t, srv) moduleData := make([]byte, sdkproto.ChunkSize*2+512) // crand.Read never returns an error as of Go 1.24. _, _ = crand.Read(moduleData) file := insertModuleFile(t, db, daemon.OrganizationID, moduleData) - stream := &mockDownloadStream{ctx: ctx} - err := server.DownloadFile(&proto.FileRequest{ + stream, err := client.DownloadFile(ctx, &proto.FileRequest{ FileId: file.ID.String(), UploadType: sdkproto.DataUploadType_UPLOAD_TYPE_MODULE_FILES, - }, stream) + }) require.NoError(t, err) - if msg, ok := stream.failure(); ok { - t.Fatalf("unexpected error on stream: %s", msg) - } - require.Equal(t, moduleData, stream.downloadedBytes(t)) + builder, err := provisionersdk.HandleReceivingDataUpload(stream) + require.NoError(t, err) + data, err := builder.Complete() + require.NoError(t, err) + require.Equal(t, moduleData, data) }) }