Skip to content
Open
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 {

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

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

Please distinguish authentication failures from validation-service failures. An OAuth token refresh rejected with HTTP 401 is wrapped in url.Error, so this classifier incorrectly allows preflight to succeed.

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 {

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

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

Please make OAuth token acquisition use the preflight context and add a test with an expired token and a stalled refresh endpoint. AuthVisitor currently refreshes with context.Background(), so preflight can exceed five seconds; the test should verify it returns when the deadline expires.

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{

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

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

This conversion drops context.Canceled from the error chain, allowing cancellation during a 404/429/503 body read to become successful preflight. Please propagate cancellation before converting the error and add coverage for cancellation while reading an error response.

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 {

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

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

Real M2M OAuth 401 errors wrap config.tokenError, not *apierr.APIError, so this still allows preflight to succeed. Please keep authentication failures blocking and add a regression test that mocks an HTTP 401 response during token refresh, letting the SDK construct the error.

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