diff --git a/core/bazel/command.go b/core/bazel/command.go index 839fd388..abf09ad7 100644 --- a/core/bazel/command.go +++ b/core/bazel/command.go @@ -1,6 +1,8 @@ package bazel -import "io" +import ( + "io" +) type commander interface { StdoutPipe() (io.ReadCloser, error) diff --git a/core/bazel/commandermock/commandermock.go b/core/bazel/commandermock/commandermock.go index 13aba3c6..f11f4c66 100644 --- a/core/bazel/commandermock/commandermock.go +++ b/core/bazel/commandermock/commandermock.go @@ -1,5 +1,10 @@ // Code generated by MockGen. DO NOT EDIT. // Source: core/bazel/command.go +// +// Generated by this command: +// +// mockgen -package=commandermock -destination=core/bazel/commandermock/commandermock.go -source=core/bazel/command.go commander +// // Package commandermock is a generated GoMock package. package commandermock @@ -15,6 +20,7 @@ import ( type Mockcommander struct { ctrl *gomock.Controller recorder *MockcommanderMockRecorder + isgomock struct{} } // MockcommanderMockRecorder is the mock recorder for Mockcommander. diff --git a/core/bazel/query.go b/core/bazel/query.go index 9efc27db..39524960 100644 --- a/core/bazel/query.go +++ b/core/bazel/query.go @@ -70,9 +70,8 @@ func (b *BazelClient) executeQueryInternal(ctx context.Context, query string, st g.Go(func() error { return streamOutput(gCtx, stderr, &stderrBuf) }) - - streamErr := g.Wait() waitErr := cmd.Wait() + streamErr := g.Wait() // The command itself failed. if waitErr != nil { b.logger.Error("Bazel query failed failed: %v", waitErr) diff --git a/core/bazel/query_test.go b/core/bazel/query_test.go index bdab8145..449c66b8 100644 --- a/core/bazel/query_test.go +++ b/core/bazel/query_test.go @@ -116,7 +116,7 @@ func TestExecuteQuery_WithStartupOptions(t *testing.T) { }, capturedArgs) } -func TestexecuteQueryInternal_ContextTimeout(t *testing.T) { +func TestExecuteQueryInternal_ContextTimeout(t *testing.T) { defer goleak.VerifyNone(t) ctrl := gomock.NewController(t) mockCmd := commandermock.NewMockcommander(ctrl) @@ -129,7 +129,10 @@ func TestexecuteQueryInternal_ContextTimeout(t *testing.T) { mockCmd.EXPECT().StdoutPipe().Return(prStdout, nil), mockCmd.EXPECT().StderrPipe().Return(prStderr, nil), mockCmd.EXPECT().Start().Return(nil), - mockCmd.EXPECT().Wait().Return(context.DeadlineExceeded), + mockCmd.EXPECT().Wait().DoAndReturn(func() error { + // Simulate process ending after timeout + return context.DeadlineExceeded + }), ) client, err := NewBazelClient(Params{ @@ -137,21 +140,15 @@ func TestexecuteQueryInternal_ContextTimeout(t *testing.T) { WorkspacePath: "/tmp/test", Logger: zap.NewNop().Sugar(), EnvVarsMap: map[string]string{}, - QueryTimeout: 1 * time.Nanosecond, // Induce timeout immediately + QueryTimeout: 10 * time.Millisecond, // Short timeout for test ExecCommandContext: func(ctx context.Context, name string, arg ...string) commander { - // This goroutine simulates the OS/exec.Cmd behavior: - // When the context is canceled, the process is "killed", - // which closes its stdout/stderr pipes. + // Simulate process behavior: when context is cancelled, close pipes go func() { - <-ctx.Done() // Wait for the timeout to fire - - // "Killing" the process: close the pipes. - // This unblocks the Read() calls in your - // streamAndParseTargets and streamOutput goroutines. - // We close with the context's error so g.Wait() sees it. - pwStdout.CloseWithError(ctx.Err()) - pwStderr.CloseWithError(ctx.Err()) + <-ctx.Done() + // Close pipes to unblock readers + pwStdout.Close() + pwStderr.Close() }() return mockCmd }, @@ -159,10 +156,12 @@ func TestexecuteQueryInternal_ContextTimeout(t *testing.T) { require.NoError(t, err) result, err := client.executeQueryInternal(context.Background(), "//...", nil) require.Nil(t, result) - assert.Contains(t, err.Error(), "context deadline exceeded") + require.Error(t, err) + // Should get timeout or deadline exceeded error + assert.Contains(t, err.Error(), "deadline exceeded") } -func TestexecuteQueryInternal_Failures(t *testing.T) { +func TestExecuteQueryInternal_Failures(t *testing.T) { tests := []struct { name string setupMock func(*commandermock.Mockcommander) diff --git a/core/bazel/stream.go b/core/bazel/stream.go index 176bf24c..6c53e06d 100644 --- a/core/bazel/stream.go +++ b/core/bazel/stream.go @@ -2,9 +2,8 @@ package bazel import ( "bufio" - "context" "io" - + "context" buildpb "github.com/bazelbuild/buildtools/build_proto" "google.golang.org/protobuf/encoding/protodelim" ) @@ -32,7 +31,7 @@ func streamAndParseTargets(ctx context.Context, src io.Reader, dst io.Writer) (* done := make(chan result, 1) go func() { - queryResult, err := getQueryResult(src, dst) + queryResult, err := getQueryResult(ctx,src, dst) done <- result{queryResult: queryResult, err: err} }() @@ -44,25 +43,32 @@ func streamAndParseTargets(ctx context.Context, src io.Reader, dst io.Writer) (* } } + + // getQueryResult reads a QueryResult containing targets from the stream and returns it. -func getQueryResult(src io.Reader, dst io.Writer) (*buildpb.QueryResult, error) { +func getQueryResult(ctx context.Context, src io.Reader, dst io.Writer) (*buildpb.QueryResult, error) { result := &buildpb.QueryResult{ Target: make([]*buildpb.Target, 0), } tr := io.TeeReader(src, dst) br := bufio.NewReader(tr) - + unmarshalOpts := protodelim.UnmarshalOptions{ + MaxSize: 64 * 1024 * 1024, // 64MB limit + } + var parseErr error for { var target buildpb.Target - err := protodelim.UnmarshalFrom(br, &target) + err := unmarshalOpts.UnmarshalFrom(br, &target) if err == io.EOF { break } if err != nil { - return result, err + parseErr = err + // Continue reading - critical to prevent Bazel from blocking on write + continue } result.Target = append(result.Target, &target) } - return result, nil + return result, parseErr } diff --git a/core/bazelrunner/native.go b/core/bazelrunner/native.go index 78fa6c3e..7711eda2 100644 --- a/core/bazelrunner/native.go +++ b/core/bazelrunner/native.go @@ -43,6 +43,7 @@ func (g *nativeGraphRunner) Compute(ctx context.Context, ws workspace.Workspace) // --noproto: parameters exclude fields from the output that are not used for hashing anyways, making // proto blob smaller and serialization/deserialization faster // TODO: pass in --enable_workspace or --enable_bzlmod based on the config + AdditionalArgs: []string{"--order_output=no", "--proto:locations", "--noproto:default_values"}, }) if err != nil { diff --git a/example/cmd/query-bench/BUILD.bazel b/example/cmd/query-bench/BUILD.bazel new file mode 100644 index 00000000..06a7986d --- /dev/null +++ b/example/cmd/query-bench/BUILD.bazel @@ -0,0 +1,19 @@ +load("@rules_go//go:def.bzl", "go_binary", "go_library") + +go_library( + name = "query-bench_lib", + srcs = ["main.go"], + importpath = "github.com/uber/tango/example/cmd/query-bench", + visibility = ["//visibility:private"], + deps = [ + "//core/bazel", + "//core/targethasher", + "@org_uber_go_zap//:zap", + ], +) + +go_binary( + name = "query-bench", + embed = [":query-bench_lib"], + visibility = ["//visibility:public"], +) diff --git a/example/cmd/query-bench/main.go b/example/cmd/query-bench/main.go new file mode 100644 index 00000000..ae0f38bf --- /dev/null +++ b/example/cmd/query-bench/main.go @@ -0,0 +1,102 @@ +// query-bench is a local benchmarking tool for the bazel query execution path and targethasher. +// It runs the same query that nativeGraphRunner uses and reports timing and target counts. +// +// Usage: +// +// bazel run //cmd/query-bench -- --workspace /path/to/repo +// bazel run //cmd/query-bench -- --workspace /path/to/repo --bazel bazelisk --runs 3 +// bazel run //cmd/query-bench -- --workspace /path/to/repo --exclude-external +// bazel run //cmd/query-bench -- --workspace /path/to/repo --query '//...:all-targets' +package main + +import ( + "context" + "flag" + "fmt" + "os" + "time" + + "github.com/uber/tango/core/bazel" + "github.com/uber/tango/core/targethasher" + "go.uber.org/zap" +) + +func main() { + if err := run(); err != nil { + fmt.Fprintf(os.Stderr, "error: %v\n", err) + os.Exit(1) + } +} + +func run() error { + bazelCmd := flag.String("bazel", "", "bazel binary to invoke (default: auto-detect bazel/bazelisk on PATH)") + workspace := flag.String("workspace", ".", "workspace root to run bazel query in") + query := flag.String("query", "", "bazel query expression (default: the standard nativeGraphRunner query)") + excludeExternal := flag.Bool("exclude-external", false, "use deps(//...:all-targets) instead of including //external:all-targets") + runs := flag.Int("runs", 1, "number of times to run the query (for benchmarking)") + timeout := flag.Duration("timeout", 30*time.Minute, "per-run timeout") + flag.Parse() + + q := *query + if q == "" { + if *excludeExternal { + q = "deps(//...:all-targets)" + } else { + q = "//external:all-targets + deps(//...:all-targets)" + } + } + + logger, err := zap.NewDevelopment() + if err != nil { + return fmt.Errorf("creating logger: %w", err) + } + defer logger.Sync() + + client, err := bazel.NewBazelClient(bazel.Params{ + BazelCommand: *bazelCmd, + WorkspacePath: *workspace, + Logger: logger.Sugar(), + QueryTimeout: *timeout, + }) + if err != nil { + return fmt.Errorf("creating bazel client: %w", err) + } + + req := &bazel.QueryRequest{ + Query: q, + AdditionalArgs: []string{"--order_output=no", "--proto:locations", "--noproto:default_values"}, + } + + fmt.Printf("workspace: %s\n", *workspace) + fmt.Printf("query: %s\n", q) + fmt.Printf("runs: %d\n\n", *runs) + + var totalDuration time.Duration + for i := range *runs { + ctx, cancel := context.WithTimeout(context.Background(), *timeout) + start := time.Now() + resp, err := client.ExecuteQuery(ctx, req) + elapsed := time.Since(start) + defer cancel() + + if err != nil { + return fmt.Errorf("run %d: query failed: %w", i+1, err) + } + + totalDuration += elapsed + fmt.Printf("run %d: bazel query: %v (%d targets)\n", i+1, elapsed.Round(time.Millisecond), len(resp.Result.Target)) + start = time.Now() + targethasherResult, err := targethasher.FromProto(ctx, resp.Result, *workspace, targethasher.HashConfig{}) + if err != nil { + return fmt.Errorf("converting result to targethasher.Result: %w", err) + } + elapsed = time.Since(start) + totalDuration += elapsed + fmt.Printf("run %d: targethasher: %v (%d targets)\n", i+1, elapsed.Round(time.Millisecond), len(targethasherResult.TargetNames)) + } + + if *runs > 1 { + fmt.Printf("\naverage: %v\n", (totalDuration / time.Duration(*runs)).Round(time.Millisecond)) + } + return nil +}