diff --git a/.nextchanges/air/preflight-validation-timeout.md b/.nextchanges/air/preflight-validation-timeout.md new file mode 100644 index 00000000000..f6a31d5fa0a --- /dev/null +++ b/.nextchanges/air/preflight-validation-timeout.md @@ -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)) diff --git a/cmd/air/validateconfig.go b/cmd/air/validateconfig.go index b72d0c72a02..62dad63f8ab 100644 --- a/cmd/air/validateconfig.go +++ b/cmd/air/validateconfig.go @@ -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 { + 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{ + 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 { + 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 { diff --git a/cmd/air/validateconfig_test.go b/cmd/air/validateconfig_test.go index c79a8208304..f0a0f3e38c8 100644 --- a/cmd/air/validateconfig_test.go +++ b/cmd/air/validateconfig_test.go @@ -4,15 +4,24 @@ import ( "context" "encoding/json" "errors" + "io" "net/http" "net/http/httptest" "net/url" + "strings" + "sync/atomic" "testing" + "time" + "github.com/databricks/cli/libs/cmdio" "github.com/databricks/databricks-sdk-go" "github.com/databricks/databricks-sdk-go/apierr" + sdkconfig "github.com/databricks/databricks-sdk-go/config" + "github.com/databricks/databricks-sdk-go/config/credentials" + sdkauth "github.com/databricks/databricks-sdk-go/config/experimental/auth" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" + "golang.org/x/oauth2" ) func baseRunConfig() *runConfig { @@ -48,6 +57,48 @@ func validationTestWorkspaceClient(t *testing.T, host string) *databricks.Worksp return w } +func validationTestWorkspaceClientWithTransport(t *testing.T, transport validationRoundTripFunc) *databricks.WorkspaceClient { + t.Helper() + w, err := databricks.NewWorkspaceClient(&databricks.Config{ + Host: "https://example.test", + Token: "token", + HTTPTransport: transport, + }) + require.NoError(t, err) + return w +} + +type validationRoundTripFunc func(*http.Request) (*http.Response, error) + +func (fn validationRoundTripFunc) RoundTrip(req *http.Request) (*http.Response, error) { + return fn(req) +} + +type validationOAuthCredentials struct { + tokenSource sdkauth.TokenSource +} + +func (validationOAuthCredentials) Name() string { return "oauth-m2m" } + +func (c validationOAuthCredentials) Configure(context.Context, *sdkconfig.Config) (credentials.CredentialsProvider, error) { + return credentials.NewOAuthCredentialsProviderFromTokenSource(c.tokenSource), nil +} + +type validationErrorReadCloser struct{} + +func (validationErrorReadCloser) Read([]byte) (int, error) { return 0, errors.New("body read failed") } +func (validationErrorReadCloser) Close() error { return nil } + +type validationContextReadCloser struct { + ctx context.Context +} + +func (r validationContextReadCloser) Read([]byte) (int, error) { + <-r.ctx.Done() + return 0, r.ctx.Err() +} +func (validationContextReadCloser) Close() error { return nil } + func TestPreflightValidateFailsOpenWhenBackendUnavailable(t *testing.T) { tests := []struct { name string @@ -67,6 +118,240 @@ func TestPreflightValidateFailsOpenWhenBackendUnavailable(t *testing.T) { } } +func TestPreflightValidationUnavailableWarnsWithoutRetry(t *testing.T) { + var requests atomic.Int32 + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.URL.Path == validateConfigPath { + requests.Add(1) + w.WriteHeader(http.StatusServiceUnavailable) + _, _ = w.Write([]byte(`{"error_code":"TEMPORARILY_UNAVAILABLE","message":"try again"}`)) + return + } + _, _ = w.Write([]byte(`{}`)) + })) + t.Cleanup(srv.Close) + ctx, stderr := cmdio.NewTestContextWithStderr(t.Context()) + + err := preflightValidate(ctx, validationTestWorkspaceClient(t, srv.URL), baseRunConfig(), "/Workspace/Users/me/cmd.sh", nil, "token") + require.NoError(t, err) + assert.Equal(t, int32(1), requests.Load()) + assert.Equal(t, "Warning: server-side config validation was unavailable; continuing with submission.\n", stderr.String()) +} + +func TestPreflightValidationDoesNotRetryTransportFailure(t *testing.T) { + var requests atomic.Int32 + w := validationTestWorkspaceClientWithTransport(t, func(req *http.Request) (*http.Response, error) { + if req.URL.Path == validateConfigPath { + requests.Add(1) + } + return nil, errors.New("connection failed") + }) + + err := preflightValidate(t.Context(), w, baseRunConfig(), "/Workspace/Users/me/cmd.sh", nil, "token") + require.NoError(t, err) + assert.Equal(t, int32(1), requests.Load()) +} + +func TestPreflightValidationFailsOpenOnInvalidSuccessResponse(t *testing.T) { + tests := []struct { + name string + body func() io.ReadCloser + }{ + {"body read failure", func() io.ReadCloser { return validationErrorReadCloser{} }}, + {"malformed JSON", func() io.ReadCloser { return io.NopCloser(strings.NewReader("{")) }}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + var requests atomic.Int32 + w := validationTestWorkspaceClientWithTransport(t, func(req *http.Request) (*http.Response, error) { + body := io.NopCloser(strings.NewReader(`{}`)) + if req.URL.Path == validateConfigPath { + requests.Add(1) + body = tt.body() + } + return &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{"Content-Type": []string{"application/json"}}, + Body: body, + Request: req, + }, nil + }) + + err := preflightValidate(t.Context(), w, baseRunConfig(), "/Workspace/Users/me/cmd.sh", nil, "token") + require.NoError(t, err) + assert.Equal(t, int32(1), requests.Load()) + }) + } +} + +func TestPreflightValidationBodyReadFailureOnCallerErrorBlocks(t *testing.T) { + w := validationTestWorkspaceClientWithTransport(t, func(req *http.Request) (*http.Response, error) { + body := io.NopCloser(strings.NewReader(`{}`)) + status := http.StatusOK + if req.URL.Path == validateConfigPath { + body = validationContextReadCloser{ctx: req.Context()} + status = http.StatusUnauthorized + } + return &http.Response{ + StatusCode: status, + Header: http.Header{"Content-Type": []string{"application/json"}}, + Body: body, + Request: req, + }, nil + }) + + ctx, cancel := context.WithTimeout(t.Context(), 50*time.Millisecond) + defer cancel() + err := preflightValidate(ctx, w, baseRunConfig(), "/Workspace/Users/me/cmd.sh", nil, "token") + require.Error(t, err) + apiErr, ok := errors.AsType[*apierr.APIError](err) + require.True(t, ok) + assert.Equal(t, http.StatusUnauthorized, apiErr.StatusCode) +} + +func TestPreflightValidationOAuthRefreshUsesDeadline(t *testing.T) { + refreshStarted := make(chan struct{}) + tokenSource := sdkauth.NewCachedTokenSource( + sdkauth.TokenSourceFn(func(ctx context.Context) (*oauth2.Token, error) { + close(refreshStarted) + <-ctx.Done() + return nil, ctx.Err() + }), + sdkauth.WithCachedToken(&oauth2.Token{ + AccessToken: "expired", + Expiry: time.Now().Add(-time.Hour), + }), + ) + w, err := databricks.NewWorkspaceClient(&databricks.Config{ + Host: "https://example.test", + Credentials: validationOAuthCredentials{tokenSource: tokenSource}, + HTTPTransport: validationRoundTripFunc(func(*http.Request) (*http.Response, error) { + t.Fatal("validation request should not be sent") + return nil, nil + }), + HostMetadataResolver: func(context.Context, string) (*sdkconfig.HostMetadata, error) { + return &sdkconfig.HostMetadata{}, nil + }, + }) + require.NoError(t, err) + + ctx, cancel := context.WithTimeout(t.Context(), 50*time.Millisecond) + defer cancel() + started := time.Now() + err = preflightValidate(ctx, w, baseRunConfig(), "/Workspace/Users/me/cmd.sh", nil, "token") + + require.NoError(t, err) + assert.Less(t, time.Since(started), time.Second) + select { + case <-refreshStarted: + default: + t.Fatal("OAuth token refresh did not start") + } +} + +func TestPreflightValidationM2MOAuthRejectionBlocks(t *testing.T) { + tests := []struct { + status int + wantErr error + }{ + {http.StatusBadRequest, apierr.ErrBadRequest}, + {http.StatusUnauthorized, apierr.ErrUnauthenticated}, + {http.StatusForbidden, apierr.ErrPermissionDenied}, + } + + for _, tt := range tests { + t.Run(http.StatusText(tt.status), func(t *testing.T) { + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + switch r.URL.Path { + case "/oidc/.well-known/oauth-authorization-server": + _ = json.NewEncoder(w).Encode(map[string]string{"token_endpoint": "http://" + r.Host + "/token"}) + case "/token": + w.WriteHeader(tt.status) + _, _ = w.Write([]byte(`{"error":"invalid_client"}`)) + } + })) + t.Cleanup(srv.Close) + + w, err := databricks.NewWorkspaceClient(&databricks.Config{ + Host: srv.URL, + ClientID: "client-id", + ClientSecret: "client-secret", + AuthType: "oauth-m2m", + DiscoveryURL: srv.URL + "/oidc/.well-known/oauth-authorization-server", + HostMetadataResolver: func(context.Context, string) (*sdkconfig.HostMetadata, error) { + return &sdkconfig.HostMetadata{}, nil + }, + }) + require.NoError(t, err) + + err = preflightValidate(t.Context(), w, baseRunConfig(), "/Workspace/Users/me/cmd.sh", nil, "token") + require.ErrorIs(t, err, tt.wantErr) + }) + } +} + +func TestPreflightValidationTimeoutFailsOpen(t *testing.T) { + releaseRequest := make(chan struct{}) + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.URL.Path == validateConfigPath { + <-releaseRequest + return + } + _, _ = w.Write([]byte(`{}`)) + })) + t.Cleanup(srv.Close) + + ctx, cancel := context.WithTimeout(t.Context(), 50*time.Millisecond) + defer cancel() + started := time.Now() + err := preflightValidate(ctx, validationTestWorkspaceClient(t, srv.URL), baseRunConfig(), "/Workspace/Users/me/cmd.sh", nil, "token") + close(releaseRequest) + require.NoError(t, err) + assert.Less(t, time.Since(started), time.Second) +} + +func TestPreflightValidationPropagatesCancellation(t *testing.T) { + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.URL.Path == validateConfigPath { + <-r.Context().Done() + return + } + _, _ = w.Write([]byte(`{}`)) + })) + t.Cleanup(srv.Close) + + ctx, cancel := context.WithCancel(t.Context()) + cancel() + err := preflightValidate(ctx, validationTestWorkspaceClient(t, srv.URL), baseRunConfig(), "/Workspace/Users/me/cmd.sh", nil, "token") + require.ErrorIs(t, err, context.Canceled) +} + +func TestPreflightValidationPropagatesCancellationWhileReadingErrorResponse(t *testing.T) { + validationStarted := make(chan struct{}) + w := validationTestWorkspaceClientWithTransport(t, func(req *http.Request) (*http.Response, error) { + body := io.NopCloser(strings.NewReader(`{}`)) + if req.URL.Path == validateConfigPath { + close(validationStarted) + body = validationContextReadCloser{ctx: req.Context()} + } + return &http.Response{ + StatusCode: http.StatusServiceUnavailable, + Header: http.Header{"Content-Type": []string{"application/json"}}, + Body: body, + Request: req, + }, nil + }) + + ctx, cancel := context.WithCancel(t.Context()) + go func() { + <-validationStarted + cancel() + }() + err := preflightValidate(ctx, w, baseRunConfig(), "/Workspace/Users/me/cmd.sh", nil, "token") + require.ErrorIs(t, err, context.Canceled) +} + func TestPreflightValidateBlocksOnCallerError(t *testing.T) { for _, status := range []int{http.StatusBadRequest, http.StatusUnauthorized, http.StatusForbidden} { t.Run(http.StatusText(status), func(t *testing.T) {