diff --git a/commands/policy/eval.go b/commands/policy/eval.go index 364dd0f15..7c77b33e2 100644 --- a/commands/policy/eval.go +++ b/commands/policy/eval.go @@ -15,6 +15,7 @@ import ( "github.com/docker/buildx/builder" "github.com/docker/buildx/policy" "github.com/docker/buildx/util/confutil" + "github.com/docker/buildx/util/sourcemeta" "github.com/docker/cli/cli/command" "github.com/moby/buildkit/client/llb/sourceresolver" "github.com/moby/buildkit/frontend/dockerui" @@ -96,11 +97,8 @@ func runEval(ctx context.Context, dockerCli command.Cli, source string, opts eva OS: defaultPlatform.OS, Variant: defaultPlatform.Variant, } - openClient, release, err := gatewayClientFactory(c) - if err != nil { - return err - } - defer release() + metaResolver := sourcemeta.NewResolver(c) + defer metaResolver.Close() platform := &pb.Platform{ Architecture: p.Architecture, @@ -148,13 +146,8 @@ func runEval(ctx context.Context, dockerCli command.Cli, source string, opts eva if err := policy.AddUnknowns(req, toReload); err != nil { return err } - gwClient, err := openClient(ctx) - if err != nil { - return err - } - opt := sourceResolverOpt(req, &p) - resp, err := gwClient.ResolveSourceMetadata(ctx, src, opt) + resp, err := metaResolver.ResolveSourceMetadata(ctx, src, opt) if err != nil { return err } @@ -240,12 +233,8 @@ func runEval(ctx context.Context, dockerCli command.Cli, source string, opts eva return evalDecisionError(decision) } - gwClient, err := openClient(ctx) - if err != nil { - return err - } opt := sourceResolverOpt(next, &p) - resp, err := gwClient.ResolveSourceMetadata(ctx, src, opt) + resp, err := metaResolver.ResolveSourceMetadata(ctx, src, opt) if err != nil { return err } diff --git a/commands/policy/gateway_client.go b/commands/policy/gateway_client.go deleted file mode 100644 index 711d02855..000000000 --- a/commands/policy/gateway_client.go +++ /dev/null @@ -1,80 +0,0 @@ -package policy - -import ( - "context" - "errors" - "sync" - "sync/atomic" - - "github.com/moby/buildkit/client" - gwclient "github.com/moby/buildkit/frontend/gateway/client" -) - -type gatewayClientOpener func(context.Context) (gwclient.Client, error) - -func gatewayClientFactory(c *client.Client) (gatewayClientOpener, func() error, error) { - var ( - once sync.Once - releaseOnce sync.Once - started atomic.Bool - ready = make(chan gwclient.Client, 1) - done = make(chan error, 1) - openErr error - releaseErr error - gwClient gwclient.Client - cancel context.CancelCauseFunc - ) - - open := func(ctx context.Context) (gwclient.Client, error) { - once.Do(func() { - started.Store(true) - buildCtx, cancelFn := context.WithCancelCause(ctx) - cancel = cancelFn - - go func() { - _, err := c.Build(buildCtx, client.SolveOpt{Internal: true}, "buildx", func(ctx context.Context, c gwclient.Client) (*gwclient.Result, error) { - ready <- c - <-buildCtx.Done() - return nil, context.Cause(buildCtx) - }, nil) - done <- err - }() - - select { - case gwClient = <-ready: - case err := <-done: - if err == nil { - err = errors.New("gateway build finished without a client") - } - openErr = err - case <-ctx.Done(): - openErr = context.Cause(ctx) - cancelFn(openErr) - } - }) - - if openErr != nil { - return nil, openErr - } - return gwClient, nil - } - - release := func() error { - releaseOnce.Do(func() { - if !started.Load() { - return - } - if cancel != nil { - cancel(context.Canceled) - } - err := <-done - if err == nil || errors.Is(err, context.Canceled) || errors.Is(err, context.DeadlineExceeded) { - return - } - releaseErr = err - }) - return releaseErr - } - - return open, release, nil -} diff --git a/commands/policy/test.go b/commands/policy/test.go index 35eae0570..f1d964de6 100644 --- a/commands/policy/test.go +++ b/commands/policy/test.go @@ -14,6 +14,7 @@ import ( "github.com/docker/buildx/policy" "github.com/docker/buildx/util/cobrautil" "github.com/docker/buildx/util/confutil" + "github.com/docker/buildx/util/sourcemeta" "github.com/docker/cli/cli/command" gwpb "github.com/moby/buildkit/frontend/gateway/pb" "github.com/moby/buildkit/solver/pb" @@ -30,9 +31,9 @@ func testCmd(dockerCli command.Cli, rootOpts RootOptions) *cobra.Command { Args: cobra.ExactArgs(1), DisableFlagsInUseLine: true, RunE: func(cmd *cobra.Command, args []string) error { - resolver := newPolicyTestResolver(dockerCli, rootOpts.Builder) - opts.Resolver = resolver.Options() - defer resolver.Close() + optionsProvider := newPolicyTestOptionsProvider(dockerCli, rootOpts.Builder) + opts.Provider = optionsProvider.TestOptionsProvider() + defer optionsProvider.Close() return runTest(cmd.Context(), cmd.OutOrStdout(), args[0], opts) }, } @@ -114,63 +115,58 @@ func withInputPrefix(keys []string) []string { return out } -type policyTestResolver struct { +type policyTestOptionsProvider struct { dockerCli command.Cli builderName *string - once sync.Once - platform *ocispecs.Platform - openClient gatewayClientOpener - release func() error - err error + once sync.Once + platform *ocispecs.Platform + metaResolver *sourcemeta.Resolver + err error } -func newPolicyTestResolver(dockerCli command.Cli, builderName *string) *policyTestResolver { - return &policyTestResolver{ +func newPolicyTestOptionsProvider(dockerCli command.Cli, builderName *string) *policyTestOptionsProvider { + return &policyTestOptionsProvider{ dockerCli: dockerCli, builderName: builderName, } } -func (r *policyTestResolver) Options() *policy.TestResolver { - return &policy.TestResolver{ +func (r *policyTestOptionsProvider) TestOptionsProvider() *policy.TestOptionsProvider { + return &policy.TestOptionsProvider{ Resolve: r.Resolve, Platform: r.Platform, VerifierProvider: policy.SignatureVerifier(confutil.NewConfig(r.dockerCli)), } } -func (r *policyTestResolver) Close() error { - if r.release == nil { +func (r *policyTestOptionsProvider) Close() error { + if r.metaResolver == nil { return nil } - return r.release() + return r.metaResolver.Close() } -func (r *policyTestResolver) Platform(ctx context.Context) (*ocispecs.Platform, error) { +func (r *policyTestOptionsProvider) Platform(ctx context.Context) (*ocispecs.Platform, error) { if err := r.init(ctx); err != nil { return nil, err } return r.platform, nil } -func (r *policyTestResolver) Resolve(ctx context.Context, source *pb.SourceOp, req *gwpb.ResolveSourceMetaRequest) (*gwpb.ResolveSourceMetaResponse, error) { +func (r *policyTestOptionsProvider) Resolve(ctx context.Context, source *pb.SourceOp, req *gwpb.ResolveSourceMetaRequest) (*gwpb.ResolveSourceMetaResponse, error) { if err := r.init(ctx); err != nil { return nil, err } - gwClient, err := r.openClient(ctx) - if err != nil { - return nil, err - } opt := sourceResolverOpt(req, r.platform) - resp, err := gwClient.ResolveSourceMetadata(ctx, source, opt) + resp, err := r.metaResolver.ResolveSourceMetadata(ctx, source, opt) if err != nil { return nil, err } return buildSourceMetaResponse(resp), nil } -func (r *policyTestResolver) init(ctx context.Context) error { +func (r *policyTestOptionsProvider) init(ctx context.Context) error { r.once.Do(func() { bopts := []builder.Option{} if r.builderName != nil { @@ -208,13 +204,7 @@ func (r *policyTestResolver) init(ctx context.Context) error { OS: defaultPlatform.OS, Variant: defaultPlatform.Variant, } - openClient, release, err := gatewayClientFactory(c) - if err != nil { - r.err = err - return - } - r.openClient = openClient - r.release = release + r.metaResolver = sourcemeta.NewResolver(c) }) return r.err } diff --git a/policy/tester.go b/policy/tester.go index e93cc3b95..f6fd14e86 100644 --- a/policy/tester.go +++ b/policy/tester.go @@ -25,7 +25,7 @@ type TestOptions struct { Run string Filename string Root fs.StatFS - Resolver *TestResolver + Provider *TestOptionsProvider } type TestSummary struct { @@ -50,7 +50,7 @@ type testDef struct { PkgPath string } -type TestResolver struct { +type TestOptionsProvider struct { Resolve func(context.Context, *pb.SourceOp, *gwpb.ResolveSourceMetaRequest) (*gwpb.ResolveSourceMetaResponse, error) Platform func(context.Context) (*ocispecs.Platform, error) VerifierProvider PolicyVerifierProvider @@ -286,8 +286,8 @@ func runPolicyTest(ctx context.Context, policyModules map[string]*ast.Module, te return result, err } effectiveInput := input - if opts.Resolver != nil { - resolvedInput, ok, err := resolveTestInput(ctx, policyFiles, opts.Resolver, policyPackageModules, input, fsProvider) + if opts.Provider != nil { + resolvedInput, ok, err := resolveTestInput(ctx, policyFiles, opts.Provider, policyPackageModules, input, fsProvider) if err != nil { return result, err } @@ -397,7 +397,7 @@ func stateFromInput(input *Input) *state { return st } -func resolveTestInput(ctx context.Context, files []File, resolver *TestResolver, policyModules []*ast.Module, input *Input, fsProvider func() (fs.StatFS, func() error, error)) (*Input, bool, error) { +func resolveTestInput(ctx context.Context, files []File, resolver *TestOptionsProvider, policyModules []*ast.Module, input *Input, fsProvider func() (fs.StatFS, func() error, error)) (*Input, bool, error) { if resolver == nil { return nil, false, nil } diff --git a/util/sourcemeta/resolver.go b/util/sourcemeta/resolver.go new file mode 100644 index 000000000..51fd70719 --- /dev/null +++ b/util/sourcemeta/resolver.go @@ -0,0 +1,126 @@ +package sourcemeta + +import ( + "context" + "errors" + "sync" + "sync/atomic" + + "github.com/moby/buildkit/client" + "github.com/moby/buildkit/client/llb/sourceresolver" + gwclient "github.com/moby/buildkit/frontend/gateway/client" + "github.com/moby/buildkit/solver/pb" +) + +var _ sourceresolver.MetaResolver = &Resolver{} + +type Resolver struct { + startOnce sync.Once + closeOnce sync.Once + started atomic.Bool + mu sync.Mutex + + ready chan sourceresolver.MetaResolver + done chan struct{} + openErr error + doneErr error + cancel context.CancelCauseFunc + + metaResolver sourceresolver.MetaResolver + run func(context.Context, chan<- sourceresolver.MetaResolver) error + closeErr error +} + +func NewResolver(c *client.Client) *Resolver { + return newWithRun(func(ctx context.Context, ready chan<- sourceresolver.MetaResolver) error { + _, err := c.Build(ctx, client.SolveOpt{Internal: true}, "buildx", func(ctx context.Context, gw gwclient.Client) (*gwclient.Result, error) { + ready <- gw + <-ctx.Done() + return nil, context.Cause(ctx) + }, nil) + return err + }) +} + +func newWithRun(run func(context.Context, chan<- sourceresolver.MetaResolver) error) *Resolver { + return &Resolver{ + ready: make(chan sourceresolver.MetaResolver, 1), + done: make(chan struct{}), + run: run, + } +} + +func (r *Resolver) ResolveSourceMetadata(ctx context.Context, op *pb.SourceOp, opt sourceresolver.Opt) (*sourceresolver.MetaResponse, error) { + mr, err := r.open(ctx) + if err != nil { + return nil, err + } + return mr.ResolveSourceMetadata(ctx, op, opt) +} + +func (r *Resolver) Close() error { + r.closeOnce.Do(func() { + if !r.started.Load() { + return + } + if r.cancel != nil { + r.cancel(context.Canceled) + } + <-r.done + err := r.doneErr + if err == nil || errors.Is(err, context.Canceled) || errors.Is(err, context.DeadlineExceeded) { + return + } + r.closeErr = err + }) + return r.closeErr +} + +func (r *Resolver) open(ctx context.Context) (sourceresolver.MetaResolver, error) { + r.startOnce.Do(func() { + r.started.Store(true) + buildCtx, cancel := context.WithCancelCause(context.Background()) + r.cancel = cancel + + go func() { + r.doneErr = r.run(buildCtx, r.ready) + close(r.done) + }() + }) + + for { + r.mu.Lock() + if r.metaResolver != nil { + mr := r.metaResolver + r.mu.Unlock() + return mr, nil + } + if r.openErr != nil { + err := r.openErr + r.mu.Unlock() + return nil, err + } + r.mu.Unlock() + + select { + case mr := <-r.ready: + r.mu.Lock() + if r.metaResolver == nil { + r.metaResolver = mr + } + r.mu.Unlock() + case <-r.done: + r.mu.Lock() + if r.metaResolver == nil && r.openErr == nil { + err := r.doneErr + if err == nil { + err = errors.New("gateway build finished without a source metadata resolver") + } + r.openErr = err + } + r.mu.Unlock() + case <-ctx.Done(): + return nil, context.Cause(ctx) + } + } +} diff --git a/util/sourcemeta/resolver_test.go b/util/sourcemeta/resolver_test.go new file mode 100644 index 000000000..a6904bf25 --- /dev/null +++ b/util/sourcemeta/resolver_test.go @@ -0,0 +1,193 @@ +package sourcemeta + +import ( + "context" + "errors" + "sync" + "sync/atomic" + "testing" + "time" + + "github.com/moby/buildkit/client/llb/sourceresolver" + "github.com/moby/buildkit/solver/pb" + "github.com/stretchr/testify/require" +) + +type fakeMetaResolver struct { + calls atomic.Int32 + resp *sourceresolver.MetaResponse + err error +} + +func (f *fakeMetaResolver) ResolveSourceMetadata(ctx context.Context, op *pb.SourceOp, opt sourceresolver.Opt) (*sourceresolver.MetaResponse, error) { + f.calls.Add(1) + return f.resp, f.err +} + +func TestResolverCloseNoopBeforeResolve(t *testing.T) { + t.Parallel() + + var called atomic.Int32 + r := newWithRun(func(ctx context.Context, ready chan<- sourceresolver.MetaResolver) error { + called.Add(1) + return nil + }) + + require.NoError(t, r.Close()) + require.EqualValues(t, 0, called.Load()) +} + +func TestResolverResolveOpensOnce(t *testing.T) { + t.Parallel() + + var runs atomic.Int32 + mr := &fakeMetaResolver{resp: &sourceresolver.MetaResponse{}} + r := newWithRun(func(ctx context.Context, ready chan<- sourceresolver.MetaResolver) error { + runs.Add(1) + ready <- mr + <-ctx.Done() + return context.Cause(ctx) + }) + + op := &pb.SourceOp{} + _, err := r.ResolveSourceMetadata(t.Context(), op, sourceresolver.Opt{}) + require.NoError(t, err) + _, err = r.ResolveSourceMetadata(t.Context(), op, sourceresolver.Opt{}) + require.NoError(t, err) + + require.EqualValues(t, 1, runs.Load()) + require.EqualValues(t, 2, mr.calls.Load()) + require.NoError(t, r.Close()) +} + +func TestResolverCloseAfterOpenCancelsBuild(t *testing.T) { + t.Parallel() + + var canceled atomic.Bool + r := newWithRun(func(ctx context.Context, ready chan<- sourceresolver.MetaResolver) error { + ready <- &fakeMetaResolver{resp: &sourceresolver.MetaResponse{}} + <-ctx.Done() + canceled.Store(true) + return context.Cause(ctx) + }) + + _, err := r.ResolveSourceMetadata(t.Context(), &pb.SourceOp{}, sourceresolver.Opt{}) + require.NoError(t, err) + require.NoError(t, r.Close()) + require.True(t, canceled.Load()) +} + +func TestResolverOpenFailureIsSticky(t *testing.T) { + t.Parallel() + + expected := errors.New("boom") + var runs atomic.Int32 + r := newWithRun(func(ctx context.Context, ready chan<- sourceresolver.MetaResolver) error { + runs.Add(1) + return expected + }) + + _, err := r.ResolveSourceMetadata(t.Context(), &pb.SourceOp{}, sourceresolver.Opt{}) + require.ErrorIs(t, err, expected) + _, err = r.ResolveSourceMetadata(t.Context(), &pb.SourceOp{}, sourceresolver.Opt{}) + require.ErrorIs(t, err, expected) + require.EqualValues(t, 1, runs.Load()) + require.ErrorIs(t, r.Close(), expected) +} + +func TestResolverCloseIgnoresTerminalContextErrors(t *testing.T) { + t.Parallel() + + testCases := []struct { + name string + err error + }{ + {name: "canceled", err: context.Canceled}, + {name: "deadline", err: context.DeadlineExceeded}, + } + + for _, tc := range testCases { + t.Run(tc.name, func(t *testing.T) { + r := newWithRun(func(ctx context.Context, ready chan<- sourceresolver.MetaResolver) error { + return tc.err + }) + _, err := r.ResolveSourceMetadata(t.Context(), &pb.SourceOp{}, sourceresolver.Opt{}) + require.ErrorIs(t, err, tc.err) + require.NoError(t, r.Close()) + }) + } +} + +func TestResolverConcurrentResolveUsesSingleOpen(t *testing.T) { + t.Parallel() + + var runs atomic.Int32 + mr := &fakeMetaResolver{resp: &sourceresolver.MetaResponse{}} + r := newWithRun(func(ctx context.Context, ready chan<- sourceresolver.MetaResolver) error { + runs.Add(1) + ready <- mr + <-ctx.Done() + return context.Cause(ctx) + }) + + const n = 16 + errCh := make(chan error, n) + var wg sync.WaitGroup + wg.Add(n) + for range n { + go func() { + defer wg.Done() + _, err := r.ResolveSourceMetadata(t.Context(), &pb.SourceOp{}, sourceresolver.Opt{}) + errCh <- err + }() + } + wg.Wait() + close(errCh) + + for err := range errCh { + require.NoError(t, err) + } + require.EqualValues(t, 1, runs.Load()) + require.EqualValues(t, n, mr.calls.Load()) + + done := make(chan struct{}) + closeErr := make(chan error, 1) + go func() { + defer close(done) + closeErr <- r.Close() + }() + select { + case <-done: + require.NoError(t, <-closeErr) + case <-time.After(2 * time.Second): + t.Fatal("close timed out") + } +} + +func TestResolverFirstCanceledContextDoesNotPoisonFutureCalls(t *testing.T) { + t.Parallel() + + mr := &fakeMetaResolver{resp: &sourceresolver.MetaResponse{}} + started := make(chan struct{}) + release := make(chan struct{}) + + r := newWithRun(func(ctx context.Context, ready chan<- sourceresolver.MetaResolver) error { + close(started) + <-release + ready <- mr + <-ctx.Done() + return context.Cause(ctx) + }) + + canceledCtx, cancel := context.WithCancelCause(t.Context()) + cancel(context.Canceled) + _, err := r.ResolveSourceMetadata(canceledCtx, &pb.SourceOp{}, sourceresolver.Opt{}) + require.ErrorIs(t, err, context.Canceled) + + <-started + close(release) + + _, err = r.ResolveSourceMetadata(t.Context(), &pb.SourceOp{}, sourceresolver.Opt{}) + require.NoError(t, err) + require.NoError(t, r.Close()) +}