chore: ensure proper rbac permissions on 'Acquire' file in the cache (#18348)

The file cache was caching the `Unauthorized` errors if a user without
the right perms opened the file first. So all future opens would fail.

Now the cache always opens with a subject that can read files. And authz
is checked on the Acquire per user.
This commit is contained in:
Steven Masley
2025-06-16 13:40:45 +00:00
committed by GitHub
parent d83706bd5b
commit 1d1070d051
16 changed files with 218 additions and 51 deletions
+39 -18
View File
@@ -13,33 +13,41 @@ import (
archivefs "github.com/coder/coder/v2/archive/fs"
"github.com/coder/coder/v2/coderd/database"
"github.com/coder/coder/v2/coderd/database/dbauthz"
"github.com/coder/coder/v2/coderd/rbac"
"github.com/coder/coder/v2/coderd/rbac/policy"
"github.com/coder/coder/v2/coderd/util/lazy"
)
// NewFromStore returns a file cache that will fetch files from the provided
// database.
func NewFromStore(store database.Store, registerer prometheus.Registerer) *Cache {
fetch := func(ctx context.Context, fileID uuid.UUID) (cacheEntryValue, error) {
file, err := store.GetFileByID(ctx, fileID)
func NewFromStore(store database.Store, registerer prometheus.Registerer, authz rbac.Authorizer) *Cache {
fetch := func(ctx context.Context, fileID uuid.UUID) (CacheEntryValue, error) {
// Make sure the read does not fail due to authorization issues.
// Authz is checked on the Acquire call, so this is safe.
//nolint:gocritic
file, err := store.GetFileByID(dbauthz.AsFileReader(ctx), fileID)
if err != nil {
return cacheEntryValue{}, xerrors.Errorf("failed to read file from database: %w", err)
return CacheEntryValue{}, xerrors.Errorf("failed to read file from database: %w", err)
}
content := bytes.NewBuffer(file.Data)
return cacheEntryValue{
FS: archivefs.FromTarReader(content),
size: int64(content.Len()),
return CacheEntryValue{
Object: file.RBACObject(),
FS: archivefs.FromTarReader(content),
Size: int64(content.Len()),
}, nil
}
return New(fetch, registerer)
return New(fetch, registerer, authz)
}
func New(fetch fetcher, registerer prometheus.Registerer) *Cache {
func New(fetch fetcher, registerer prometheus.Registerer, authz rbac.Authorizer) *Cache {
return (&Cache{
lock: sync.Mutex{},
data: make(map[uuid.UUID]*cacheEntry),
fetcher: fetch,
authz: authz,
}).registerMetrics(registerer)
}
@@ -101,6 +109,7 @@ type Cache struct {
lock sync.Mutex
data map[uuid.UUID]*cacheEntry
fetcher
authz rbac.Authorizer
// metrics
cacheMetrics
@@ -117,18 +126,19 @@ type cacheMetrics struct {
totalCacheSize prometheus.Counter
}
type cacheEntryValue struct {
type CacheEntryValue struct {
fs.FS
size int64
Object rbac.Object
Size int64
}
type cacheEntry struct {
// refCount must only be accessed while the Cache lock is held.
refCount int
value *lazy.ValueWithError[cacheEntryValue]
value *lazy.ValueWithError[CacheEntryValue]
}
type fetcher func(context.Context, uuid.UUID) (cacheEntryValue, error)
type fetcher func(context.Context, uuid.UUID) (CacheEntryValue, error)
// Acquire will load the fs.FS for the given file. It guarantees that parallel
// calls for the same fileID will only result in one fetch, and that parallel
@@ -146,22 +156,33 @@ func (c *Cache) Acquire(ctx context.Context, fileID uuid.UUID) (fs.FS, error) {
c.Release(fileID)
return nil, err
}
subject, ok := dbauthz.ActorFromContext(ctx)
if !ok {
return nil, dbauthz.ErrNoActor
}
// Always check the caller can actually read the file.
if err := c.authz.Authorize(ctx, subject, policy.ActionRead, it.Object); err != nil {
c.Release(fileID)
return nil, err
}
return it.FS, err
}
func (c *Cache) prepare(ctx context.Context, fileID uuid.UUID) *lazy.ValueWithError[cacheEntryValue] {
func (c *Cache) prepare(ctx context.Context, fileID uuid.UUID) *lazy.ValueWithError[CacheEntryValue] {
c.lock.Lock()
defer c.lock.Unlock()
entry, ok := c.data[fileID]
if !ok {
value := lazy.NewWithError(func() (cacheEntryValue, error) {
value := lazy.NewWithError(func() (CacheEntryValue, error) {
val, err := c.fetcher(ctx, fileID)
// Always add to the cache size the bytes of the file loaded.
if err == nil {
c.currentCacheSize.Add(float64(val.size))
c.totalCacheSize.Add(float64(val.size))
c.currentCacheSize.Add(float64(val.Size))
c.totalCacheSize.Add(float64(val.Size))
}
return val, err
@@ -206,7 +227,7 @@ func (c *Cache) Release(fileID uuid.UUID) {
ev, err := entry.value.Load()
if err == nil {
c.currentCacheSize.Add(-1 * float64(ev.size))
c.currentCacheSize.Add(-1 * float64(ev.Size))
}
delete(c.data, fileID)
@@ -1,4 +1,4 @@
package files
package files_test
import (
"context"
@@ -12,28 +12,114 @@ import (
"github.com/stretchr/testify/require"
"golang.org/x/sync/errgroup"
"cdr.dev/slog/sloggers/slogtest"
"github.com/coder/coder/v2/coderd/coderdtest"
"github.com/coder/coder/v2/coderd/coderdtest/promhelp"
"github.com/coder/coder/v2/coderd/database"
"github.com/coder/coder/v2/coderd/database/dbauthz"
"github.com/coder/coder/v2/coderd/database/dbgen"
"github.com/coder/coder/v2/coderd/database/dbtestutil"
"github.com/coder/coder/v2/coderd/files"
"github.com/coder/coder/v2/coderd/rbac"
"github.com/coder/coder/v2/coderd/rbac/policy"
"github.com/coder/coder/v2/testutil"
)
// nolint:paralleltest,tparallel // Serially testing is easier
func TestCacheRBAC(t *testing.T) {
t.Parallel()
db, cache, rec := cacheAuthzSetup(t)
ctx := testutil.Context(t, testutil.WaitMedium)
file := dbgen.File(t, db, database.File{})
nobodyID := uuid.New()
nobody := dbauthz.As(ctx, rbac.Subject{
ID: nobodyID.String(),
Roles: rbac.Roles{},
Scope: rbac.ScopeAll,
})
userID := uuid.New()
userReader := dbauthz.As(ctx, rbac.Subject{
ID: userID.String(),
Roles: rbac.Roles{
must(rbac.RoleByName(rbac.RoleTemplateAdmin())),
},
Scope: rbac.ScopeAll,
})
//nolint:gocritic // Unit testing
cacheReader := dbauthz.AsFileReader(ctx)
t.Run("NoRolesOpen", func(t *testing.T) {
// Ensure start is clean
require.Equal(t, 0, cache.Count())
rec.Reset()
_, err := cache.Acquire(nobody, file.ID)
require.Error(t, err)
require.True(t, rbac.IsUnauthorizedError(err))
// Ensure that the cache is empty
require.Equal(t, 0, cache.Count())
// Check the assertions
rec.AssertActorID(t, nobodyID.String(), rec.Pair(policy.ActionRead, file))
rec.AssertActorID(t, rbac.SubjectTypeFileReaderID, rec.Pair(policy.ActionRead, file))
})
t.Run("CacheHasFile", func(t *testing.T) {
rec.Reset()
require.Equal(t, 0, cache.Count())
// Read the file with a file reader to put it into the cache.
_, err := cache.Acquire(cacheReader, file.ID)
require.NoError(t, err)
require.Equal(t, 1, cache.Count())
// "nobody" should not be able to read the file.
_, err = cache.Acquire(nobody, file.ID)
require.Error(t, err)
require.True(t, rbac.IsUnauthorizedError(err))
require.Equal(t, 1, cache.Count())
// UserReader can
_, err = cache.Acquire(userReader, file.ID)
require.NoError(t, err)
require.Equal(t, 1, cache.Count())
cache.Release(file.ID)
cache.Release(file.ID)
require.Equal(t, 0, cache.Count())
rec.AssertActorID(t, nobodyID.String(), rec.Pair(policy.ActionRead, file))
rec.AssertActorID(t, rbac.SubjectTypeFileReaderID, rec.Pair(policy.ActionRead, file))
rec.AssertActorID(t, userID.String(), rec.Pair(policy.ActionRead, file))
})
}
func cachePromMetricName(metric string) string {
return "coderd_file_cache_" + metric
}
func TestConcurrency(t *testing.T) {
t.Parallel()
//nolint:gocritic // Unit testing
ctx := dbauthz.AsFileReader(t.Context())
const fileSize = 10
emptyFS := afero.NewIOFS(afero.NewReadOnlyFs(afero.NewMemMapFs()))
var fetches atomic.Int64
reg := prometheus.NewRegistry()
c := New(func(_ context.Context, _ uuid.UUID) (cacheEntryValue, error) {
c := files.New(func(_ context.Context, _ uuid.UUID) (files.CacheEntryValue, error) {
fetches.Add(1)
// Wait long enough before returning to make sure that all of the goroutines
// will be waiting in line, ensuring that no one duplicated a fetch.
time.Sleep(testutil.IntervalMedium)
return cacheEntryValue{FS: emptyFS, size: fileSize}, nil
}, reg)
return files.CacheEntryValue{FS: emptyFS, Size: fileSize}, nil
}, reg, &coderdtest.FakeAuthorizer{})
batches := 1000
groups := make([]*errgroup.Group, 0, batches)
@@ -51,7 +137,7 @@ func TestConcurrency(t *testing.T) {
g.Go(func() error {
// We don't bother to Release these references because the Cache will be
// released at the end of the test anyway.
_, err := c.Acquire(t.Context(), id)
_, err := c.Acquire(ctx, id)
return err
})
}
@@ -74,16 +160,18 @@ func TestConcurrency(t *testing.T) {
func TestRelease(t *testing.T) {
t.Parallel()
//nolint:gocritic // Unit testing
ctx := dbauthz.AsFileReader(t.Context())
const fileSize = 10
emptyFS := afero.NewIOFS(afero.NewReadOnlyFs(afero.NewMemMapFs()))
reg := prometheus.NewRegistry()
c := New(func(_ context.Context, _ uuid.UUID) (cacheEntryValue, error) {
return cacheEntryValue{
c := files.New(func(_ context.Context, _ uuid.UUID) (files.CacheEntryValue, error) {
return files.CacheEntryValue{
FS: emptyFS,
size: fileSize,
Size: fileSize,
}, nil
}, reg)
}, reg, &coderdtest.FakeAuthorizer{})
batches := 100
ids := make([]uuid.UUID, 0, batches)
@@ -95,7 +183,7 @@ func TestRelease(t *testing.T) {
batchSize := 10
for openedIdx, id := range ids {
for batchIdx := range batchSize {
it, err := c.Acquire(t.Context(), id)
it, err := c.Acquire(ctx, id)
require.NoError(t, err)
require.Equal(t, emptyFS, it)
@@ -112,7 +200,7 @@ func TestRelease(t *testing.T) {
}
// Make sure cache is fully loaded
require.Equal(t, len(c.data), batches)
require.Equal(t, c.Count(), batches)
// Now release all of the references
for closedIdx, id := range ids {
@@ -136,7 +224,7 @@ func TestRelease(t *testing.T) {
}
// ...and make sure that the cache has emptied itself.
require.Equal(t, len(c.data), 0)
require.Equal(t, c.Count(), 0)
// Verify all the counts & metrics are correct.
// All existing files are closed
@@ -150,3 +238,29 @@ func TestRelease(t *testing.T) {
require.Equal(t, batches, promhelp.CounterValue(t, reg, cachePromMetricName("open_files_total"), nil))
require.Equal(t, batches*batchSize, promhelp.CounterValue(t, reg, cachePromMetricName("open_file_refs_total"), nil))
}
func cacheAuthzSetup(t *testing.T) (database.Store, *files.Cache, *coderdtest.RecordingAuthorizer) {
t.Helper()
logger := slogtest.Make(t, &slogtest.Options{})
reg := prometheus.NewRegistry()
db, _ := dbtestutil.NewDB(t)
authz := rbac.NewAuthorizer(reg)
rec := &coderdtest.RecordingAuthorizer{
Called: nil,
Wrapped: authz,
}
// Dbauthz wrap the db
db = dbauthz.New(db, rec, logger, coderdtest.AccessControlStorePointer())
c := files.NewFromStore(db, reg, rec)
return db, c, rec
}
func must[T any](t T, err error) T {
if err != nil {
panic(err)
}
return t
}