mirror of
https://github.com/coder/coder.git
synced 2026-09-24 15:04:27 +08:00
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:
+39
-18
@@ -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
|
||||
}
|
||||
Reference in New Issue
Block a user