Skip to content
Merged
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/preflight-validation-timeout.md
Original file line number Diff line number Diff line change
@@ -0,0 +1 @@
* Limit `databricks air run` submission preflight validation to one five-second attempt before continuing when the service is unavailable. ([#6994](https://github.com/databricks/cli/pull/6994))
137 changes: 115 additions & 22 deletions cmd/air/validateconfig.go
Original file line number Diff line number Diff line change
Expand Up @@ -2,6 +2,8 @@ package aircmd

import (
"context"
"crypto/tls"
"encoding/json"
"errors"
"fmt"
"net/http"
Expand All @@ -10,11 +12,15 @@ import (
"time"

"github.com/databricks/cli/libs/auth"
"github.com/databricks/cli/libs/cmdio"
"github.com/databricks/databricks-sdk-go"
"github.com/databricks/databricks-sdk-go/apierr"
"github.com/databricks/databricks-sdk-go/client"
"github.com/databricks/databricks-sdk-go/common"
"github.com/databricks/databricks-sdk-go/config"
"github.com/databricks/databricks-sdk-go/httpclient"
"github.com/databricks/databricks-sdk-go/httpclient/traceparent"
"github.com/databricks/databricks-sdk-go/useragent"
)

// validateConfigPath is AiTrainingService's pre-flight: it checks a training
Expand All @@ -24,6 +30,8 @@ const validateConfigPath = "/api/2.0/ai-training/config:validate"

const dryRunValidationMaxAttempts = 2

const submissionValidationTimeout = 5 * time.Second

type dryRunValidationAttemptBudgetKey struct{}

func armDryRunValidationAttemptBudget(ctx context.Context) context.Context {
Expand Down Expand Up @@ -51,6 +59,13 @@ type validateConfigResponse struct {
Errors []configFieldError `json:"errors"`
}

type validationAuthenticationError struct {
err error
}

func (e *validationAuthenticationError) Error() string { return e.err.Error() }
func (e *validationAuthenticationError) Unwrap() error { return e.err }

// validationUnavailableError means the backend check could not finish.
type validationUnavailableError struct {
err error
Expand All @@ -74,18 +89,103 @@ func validationUnavailable(err error) (*validationUnavailableError, bool) {
// preflightValidate checks the config against the backend before any upload, so
// a bad config fails fast with the server's own field-level errors.
//
// It preserves the existing submission behavior: a missing endpoint or 5xx
// fails open, while other request failures and field errors block.
// Availability failures fail open, while caller failures and field errors block.
func preflightValidate(ctx context.Context, w *databricks.WorkspaceClient, cfg *runConfig, commandPath string, containers []submittedContainer, idempotencyToken string) error {
apiClient, err := client.New(w.Config)
validationCtx, cancel := context.WithTimeout(ctx, submissionValidationTimeout)
defer cancel()

err := validateConfigOnce(validationCtx, w, cfg, commandPath, containers, idempotencyToken)
if unavailable, _ := classifyValidationFailure(err); unavailable {
Comment thread
caroline-db marked this conversation as resolved.
if cmdio.HasIO(ctx) {
cmdio.LogString(ctx, "Warning: server-side config validation was unavailable; continuing with submission.")
}
return nil
}
return err
}

func validateConfigOnce(ctx context.Context, w *databricks.WorkspaceClient, cfg *runConfig, commandPath string, containers []submittedContainer, idempotencyToken string) error {
clientCfg, err := config.HTTPClientConfigFromConfig(w.Config)
if err != nil {
return fmt.Errorf("failed to create API client: %w", err)
}
err = validateConfig(ctx, apiClient, cfg, commandPath, containers, idempotencyToken)
if endpointUnavailable(err) || serverError(err) {

requestBody, err := common.NewRequestBody(validateConfigRequest(ctx, cfg, commandPath, containers, idempotencyToken))
if err != nil {
return fmt.Errorf("failed to validate config: %w", err)
}
req, err := http.NewRequestWithContext(ctx, http.MethodPost, validateConfigPath, requestBody.Reader)
if err != nil {
return fmt.Errorf("failed to validate config: %w", err)
}
req.Header.Set("Accept", "application/json")
req.Header.Set("Content-Type", requestBody.ContentType)
if clientCfg.AuthVisitor != nil {
// Some SDK OAuth visitors refresh with context.Background(). For those
// providers, acquire the token with the preflight context first.
switch w.Config.AuthType {
case auth.AuthTypePat, auth.AuthTypeBasic, "noop", auth.AuthTypeAzureCli, auth.AuthTypeAzureMSI, auth.AuthTypeAzureSecret, auth.AuthTypeGoogleCreds, auth.AuthTypeGoogleID:
default:
token, err := w.Config.GetTokenSource().Token(req.Context())
if err != nil {
return &validationAuthenticationError{fmt.Errorf("failed to validate config: %w", err)}
}
token.SetAuthHeader(req)
}
if err := clientCfg.AuthVisitor(req); err != nil {
Comment thread
caroline-db marked this conversation as resolved.
return &validationAuthenticationError{fmt.Errorf("failed to validate config: %w", err)}
}
}
for _, visitor := range clientCfg.Visitors {
if err := visitor(req); err != nil {
return fmt.Errorf("failed to validate config: %w", err)
}
}
for key, value := range auth.WorkspaceIDHeaders(w.Config) {
req.Header.Set(key, value)
}
req.Header.Set("User-Agent", useragent.FromContext(req.Context()))
traceparent.AddTraceparent(req)

transport := clientCfg.Transport
if transport == nil {
defaultTransport := http.DefaultTransport.(*http.Transport).Clone()
if clientCfg.InsecureSkipVerify {
defaultTransport.TLSClientConfig = &tls.Config{InsecureSkipVerify: true}
}
transport = defaultTransport
}
resp, err := (&http.Client{Transport: transport}).Do(req)
if err != nil {
return fmt.Errorf("failed to validate config: %w", err)
}
defer resp.Body.Close()
responseBody, err := common.NewResponseWrapper(resp, requestBody)
if err != nil {
err = fmt.Errorf("failed to validate config: %w", err)
if errors.Is(err, context.Canceled) {
return err
}
if resp.StatusCode >= 400 {
return &apierr.APIError{
Comment thread
caroline-db marked this conversation as resolved.
StatusCode: resp.StatusCode,
Message: err.Error(),
}
}
return asValidationUnavailable(err, true)
}
if err := apierr.GetAPIError(ctx, responseBody); err != nil {
return fmt.Errorf("failed to validate config: %w", err)
}

var result validateConfigResponse
if err := json.Unmarshal(responseBody.DebugBytes, &result); err != nil {
return asValidationUnavailable(fmt.Errorf("failed to validate config: %w", err), true)
}
if len(result.Errors) == 0 {
return nil
}
return err
return errors.New(formatConfigErrors(result.Errors))
}

func newDryRunValidationClient(w *databricks.WorkspaceClient) (*client.DatabricksClient, error) {
Expand Down Expand Up @@ -145,11 +245,19 @@ func classifyValidationFailure(err error) (unavailable, retryable bool) {
if errors.Is(err, context.Canceled) {
return false, false
}
if unavailable, ok := validationUnavailable(err); ok {
return true, unavailable.retryable
}
if errors.Is(err, context.DeadlineExceeded) {
return true, true
}
if _, ok := errors.AsType[*validationAuthenticationError](err); ok {
return false, false
}
if _, ok := errors.AsType[*url.Error](err); ok {
return true, true
if _, ok := errors.AsType[*apierr.APIError](err); !ok {
Comment thread
caroline-db marked this conversation as resolved.
return true, true
}
}
apiErr, ok := errors.AsType[*apierr.APIError](err)
if !ok {
Expand Down Expand Up @@ -246,21 +354,6 @@ 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.
func endpointUnavailable(err error) bool {
apiErr, ok := errors.AsType[*apierr.APIError](err)
return ok && (apiErr.ErrorCode == "FEATURE_DISABLED" ||
apiErr.StatusCode == http.StatusNotFound ||
apiErr.StatusCode == http.StatusNotImplemented)
}

// serverError reports a backend 5xx rather than a caller error.
func serverError(err error) bool {
apiErr, ok := errors.AsType[*apierr.APIError](err)
return ok && apiErr.StatusCode >= 500
}

// formatConfigErrors renders the field errors as one message, one problem per
// line, each pointing at the config field the user wrote.
func formatConfigErrors(fieldErrors []configFieldError) string {
Expand Down
Loading
Loading