Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
1 change: 1 addition & 0 deletions .nextchanges/air/workspace-compute-help.md
Original file line number Diff line number Diff line change
@@ -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))
4 changes: 4 additions & 0 deletions acceptance/air/config-help/output.txt
Original file line number Diff line number Diff line change
Expand Up @@ -74,6 +74,10 @@ Use "-h config.<field>" 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.
Expand Down
16 changes: 16 additions & 0 deletions acceptance/air/config-help/test.toml
Original file line number Diff line number Diff line change
@@ -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"}
]
}
'''
28 changes: 28 additions & 0 deletions cmd/air/aitraining.go
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
25 changes: 25 additions & 0 deletions cmd/air/aitraining_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -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 {
Expand Down
39 changes: 38 additions & 1 deletion cmd/air/run.go
Original file line number Diff line number Diff line change
Expand Up @@ -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"`
Expand Down Expand Up @@ -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")
Expand Down Expand Up @@ -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()) {

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

make sure if any of these are RPC calls that they have retries

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()))

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

ditto

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 != "" {
Expand Down
57 changes: 55 additions & 2 deletions cmd/air/runconfig.go
Original file line number Diff line number Diff line change
Expand Up @@ -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
}

Expand Down Expand Up @@ -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)
Expand Down Expand Up @@ -890,6 +896,53 @@ func renderConfigField(w io.Writer, f configField) {
fmt.Fprintf(w, "\nUse \"-h %s.<field>\" 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 {
Expand Down
79 changes: 79 additions & 0 deletions cmd/air/runconfig_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -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"
)
Expand Down Expand Up @@ -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")
Expand Down
3 changes: 1 addition & 2 deletions cmd/air/validateconfig.go
Original file line number Diff line number Diff line change
Expand Up @@ -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" ||
Expand Down
Loading