diff --git a/.nextchanges/air/workspace-compute-help.md b/.nextchanges/air/workspace-compute-help.md new file mode 100644 index 00000000000..4cd5dfcd3c2 --- /dev/null +++ b/.nextchanges/air/workspace-compute-help.md @@ -0,0 +1 @@ +* Show only the accelerator types available in the current workspace in `databricks air run -h config.compute`. ([#6993](https://github.com/databricks/cli/pull/6993)) diff --git a/acceptance/air/config-help/output.txt b/acceptance/air/config-help/output.txt index f73f9618530..be048a1e392 100644 --- a/acceptance/air/config-help/output.txt +++ b/acceptance/air/config-help/output.txt @@ -74,6 +74,10 @@ Use "-h config." for details on a field. config.compute Which accelerators to run on and how many. + Supported accelerator types: + GPU_1xH100 - 1 accelerator per node [Private Preview] + GPU_8xH100 - 8 accelerators per node + Fields: num_accelerators Total number of GPUs to allocate. accelerator_type Which accelerator to run on, e.g. GPU_1xA10. diff --git a/acceptance/air/config-help/test.toml b/acceptance/air/config-help/test.toml new file mode 100644 index 00000000000..965694857c9 --- /dev/null +++ b/acceptance/air/config-help/test.toml @@ -0,0 +1,16 @@ +# The SDK occasionally probes host reachability with a HEAD request; stub it so +# the test is deterministic. +[[Server]] +Pattern = "HEAD /" +Response.Body = '' + +[[Server]] +Pattern = "GET /api/2.0/ai-training/compute-options" +Response.Body = ''' +{ + "compute_options": [ + {"hardware_accelerator": "GPU_1xH100", "per_node_accelerator_count": 1, "launch_stage": "PRIVATE_PREVIEW"}, + {"hardware_accelerator": "GPU_8xH100", "per_node_accelerator_count": 8, "launch_stage": "GA"} + ] +} +''' diff --git a/cmd/air/aitraining.go b/cmd/air/aitraining.go index 594489a5d74..c13db8e2538 100644 --- a/cmd/air/aitraining.go +++ b/cmd/air/aitraining.go @@ -18,6 +18,34 @@ import ( // with a raw client.Do because the SDK does not model the AiTrainingService. const aiTrainingWorkflowsPath = "/api/2.0/ai-training/workflows" +const computeOptionsPath = "/api/2.0/ai-training/compute-options" + +type computeOption struct { + HardwareAccelerator string `json:"hardware_accelerator"` + DisplayName string `json:"display_name"` + MultiNodeSupported *bool `json:"multi_node_supported"` + PerNodeAcceleratorCount *int `json:"per_node_accelerator_count"` + LaunchStage string `json:"launch_stage"` +} + +type computeOptionsResponse struct { + ComputeOptions []computeOption `json:"compute_options"` +} + +func listWorkspaceComputeOptions(ctx context.Context, w *databricks.WorkspaceClient) ([]computeOption, error) { + apiClient, err := client.New(w.Config) + if err != nil { + return nil, fmt.Errorf("failed to create API client: %w", err) + } + + var resp computeOptionsResponse + err = apiClient.Do(ctx, http.MethodGet, computeOptionsPath, auth.WorkspaceIDHeaders(w.Config), nil, nil, &resp) + if err != nil { + return nil, fmt.Errorf("failed to list compute options: %w", err) + } + return resp.ComputeOptions, nil +} + // workflowRef is one run from the index: its Jobs run id and submission time. type workflowRef struct { jobRunID int64 diff --git a/cmd/air/aitraining_test.go b/cmd/air/aitraining_test.go index 4c44266d8cb..fdcb250cbab 100644 --- a/cmd/air/aitraining_test.go +++ b/cmd/air/aitraining_test.go @@ -30,6 +30,31 @@ func TestParseSubmitTimeMs(t *testing.T) { } } +func TestListWorkspaceComputeOptions(t *testing.T) { + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.URL.Path == "/.well-known/databricks-config" { + _, _ = w.Write([]byte(`{}`)) + return + } + assert.Equal(t, http.MethodGet, r.Method) + assert.Equal(t, computeOptionsPath, r.URL.Path) + assert.Equal(t, "123", r.Header.Get("X-Databricks-Workspace-Id")) + _, _ = w.Write([]byte(`{"compute_options":[{"hardware_accelerator":"GPU_8xH100","display_name":"8x H100","multi_node_supported":true,"per_node_accelerator_count":8,"launch_stage":"PUBLIC_PREVIEW"}]}`)) + })) + t.Cleanup(srv.Close) + w := newTestWorkspaceClient(t, srv.URL) + w.Config.WorkspaceID = "123" + + options, err := listWorkspaceComputeOptions(t.Context(), w) + require.NoError(t, err) + require.Len(t, options, 1) + assert.Equal(t, "GPU_8xH100", options[0].HardwareAccelerator) + assert.Equal(t, "8x H100", options[0].DisplayName) + assert.True(t, *options[0].MultiNodeSupported) + assert.Equal(t, 8, *options[0].PerNodeAcceleratorCount) + assert.Equal(t, "PUBLIC_PREVIEW", options[0].LaunchStage) +} + // indexServer serves paginated AiTrainingService responses, one body per call, // tracking whether the index was hit. func indexServer(t *testing.T, hit *bool, bodies ...string) *httptest.Server { diff --git a/cmd/air/run.go b/cmd/air/run.go index c8027b09e9a..5fa2dd0641c 100644 --- a/cmd/air/run.go +++ b/cmd/air/run.go @@ -24,6 +24,9 @@ import ( // dryRunValidationTimeout bounds only the config:validate request. const dryRunValidationTimeout = 15 * time.Second +// Bounds the workspace lookup performed by config help. +const computeOptionsHelpTimeout = 5 * time.Second + // runResult is the JSON payload for `air run`. type runResult struct { Status string `json:"status"` @@ -112,9 +115,17 @@ The path must be a separate argument: cobra reserves -h as a boolean, so _ = c.Usage() return } - if err := writeConfigFieldHelp(c.OutOrStdout(), fields[0]); err != nil { + field, err := resolveConfigField(fields[0]) + if err != nil { c.PrintErrln("Error:", err) + return + } + + var computeOptions []computeOption + if field.path == "config.compute" { + computeOptions = loadComputeOptionsForHelp(c, args) } + renderConfigField(c.OutOrStdout(), field, computeOptions) }) cmd.Flags().StringVarP(&file, "file", "f", "", "Path to the workload YAML config") @@ -265,6 +276,32 @@ The path must be a separate argument: cobra reserves -h as a boolean, so return cmd } +func loadComputeOptionsForHelp(cmd *cobra.Command, args []string) []computeOption { + ctx, cancel := context.WithTimeout(cmd.Context(), computeOptionsHelpTimeout) + defer cancel() + cmd.SetContext(root.SkipLoadBundle(root.SkipPrompt(ctx))) + + if !cmdctx.HasWorkspaceClient(cmd.Context()) { + if err := root.MustWorkspaceClient(cmd, args); err != nil { + cmd.PrintErrf("Warning: couldn't resolve a workspace to list accelerator types; showing the built-in list instead: %v\n", err) + return fallbackComputeOptions() + } + } + + options, err := listWorkspaceComputeOptions(cmd.Context(), cmdctx.WorkspaceClient(cmd.Context())) + if err != nil { + if !endpointUnavailable(err) { + cmd.PrintErrf("Warning: couldn't fetch the accelerator types available in this workspace; showing the built-in list instead: %v\n", err) + } + return fallbackComputeOptions() + } + if len(options) == 0 { + cmd.PrintErrln("Warning: the workspace reported no supported accelerator types; showing the built-in list instead") + return fallbackComputeOptions() + } + return options +} + func airLogsCommand(profile, runID string) string { args := []string{"databricks", "air", "logs", shellquote.BashArg(runID)} if profile != "" { diff --git a/cmd/air/runconfig.go b/cmd/air/runconfig.go index 53846400a68..6f3c9e8970e 100644 --- a/cmd/air/runconfig.go +++ b/cmd/air/runconfig.go @@ -610,7 +610,7 @@ func writeConfigFieldHelp(w io.Writer, path string) error { if err != nil { return err } - renderConfigField(w, field) + renderConfigField(w, field, nil) return nil } @@ -857,11 +857,17 @@ func underlyingConfigStruct(t reflect.Type) reflect.Type { // renderConfigField writes a resolved field's documentation. An object lists its // immediate children; a leaf gets its type, required-ness, and description. -func renderConfigField(w io.Writer, f configField) { +func renderConfigField(w io.Writer, f configField, computeOptions []computeOption) { fmt.Fprintf(w, "%s\n", f.path) if f.help != "" { fmt.Fprintf(w, " %s\n", f.help) } + if f.path == "config.compute" && len(computeOptions) > 0 { + fmt.Fprintln(w, "\n Supported accelerator types:") + for _, option := range computeOptions { + fmt.Fprintf(w, " %s - %s%s\n", option.HardwareAccelerator, formatPerNodeAcceleratorCount(option), launchStageBadge(option.LaunchStage)) + } + } if len(f.children) == 0 { fmt.Fprintf(w, "\n Type: %s\n", f.typeName) @@ -890,6 +896,53 @@ func renderConfigField(w io.Writer, f configField) { fmt.Fprintf(w, "\nUse \"-h %s.\" for details on a field.\n", f.path) } +func fallbackComputeOptions() []computeOption { + options := make([]computeOption, 0, len(gpuTypes)) + for _, acceleratorType := range gpuTypes { + perNode, err := gpusPerNode(acceleratorType) + if err != nil { + continue + } + options = append(options, computeOption{ + HardwareAccelerator: string(acceleratorType), + PerNodeAcceleratorCount: &perNode, + }) + } + return options +} + +func formatPerNodeAcceleratorCount(option computeOption) string { + count := option.PerNodeAcceleratorCount + if count == nil { + acceleratorType, err := parseGPUType(option.HardwareAccelerator) + if err == nil { + if fallback, err := gpusPerNode(acceleratorType); err == nil { + count = &fallback + } + } + } + if count == nil { + return "accelerator count per node unavailable" + } + unit := "accelerators" + if *count == 1 { + unit = "accelerator" + } + return fmt.Sprintf("%d %s per node", *count, unit) +} + +func launchStageBadge(stage string) string { + labels := map[string]string{ + "PRIVATE_PREVIEW": "Private Preview", + "PUBLIC_BETA": "Beta", + "PUBLIC_PREVIEW": "Public Preview", + } + if label := labels[stage]; label != "" { + return " [" + label + "]" + } + return "" +} + // configFieldSummary is the one-line description used in a field listing: the // first sentence of the help text, annotated when the field is required. func configFieldSummary(f configField) string { diff --git a/cmd/air/runconfig_test.go b/cmd/air/runconfig_test.go index 798d23a6b5f..f360fe4f56f 100644 --- a/cmd/air/runconfig_test.go +++ b/cmd/air/runconfig_test.go @@ -2,11 +2,14 @@ package aircmd import ( "fmt" + "net/http" + "net/http/httptest" "os" "path/filepath" "strings" "testing" + "github.com/databricks/cli/libs/cmdctx" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" ) @@ -756,6 +759,82 @@ func TestWriteConfigFieldHelp(t *testing.T) { require.Error(t, writeConfigFieldHelp(&strings.Builder{}, "config.nope")) } +func TestRunCommandHelpComputeOptions(t *testing.T) { + builtInOptions := ` GPU_1xA10 - 1 accelerator per node + GPU_1xH100 - 1 accelerator per node + GPU_8xH100 - 8 accelerators per node + GPU_8xB300 - 8 accelerators per node` + tests := []struct { + name string + statusCode int + response string + wantOptions string + wantWarning bool + }{ + { + name: "workspace subset", + response: `{"compute_options":[{"hardware_accelerator":"GPU_8xH100","per_node_accelerator_count":4,"launch_stage":"PUBLIC_PREVIEW"},{"hardware_accelerator":"GPU_FUTURE","per_node_accelerator_count":16},{"hardware_accelerator":"GPU_1xA10"}]}`, + wantOptions: " GPU_8xH100 - 4 accelerators per node [Public Preview]\n GPU_FUTURE - 16 accelerators per node\n GPU_1xA10 - 1 accelerator per node", + }, + { + name: "empty response falls back", + response: `{}`, + wantOptions: builtInOptions, + wantWarning: true, + }, + { + name: "disabled endpoint falls back", + statusCode: http.StatusBadRequest, + response: `{"error_code":"FEATURE_DISABLED","message":"ListComputeOptions is not yet enabled."}`, + wantOptions: builtInOptions, + }, + { + name: "missing endpoint falls back", + statusCode: http.StatusNotFound, + response: `{"error_code":"ENDPOINT_NOT_FOUND","message":"Not found."}`, + wantOptions: builtInOptions, + }, + { + name: "unexpected API error warns and falls back", + statusCode: http.StatusInternalServerError, + response: `{"error_code":"INTERNAL_ERROR","message":"Internal error."}`, + wantOptions: builtInOptions, + wantWarning: true, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if tt.statusCode != 0 { + w.WriteHeader(tt.statusCode) + } + _, _ = w.Write([]byte(tt.response)) + })) + t.Cleanup(srv.Close) + + var out, errOut strings.Builder + cmd := newRunCommand() + cmd.SetContext(cmdctx.SetWorkspaceClient(t.Context(), newTestWorkspaceClient(t, srv.URL))) + cmd.SetOut(&out) + cmd.SetErr(&errOut) + cmd.SetArgs([]string{"-h", "config.compute"}) + require.NoError(t, cmd.Execute()) + + options, found := strings.CutPrefix(out.String(), "config.compute\n Which accelerators to run on and how many.\n\n Supported accelerator types:\n") + require.True(t, found) + options, _, found = strings.Cut(options, "\n\n Fields:") + require.True(t, found) + assert.Equal(t, tt.wantOptions, options) + if tt.wantWarning { + assert.Contains(t, errOut.String(), "showing the built-in list") + } else { + assert.Empty(t, errOut.String()) + } + }) + } +} + // Guards against adding a schema field without a help: tag. func TestConfigFieldsAllDocumented(t *testing.T) { root, err := resolveConfigField("config") diff --git a/cmd/air/validateconfig.go b/cmd/air/validateconfig.go index b72d0c72a02..e5043440bf4 100644 --- a/cmd/air/validateconfig.go +++ b/cmd/air/validateconfig.go @@ -246,8 +246,7 @@ func putOpt[T any](m map[string]any, key string, value *T) { } } -// endpointUnavailable reports that the validation endpoint could not answer -// because it is disabled or absent. +// endpointUnavailable reports that an API endpoint is disabled or absent. func endpointUnavailable(err error) bool { apiErr, ok := errors.AsType[*apierr.APIError](err) return ok && (apiErr.ErrorCode == "FEATURE_DISABLED" ||