diff --git a/build/policy_loader.go b/build/policy_loader.go index f12835ec4..61c680181 100644 --- a/build/policy_loader.go +++ b/build/policy_loader.go @@ -228,28 +228,48 @@ func normalizeLocalPolicyPath(name, contextDir string) string { type memoizedPolicyFS struct { init func() (fs.StatFS, func() error, error) - once sync.Once + mu sync.Mutex + loaded bool + refs int fs fs.StatFS closeFn func() error err error } func (m *memoizedPolicyFS) get() (fs.StatFS, error) { - m.once.Do(func() { - if m.init == nil { - return + m.mu.Lock() + defer m.mu.Unlock() + if !m.loaded { + m.loaded = true + if m.init != nil { + m.fs, m.closeFn, m.err = m.init() } - m.fs, m.closeFn, m.err = m.init() - }) + } if m.err != nil { return nil, m.err } + m.refs++ return m.fs, nil } func (m *memoizedPolicyFS) close() error { - if m.closeFn != nil { - return m.closeFn() + m.mu.Lock() + if m.refs > 0 { + m.refs-- + } + if m.refs > 0 { + m.mu.Unlock() + return nil + } + closeFn := m.closeFn + m.fs = nil + m.closeFn = nil + m.err = nil + m.loaded = false + m.mu.Unlock() + + if closeFn != nil { + return closeFn() } return nil } diff --git a/build/policy_loader_test.go b/build/policy_loader_test.go new file mode 100644 index 000000000..584f31504 --- /dev/null +++ b/build/policy_loader_test.go @@ -0,0 +1,85 @@ +package build + +import ( + "io/fs" + "testing" + "testing/fstest" + + "github.com/stretchr/testify/require" +) + +func TestMemoizedPolicyFSRefCountedClose(t *testing.T) { + var initCalls int + var closeCalls int + + m := &memoizedPolicyFS{ + init: func() (fs.StatFS, func() error, error) { + initCalls++ + root := fstest.MapFS{ + "policy.rego": &fstest.MapFile{Data: []byte("package docker\n")}, + } + return root, func() error { + closeCalls++ + return nil + }, nil + }, + } + + first, err := m.get() + require.NoError(t, err) + require.NotNil(t, first) + require.Equal(t, 1, initCalls) + + second, err := m.get() + require.NoError(t, err) + require.NotNil(t, second) + require.Equal(t, 1, initCalls) + + require.NoError(t, m.close()) + require.Equal(t, 0, closeCalls) + + third, err := m.get() + require.NoError(t, err) + require.NotNil(t, third) + require.Equal(t, 1, initCalls) + + require.NoError(t, m.close()) + require.Equal(t, 0, closeCalls) + + require.NoError(t, m.close()) + require.Equal(t, 1, closeCalls) +} + +func TestMemoizedPolicyFSReinitializesAfterAllRefsClosed(t *testing.T) { + var initCalls int + var closeCalls int + + m := &memoizedPolicyFS{ + init: func() (fs.StatFS, func() error, error) { + initCalls++ + root := fstest.MapFS{ + "policy.rego": &fstest.MapFile{Data: []byte("package docker\n")}, + } + return root, func() error { + closeCalls++ + return nil + }, nil + }, + } + + first, err := m.get() + require.NoError(t, err) + require.NotNil(t, first) + require.Equal(t, 1, initCalls) + + require.NoError(t, m.close()) + require.Equal(t, 1, closeCalls) + + second, err := m.get() + require.NoError(t, err) + require.NotNil(t, second) + require.Equal(t, 2, initCalls) + + require.NoError(t, m.close()) + require.Equal(t, 2, closeCalls) +} diff --git a/policy/funcs.go b/policy/funcs.go index 4af3bfd82..4d5d46525 100644 --- a/policy/funcs.go +++ b/policy/funcs.go @@ -486,13 +486,15 @@ func (p *Policy) readFile(path string, limit int64) ([]byte, error) { if p.opt.FS == nil { return nil, errors.Errorf("no policy FS defined for reading context files") } - fs, cf, err := p.opt.FS() + root, closeFS, err := p.opt.FS() if err != nil { return nil, errors.Wrapf(err, "failed to get policy FS for reading context files") } - defer cf() + if closeFS != nil { + defer closeFS() + } - f, err := fs.Open(path) + f, err := root.Open(path) if err != nil { return nil, errors.Wrapf(err, "failed opening file %q", path) }