Repository navigation
air: bound submission config validation timeout and retry #6994
New issue
Have a question about this project? Sign up for a free GitHub account to open an issue and contact its maintainers and the community.
By clicking “Sign up for GitHub”, you agree to our terms of service and privacy statement. We’ll occasionally send you account related emails.
Already on GitHub? Sign in to your account
base: main
Are you sure you want to change the base?
Changes from all commits
File filter
Filter by extension
Conversations
Jump to
Diff view
Diff view
There are no files selected for viewing
| 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)) |
| Original file line number | Diff line number | Diff line change |
|---|---|---|
|
|
@@ -2,6 +2,8 @@ package aircmd | |
|
|
||
| import ( | ||
| "context" | ||
| "crypto/tls" | ||
| "encoding/json" | ||
| "errors" | ||
| "fmt" | ||
| "net/http" | ||
|
|
@@ -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 | ||
|
|
@@ -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 { | ||
|
|
@@ -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 | ||
|
|
@@ -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 { | ||
| 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 { | ||
|
Contributor
There was a problem hiding this comment. Choose a reason for hiding this commentThe 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. |
||
| 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{ | ||
|
Contributor
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. This conversion drops |
||
| 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) { | ||
|
|
@@ -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 { | ||
|
Contributor
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. Real M2M OAuth 401 errors wrap |
||
| return true, true | ||
| } | ||
| } | ||
| apiErr, ok := errors.AsType[*apierr.APIError](err) | ||
| if !ok { | ||
|
|
@@ -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 { | ||
|
|
||
There was a problem hiding this comment.
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.