From 07900aa6c71448fec055567b15c9f9924b147257 Mon Sep 17 00:00:00 2001 From: Caroline Chen <324939130+caroline-db@users.noreply.github.com> Date: Fri, 11 Sep 2026 19:07:25 +0000 Subject: [PATCH 1/2] Apply AIR permissions to submitted workloads --- .../experimental/air/run-submit/output.txt | 51 +++++- .../experimental/air/run-submit/run.yaml.tmpl | 7 + acceptance/experimental/air/run-submit/script | 4 +- .../experimental/air/run-submit/test.toml | 27 +++ experimental/air/cmd/runpermissions.go | 164 ++++++++++++++++++ experimental/air/cmd/runpermissions_test.go | 113 ++++++++++++ experimental/air/cmd/runsubmit.go | 2 + 7 files changed, 364 insertions(+), 4 deletions(-) create mode 100644 experimental/air/cmd/runpermissions.go create mode 100644 experimental/air/cmd/runpermissions_test.go diff --git a/acceptance/experimental/air/run-submit/output.txt b/acceptance/experimental/air/run-submit/output.txt index 5b54d70cc4d..a3a5a6acfb1 100644 --- a/acceptance/experimental/air/run-submit/output.txt +++ b/acceptance/experimental/air/run-submit/output.txt @@ -10,8 +10,55 @@ Tip: use --watch when submitting a run to stream logs to your terminal. Stream logs after submission using: databricks experimental air logs 555 -=== the ai_runtime_task carries the code_source_path ->>> print_requests.py //api/2.2/jobs/runs/submit +=== the ai_runtime_task and additive permission grants +>>> print_requests.py //api/2.2/jobs/runs/submit //api/2.0/mlflow/experiments/create //api/2.0/permissions --sort +{ + "method": "PATCH", + "path": "/api/2.0/permissions/experiments/exp-456", + "body": { + "access_control_list": [ + { + "permission_level": "CAN_MANAGE", + "user_name": "alice@example.com" + }, + { + "group_name": "data-team", + "permission_level": "CAN_READ" + }, + { + "permission_level": "CAN_EDIT", + "service_principal_name": "training-sp" + } + ] + } +} +{ + "method": "PATCH", + "path": "/api/2.0/permissions/jobs/123", + "body": { + "access_control_list": [ + { + "permission_level": "CAN_MANAGE", + "user_name": "alice@example.com" + }, + { + "group_name": "data-team", + "permission_level": "CAN_VIEW" + }, + { + "permission_level": "CAN_MANAGE_RUN", + "service_principal_name": "training-sp" + } + ] + } +} +{ + "method": "POST", + "path": "/api/2.0/mlflow/experiments/create", + "body": { + "name": "/Users/[USERNAME]/submit-smoke" + } +} { "method": "POST", "path": "/api/2.2/jobs/runs/submit", diff --git a/acceptance/experimental/air/run-submit/run.yaml.tmpl b/acceptance/experimental/air/run-submit/run.yaml.tmpl index 3fdbf48eb85..aab6a868e63 100644 --- a/acceptance/experimental/air/run-submit/run.yaml.tmpl +++ b/acceptance/experimental/air/run-submit/run.yaml.tmpl @@ -3,6 +3,13 @@ command: python train.py compute: accelerator_type: GPU_1xH100 num_accelerators: 1 +permissions: + - user_name: alice@example.com + level: CAN_MANAGE + - group_name: data-team + level: CAN_VIEW + - service_principal_name: training-sp + level: CAN_MANAGE_RUN code_source: type: snapshot snapshot: diff --git a/acceptance/experimental/air/run-submit/script b/acceptance/experimental/air/run-submit/script index 1f88a7f56d6..ae87b507984 100644 --- a/acceptance/experimental/air/run-submit/script +++ b/acceptance/experimental/air/run-submit/script @@ -12,7 +12,7 @@ sed "s/COMMIT_SHA/$(git rev-parse HEAD)/" run.yaml.tmpl > run.yaml title "submit with a git code_source" trace $CLI experimental air run -f run.yaml -title "the ai_runtime_task carries the code_source_path" -trace print_requests.py //api/2.2/jobs/runs/submit +title "the ai_runtime_task and additive permission grants" +trace print_requests.py //api/2.2/jobs/runs/submit //api/2.0/mlflow/experiments/create //api/2.0/permissions --sort rm -fr .git diff --git a/acceptance/experimental/air/run-submit/test.toml b/acceptance/experimental/air/run-submit/test.toml index 2e641379092..19238d1ea0b 100644 --- a/acceptance/experimental/air/run-submit/test.toml +++ b/acceptance/experimental/air/run-submit/test.toml @@ -21,6 +21,33 @@ Response.Body = ''' {"run_id": 555} ''' +[[Server]] +Pattern = "GET /api/2.0/mlflow/experiments/get-by-name" +Response.StatusCode = 404 +Response.Body = ''' +{"error_code":"RESOURCE_DOES_NOT_EXIST","message":"experiment does not exist"} +''' + +[[Server]] +Pattern = "POST /api/2.0/mlflow/experiments/create" +Response.Body = ''' +{"experiment_id":"exp-456"} +''' + +[[Server]] +Pattern = "GET /api/2.2/jobs/runs/get" +Response.Body = ''' +{"job_id":123,"run_id":555} +''' + +[[Server]] +Pattern = "PATCH /api/2.0/permissions/jobs/123" +Response.Body = '{}' + +[[Server]] +Pattern = "PATCH /api/2.0/permissions/experiments/exp-456" +Response.Body = '{}' + # The snapshot tarball is named _.tar.gz, where is the # test's temp-dir basename and the cache key derives from the pinned commit SHA. # Both are stable given the pinned commit dates in the script, but the temp-dir diff --git a/experimental/air/cmd/runpermissions.go b/experimental/air/cmd/runpermissions.go new file mode 100644 index 00000000000..1cd826696f0 --- /dev/null +++ b/experimental/air/cmd/runpermissions.go @@ -0,0 +1,164 @@ +package aircmd + +import ( + "context" + "errors" + "fmt" + "strconv" + "strings" + + "github.com/databricks/cli/libs/log" + "github.com/databricks/databricks-sdk-go" + "github.com/databricks/databricks-sdk-go/apierr" + "github.com/databricks/databricks-sdk-go/service/iam" + "github.com/databricks/databricks-sdk-go/service/jobs" + "github.com/databricks/databricks-sdk-go/service/ml" +) + +// permissionExperimentName returns the full MLflow experiment path for a run. +func permissionExperimentName(ctx context.Context, w *databricks.WorkspaceClient, cfg *runConfig) (string, error) { + if cfg.MLflowExperimentDirectory != nil { + return strings.TrimRight(*cfg.MLflowExperimentDirectory, "/") + "/" + cfg.ExperimentName, nil + } + + email, err := currentUserEmail(ctx, w) + if err != nil { + return "", err + } + return "/Users/" + email + "/" + cfg.ExperimentName, nil +} + +// getOrCreateMLflowExperiment resolves the experiment ID, creating it when absent. +func getOrCreateMLflowExperiment(ctx context.Context, w *databricks.WorkspaceClient, name, artifactLocation string) (string, error) { + existing, err := w.Experiments.GetByName(ctx, ml.GetByNameRequest{ExperimentName: name}) + if err == nil && existing.Experiment != nil && existing.Experiment.ExperimentId != "" { + return existing.Experiment.ExperimentId, nil + } + if err != nil && !errors.Is(err, apierr.ErrNotFound) { + return "", fmt.Errorf("failed to get MLflow experiment %q: %w", name, err) + } + + created, err := w.Experiments.CreateExperiment(ctx, ml.CreateExperiment{ + Name: name, + ArtifactLocation: artifactLocation, + }) + if err == nil { + return created.ExperimentId, nil + } + if !errors.Is(err, apierr.ErrAlreadyExists) && !errors.Is(err, apierr.ErrResourceAlreadyExists) { + return "", fmt.Errorf("failed to create MLflow experiment %q: %w", name, err) + } + + existing, err = w.Experiments.GetByName(ctx, ml.GetByNameRequest{ExperimentName: name}) + if err != nil { + return "", fmt.Errorf("failed to get concurrently created MLflow experiment %q: %w", name, err) + } + if existing.Experiment == nil || existing.Experiment.ExperimentId == "" { + return "", fmt.Errorf("MLflow experiment %q exists but has no experiment ID", name) + } + return existing.Experiment.ExperimentId, nil +} + +// experimentPermissionLevel maps a Jobs permission level to its MLflow equivalent. +func experimentPermissionLevel(level string) (iam.PermissionLevel, error) { + switch iam.PermissionLevel(level) { + case iam.PermissionLevelCanView: + return iam.PermissionLevelCanRead, nil + case iam.PermissionLevelCanManageRun: + return iam.PermissionLevelCanEdit, nil + case iam.PermissionLevelCanManage, iam.PermissionLevelIsOwner: + return iam.PermissionLevelCanManage, nil + default: + return "", fmt.Errorf("unsupported AIR permission level %q", level) + } +} + +// permissionAccessControl builds an ACL entry for a validated permission grant. +func permissionAccessControl(p permission, level iam.PermissionLevel) iam.AccessControlRequest { + acl := iam.AccessControlRequest{PermissionLevel: level} + switch { + case p.UserName != nil: + acl.UserName = *p.UserName + case p.GroupName != nil: + acl.GroupName = *p.GroupName + case p.ServicePrincipalName != nil: + acl.ServicePrincipalName = *p.ServicePrincipalName + } + return acl +} + +// grantWorkloadPermissions adds the configured ACLs to a job and its experiment. +func grantWorkloadPermissions(ctx context.Context, w *databricks.WorkspaceClient, jobID, experimentID string, permissions []permission) error { + if len(permissions) == 0 { + return nil + } + + jobACL := make([]iam.AccessControlRequest, 0, len(permissions)) + experimentACL := make([]iam.AccessControlRequest, 0, len(permissions)) + for _, p := range permissions { + experimentLevel, err := experimentPermissionLevel(p.Level) + if err != nil { + return err + } + jobACL = append(jobACL, permissionAccessControl(p, iam.PermissionLevel(p.Level))) + experimentACL = append(experimentACL, permissionAccessControl(p, experimentLevel)) + } + + _, err := w.Permissions.Update(ctx, iam.UpdateObjectPermissions{ + RequestObjectType: "jobs", + RequestObjectId: jobID, + AccessControlList: jobACL, + }) + if err != nil { + return fmt.Errorf("failed to grant job permissions: %w", err) + } + + _, err = w.Permissions.Update(ctx, iam.UpdateObjectPermissions{ + RequestObjectType: "experiments", + RequestObjectId: experimentID, + AccessControlList: experimentACL, + }) + if err != nil { + return fmt.Errorf("failed to grant MLflow experiment permissions: %w", err) + } + return nil +} + +// preparePermissionExperiment resolves the experiment before the workload is submitted. +func preparePermissionExperiment(ctx context.Context, w *databricks.WorkspaceClient, cfg *runConfig) string { + if len(cfg.Permissions) == 0 { + return "" + } + + name, err := permissionExperimentName(ctx, w, cfg) + if err != nil { + log.Warnf(ctx, "unable to resolve MLflow experiment name; skipping permission grants: %v", err) + return "" + } + artifactLocation := "" + if cfg.MLflowArtifactLocation != nil { + artifactLocation = *cfg.MLflowArtifactLocation + } + experimentID, err := getOrCreateMLflowExperiment(ctx, w, name, artifactLocation) + if err != nil { + log.Warnf(ctx, "unable to get or create MLflow experiment; skipping permission grants: %v", err) + return "" + } + return experimentID +} + +// applySubmittedPermissions resolves the submitted job and adds its configured ACLs. +func applySubmittedPermissions(ctx context.Context, w *databricks.WorkspaceClient, runID int64, experimentID string, permissions []permission) { + if len(permissions) == 0 || experimentID == "" { + return + } + + run, err := w.Jobs.GetRun(ctx, jobs.GetRunRequest{RunId: runID}) + if err == nil { + err = grantWorkloadPermissions(ctx, w, strconv.FormatInt(run.JobId, 10), experimentID, permissions) + } + if err != nil { + log.Warnf(ctx, "failed to grant permissions on workload: %v", err) + log.Warnf(ctx, "job was created successfully, but permissions could not be granted") + } +} diff --git a/experimental/air/cmd/runpermissions_test.go b/experimental/air/cmd/runpermissions_test.go new file mode 100644 index 00000000000..daa2a6e4144 --- /dev/null +++ b/experimental/air/cmd/runpermissions_test.go @@ -0,0 +1,113 @@ +package aircmd + +import ( + "encoding/json" + "net/http" + "testing" + + "github.com/databricks/cli/libs/testserver" + "github.com/databricks/databricks-sdk-go" + "github.com/databricks/databricks-sdk-go/service/iam" + "github.com/databricks/databricks-sdk-go/service/ml" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func TestGetOrCreateMLflowExperimentCreatesWithArtifactLocation(t *testing.T) { + server := testserver.New(t) + t.Cleanup(server.Close) + + server.Handle("GET", "/api/2.0/mlflow/experiments/get-by-name", func(req testserver.Request) any { + return testserver.Response{ + StatusCode: http.StatusNotFound, + Body: map[string]string{ + "error_code": "RESOURCE_DOES_NOT_EXIST", + "message": "experiment does not exist", + }, + } + }) + server.Handle("POST", "/api/2.0/mlflow/experiments/create", func(req testserver.Request) any { + var got ml.CreateExperiment + require.NoError(t, json.Unmarshal(req.Body, &got)) + assert.Equal(t, "/Users/alice@example.com/training", got.Name) + assert.Equal(t, "dbfs:/Volumes/main/default/artifacts", got.ArtifactLocation) + return ml.CreateExperimentResponse{ExperimentId: "exp-456"} + }) + + w, err := databricks.NewWorkspaceClient(&databricks.Config{Host: server.URL, Token: "token"}) + require.NoError(t, err) + experimentID, err := getOrCreateMLflowExperiment(t.Context(), w, "/Users/alice@example.com/training", "dbfs:/Volumes/main/default/artifacts") + require.NoError(t, err) + assert.Equal(t, "exp-456", experimentID) +} + +func TestExperimentPermissionLevel(t *testing.T) { + tests := []struct { + job string + experiment iam.PermissionLevel + }{ + {"CAN_VIEW", iam.PermissionLevelCanRead}, + {"CAN_MANAGE_RUN", iam.PermissionLevelCanEdit}, + {"CAN_MANAGE", iam.PermissionLevelCanManage}, + {"IS_OWNER", iam.PermissionLevelCanManage}, + } + for _, tt := range tests { + t.Run(tt.job, func(t *testing.T) { + got, err := experimentPermissionLevel(tt.job) + require.NoError(t, err) + assert.Equal(t, tt.experiment, got) + }) + } +} + +func TestGrantWorkloadPermissionsRejectsUnsupportedLevelBeforeRequests(t *testing.T) { + server := testserver.New(t) + t.Cleanup(server.Close) + + w, err := databricks.NewWorkspaceClient(&databricks.Config{Host: server.URL, Token: "token"}) + require.NoError(t, err) + server.RequestCallback = func(req *testserver.Request) { + t.Errorf("unexpected permission request: %s %s", req.Method, req.URL.Path) + } + err = grantWorkloadPermissions(t.Context(), w, "123", "exp-456", []permission{ + {GroupName: new("data-team"), Level: "CAN_USE"}, + }) + require.EqualError(t, err, `unsupported AIR permission level "CAN_USE"`) +} + +func TestGrantWorkloadPermissionsUpdatesJobAndExperiment(t *testing.T) { + server := testserver.New(t) + t.Cleanup(server.Close) + + requests := make(map[string]string) + for _, objectPath := range []string{"jobs/123", "experiments/exp-456"} { + server.Handle("PATCH", "/api/2.0/permissions/"+objectPath, func(req testserver.Request) any { + requests[objectPath] = string(req.Body) + return map[string]any{} + }) + } + + w, err := databricks.NewWorkspaceClient(&databricks.Config{Host: server.URL, Token: "token"}) + require.NoError(t, err) + err = grantWorkloadPermissions(t.Context(), w, "123", "exp-456", []permission{ + {UserName: new("alice@example.com"), Level: "CAN_MANAGE"}, + {GroupName: new("data-team"), Level: "CAN_VIEW"}, + {ServicePrincipalName: new("training-sp"), Level: "CAN_MANAGE_RUN"}, + }) + require.NoError(t, err) + + assert.JSONEq(t, `{ + "access_control_list": [ + {"user_name": "alice@example.com", "permission_level": "CAN_MANAGE"}, + {"group_name": "data-team", "permission_level": "CAN_VIEW"}, + {"service_principal_name": "training-sp", "permission_level": "CAN_MANAGE_RUN"} + ] + }`, requests["jobs/123"]) + assert.JSONEq(t, `{ + "access_control_list": [ + {"user_name": "alice@example.com", "permission_level": "CAN_MANAGE"}, + {"group_name": "data-team", "permission_level": "CAN_READ"}, + {"service_principal_name": "training-sp", "permission_level": "CAN_EDIT"} + ] + }`, requests["experiments/exp-456"]) +} diff --git a/experimental/air/cmd/runsubmit.go b/experimental/air/cmd/runsubmit.go index c580fec8c47..2474ed8a0ac 100644 --- a/experimental/air/cmd/runsubmit.go +++ b/experimental/air/cmd/runsubmit.go @@ -363,6 +363,7 @@ func submitWorkload(ctx context.Context, w *databricks.WorkspaceClient, cfg *run runtimeVersion, _ := cfg.runtimeVersion() payload := buildSubmitPayload(cfg, commandPath, dlRuntimeImage(ctx, runtimeVersion), usagePolicyID, snap, deps) payload.IdempotencyToken = token + experimentID := preparePermissionExperiment(ctx, w, cfg) provisionedCapacityID := "" if cfg.Compute.ProvisionedCapacityID != nil { @@ -377,6 +378,7 @@ func submitWorkload(ctx context.Context, w *databricks.WorkspaceClient, cfg *run if err != nil { return 0, "", err } + applySubmittedPermissions(ctx, w, runID, experimentID, cfg.Permissions) dashboardURL := strings.TrimRight(w.Config.Host, "/") + "/jobs/runs/" + strconv.FormatInt(runID, 10) return runID, dashboardURL, nil From 03b16debe2919fbffdc49034f91da6e59da48a33 Mon Sep 17 00:00:00 2001 From: Caroline Chen <324939130+caroline-db@users.noreply.github.com> Date: Wed, 16 Sep 2026 15:05:34 +0000 Subject: [PATCH 2/2] Apply AIR permissions only to jobs --- .../experimental/air/run-submit/output.txt | 32 +---- acceptance/experimental/air/run-submit/script | 5 +- .../experimental/air/run-submit/test.toml | 17 --- experimental/air/cmd/runpermissions.go | 116 ++---------------- experimental/air/cmd/runpermissions_test.go | 89 ++------------ experimental/air/cmd/runsubmit.go | 3 +- 6 files changed, 25 insertions(+), 237 deletions(-) diff --git a/acceptance/experimental/air/run-submit/output.txt b/acceptance/experimental/air/run-submit/output.txt index a3a5a6acfb1..c20c8467f1e 100644 --- a/acceptance/experimental/air/run-submit/output.txt +++ b/acceptance/experimental/air/run-submit/output.txt @@ -10,26 +10,13 @@ Tip: use --watch when submitting a run to stream logs to your terminal. Stream logs after submission using: databricks experimental air logs 555 -=== the ai_runtime_task and additive permission grants ->>> print_requests.py //api/2.2/jobs/runs/submit //api/2.0/mlflow/experiments/create //api/2.0/permissions --sort +=== the ai_runtime_task and additive job permission grants +>>> print_requests.py --get //api/2.2/jobs/runs/submit //api/2.2/jobs/runs/get //api/2.0/permissions //api/2.0/mlflow --sort --unique { - "method": "PATCH", - "path": "/api/2.0/permissions/experiments/exp-456", - "body": { - "access_control_list": [ - { - "permission_level": "CAN_MANAGE", - "user_name": "alice@example.com" - }, - { - "group_name": "data-team", - "permission_level": "CAN_READ" - }, - { - "permission_level": "CAN_EDIT", - "service_principal_name": "training-sp" - } - ] + "method": "GET", + "path": "/api/2.2/jobs/runs/get", + "q": { + "run_id": "555" } } { @@ -52,13 +39,6 @@ Stream logs after submission using: ] } } -{ - "method": "POST", - "path": "/api/2.0/mlflow/experiments/create", - "body": { - "name": "/Users/[USERNAME]/submit-smoke" - } -} { "method": "POST", "path": "/api/2.2/jobs/runs/submit", diff --git a/acceptance/experimental/air/run-submit/script b/acceptance/experimental/air/run-submit/script index ae87b507984..26161d63ad4 100644 --- a/acceptance/experimental/air/run-submit/script +++ b/acceptance/experimental/air/run-submit/script @@ -12,7 +12,8 @@ sed "s/COMMIT_SHA/$(git rev-parse HEAD)/" run.yaml.tmpl > run.yaml title "submit with a git code_source" trace $CLI experimental air run -f run.yaml -title "the ai_runtime_task and additive permission grants" -trace print_requests.py //api/2.2/jobs/runs/submit //api/2.0/mlflow/experiments/create //api/2.0/permissions --sort +title "the ai_runtime_task and additive job permission grants" +# Include MLflow paths so an accidental experiment request appears in the golden output. +trace print_requests.py --get //api/2.2/jobs/runs/submit //api/2.2/jobs/runs/get //api/2.0/permissions //api/2.0/mlflow --sort --unique rm -fr .git diff --git a/acceptance/experimental/air/run-submit/test.toml b/acceptance/experimental/air/run-submit/test.toml index 19238d1ea0b..8a55a09605c 100644 --- a/acceptance/experimental/air/run-submit/test.toml +++ b/acceptance/experimental/air/run-submit/test.toml @@ -21,19 +21,6 @@ Response.Body = ''' {"run_id": 555} ''' -[[Server]] -Pattern = "GET /api/2.0/mlflow/experiments/get-by-name" -Response.StatusCode = 404 -Response.Body = ''' -{"error_code":"RESOURCE_DOES_NOT_EXIST","message":"experiment does not exist"} -''' - -[[Server]] -Pattern = "POST /api/2.0/mlflow/experiments/create" -Response.Body = ''' -{"experiment_id":"exp-456"} -''' - [[Server]] Pattern = "GET /api/2.2/jobs/runs/get" Response.Body = ''' @@ -44,10 +31,6 @@ Response.Body = ''' Pattern = "PATCH /api/2.0/permissions/jobs/123" Response.Body = '{}' -[[Server]] -Pattern = "PATCH /api/2.0/permissions/experiments/exp-456" -Response.Body = '{}' - # The snapshot tarball is named _.tar.gz, where is the # test's temp-dir basename and the cache key derives from the pinned commit SHA. # Both are stable given the pinned commit dates in the script, but the temp-dir diff --git a/experimental/air/cmd/runpermissions.go b/experimental/air/cmd/runpermissions.go index 1cd826696f0..61ac4a0bc68 100644 --- a/experimental/air/cmd/runpermissions.go +++ b/experimental/air/cmd/runpermissions.go @@ -2,80 +2,18 @@ package aircmd import ( "context" - "errors" "fmt" "strconv" - "strings" "github.com/databricks/cli/libs/log" "github.com/databricks/databricks-sdk-go" - "github.com/databricks/databricks-sdk-go/apierr" "github.com/databricks/databricks-sdk-go/service/iam" "github.com/databricks/databricks-sdk-go/service/jobs" - "github.com/databricks/databricks-sdk-go/service/ml" ) -// permissionExperimentName returns the full MLflow experiment path for a run. -func permissionExperimentName(ctx context.Context, w *databricks.WorkspaceClient, cfg *runConfig) (string, error) { - if cfg.MLflowExperimentDirectory != nil { - return strings.TrimRight(*cfg.MLflowExperimentDirectory, "/") + "/" + cfg.ExperimentName, nil - } - - email, err := currentUserEmail(ctx, w) - if err != nil { - return "", err - } - return "/Users/" + email + "/" + cfg.ExperimentName, nil -} - -// getOrCreateMLflowExperiment resolves the experiment ID, creating it when absent. -func getOrCreateMLflowExperiment(ctx context.Context, w *databricks.WorkspaceClient, name, artifactLocation string) (string, error) { - existing, err := w.Experiments.GetByName(ctx, ml.GetByNameRequest{ExperimentName: name}) - if err == nil && existing.Experiment != nil && existing.Experiment.ExperimentId != "" { - return existing.Experiment.ExperimentId, nil - } - if err != nil && !errors.Is(err, apierr.ErrNotFound) { - return "", fmt.Errorf("failed to get MLflow experiment %q: %w", name, err) - } - - created, err := w.Experiments.CreateExperiment(ctx, ml.CreateExperiment{ - Name: name, - ArtifactLocation: artifactLocation, - }) - if err == nil { - return created.ExperimentId, nil - } - if !errors.Is(err, apierr.ErrAlreadyExists) && !errors.Is(err, apierr.ErrResourceAlreadyExists) { - return "", fmt.Errorf("failed to create MLflow experiment %q: %w", name, err) - } - - existing, err = w.Experiments.GetByName(ctx, ml.GetByNameRequest{ExperimentName: name}) - if err != nil { - return "", fmt.Errorf("failed to get concurrently created MLflow experiment %q: %w", name, err) - } - if existing.Experiment == nil || existing.Experiment.ExperimentId == "" { - return "", fmt.Errorf("MLflow experiment %q exists but has no experiment ID", name) - } - return existing.Experiment.ExperimentId, nil -} - -// experimentPermissionLevel maps a Jobs permission level to its MLflow equivalent. -func experimentPermissionLevel(level string) (iam.PermissionLevel, error) { - switch iam.PermissionLevel(level) { - case iam.PermissionLevelCanView: - return iam.PermissionLevelCanRead, nil - case iam.PermissionLevelCanManageRun: - return iam.PermissionLevelCanEdit, nil - case iam.PermissionLevelCanManage, iam.PermissionLevelIsOwner: - return iam.PermissionLevelCanManage, nil - default: - return "", fmt.Errorf("unsupported AIR permission level %q", level) - } -} - // permissionAccessControl builds an ACL entry for a validated permission grant. -func permissionAccessControl(p permission, level iam.PermissionLevel) iam.AccessControlRequest { - acl := iam.AccessControlRequest{PermissionLevel: level} +func permissionAccessControl(p permission) iam.AccessControlRequest { + acl := iam.AccessControlRequest{PermissionLevel: iam.PermissionLevel(p.Level)} switch { case p.UserName != nil: acl.UserName = *p.UserName @@ -87,21 +25,15 @@ func permissionAccessControl(p permission, level iam.PermissionLevel) iam.Access return acl } -// grantWorkloadPermissions adds the configured ACLs to a job and its experiment. -func grantWorkloadPermissions(ctx context.Context, w *databricks.WorkspaceClient, jobID, experimentID string, permissions []permission) error { +// grantJobPermissions adds the configured ACLs to a job. +func grantJobPermissions(ctx context.Context, w *databricks.WorkspaceClient, jobID string, permissions []permission) error { if len(permissions) == 0 { return nil } jobACL := make([]iam.AccessControlRequest, 0, len(permissions)) - experimentACL := make([]iam.AccessControlRequest, 0, len(permissions)) for _, p := range permissions { - experimentLevel, err := experimentPermissionLevel(p.Level) - if err != nil { - return err - } - jobACL = append(jobACL, permissionAccessControl(p, iam.PermissionLevel(p.Level))) - experimentACL = append(experimentACL, permissionAccessControl(p, experimentLevel)) + jobACL = append(jobACL, permissionAccessControl(p)) } _, err := w.Permissions.Update(ctx, iam.UpdateObjectPermissions{ @@ -112,50 +44,18 @@ func grantWorkloadPermissions(ctx context.Context, w *databricks.WorkspaceClient if err != nil { return fmt.Errorf("failed to grant job permissions: %w", err) } - - _, err = w.Permissions.Update(ctx, iam.UpdateObjectPermissions{ - RequestObjectType: "experiments", - RequestObjectId: experimentID, - AccessControlList: experimentACL, - }) - if err != nil { - return fmt.Errorf("failed to grant MLflow experiment permissions: %w", err) - } return nil } -// preparePermissionExperiment resolves the experiment before the workload is submitted. -func preparePermissionExperiment(ctx context.Context, w *databricks.WorkspaceClient, cfg *runConfig) string { - if len(cfg.Permissions) == 0 { - return "" - } - - name, err := permissionExperimentName(ctx, w, cfg) - if err != nil { - log.Warnf(ctx, "unable to resolve MLflow experiment name; skipping permission grants: %v", err) - return "" - } - artifactLocation := "" - if cfg.MLflowArtifactLocation != nil { - artifactLocation = *cfg.MLflowArtifactLocation - } - experimentID, err := getOrCreateMLflowExperiment(ctx, w, name, artifactLocation) - if err != nil { - log.Warnf(ctx, "unable to get or create MLflow experiment; skipping permission grants: %v", err) - return "" - } - return experimentID -} - // applySubmittedPermissions resolves the submitted job and adds its configured ACLs. -func applySubmittedPermissions(ctx context.Context, w *databricks.WorkspaceClient, runID int64, experimentID string, permissions []permission) { - if len(permissions) == 0 || experimentID == "" { +func applySubmittedPermissions(ctx context.Context, w *databricks.WorkspaceClient, runID int64, permissions []permission) { + if len(permissions) == 0 { return } run, err := w.Jobs.GetRun(ctx, jobs.GetRunRequest{RunId: runID}) if err == nil { - err = grantWorkloadPermissions(ctx, w, strconv.FormatInt(run.JobId, 10), experimentID, permissions) + err = grantJobPermissions(ctx, w, strconv.FormatInt(run.JobId, 10), permissions) } if err != nil { log.Warnf(ctx, "failed to grant permissions on workload: %v", err) diff --git a/experimental/air/cmd/runpermissions_test.go b/experimental/air/cmd/runpermissions_test.go index daa2a6e4144..a32568a2788 100644 --- a/experimental/air/cmd/runpermissions_test.go +++ b/experimental/air/cmd/runpermissions_test.go @@ -1,95 +1,27 @@ package aircmd import ( - "encoding/json" - "net/http" "testing" "github.com/databricks/cli/libs/testserver" "github.com/databricks/databricks-sdk-go" - "github.com/databricks/databricks-sdk-go/service/iam" - "github.com/databricks/databricks-sdk-go/service/ml" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" ) -func TestGetOrCreateMLflowExperimentCreatesWithArtifactLocation(t *testing.T) { +func TestGrantJobPermissions(t *testing.T) { server := testserver.New(t) t.Cleanup(server.Close) - server.Handle("GET", "/api/2.0/mlflow/experiments/get-by-name", func(req testserver.Request) any { - return testserver.Response{ - StatusCode: http.StatusNotFound, - Body: map[string]string{ - "error_code": "RESOURCE_DOES_NOT_EXIST", - "message": "experiment does not exist", - }, - } + var requestBody string + server.Handle("PATCH", "/api/2.0/permissions/jobs/123", func(req testserver.Request) any { + requestBody = string(req.Body) + return map[string]any{} }) - server.Handle("POST", "/api/2.0/mlflow/experiments/create", func(req testserver.Request) any { - var got ml.CreateExperiment - require.NoError(t, json.Unmarshal(req.Body, &got)) - assert.Equal(t, "/Users/alice@example.com/training", got.Name) - assert.Equal(t, "dbfs:/Volumes/main/default/artifacts", got.ArtifactLocation) - return ml.CreateExperimentResponse{ExperimentId: "exp-456"} - }) - - w, err := databricks.NewWorkspaceClient(&databricks.Config{Host: server.URL, Token: "token"}) - require.NoError(t, err) - experimentID, err := getOrCreateMLflowExperiment(t.Context(), w, "/Users/alice@example.com/training", "dbfs:/Volumes/main/default/artifacts") - require.NoError(t, err) - assert.Equal(t, "exp-456", experimentID) -} - -func TestExperimentPermissionLevel(t *testing.T) { - tests := []struct { - job string - experiment iam.PermissionLevel - }{ - {"CAN_VIEW", iam.PermissionLevelCanRead}, - {"CAN_MANAGE_RUN", iam.PermissionLevelCanEdit}, - {"CAN_MANAGE", iam.PermissionLevelCanManage}, - {"IS_OWNER", iam.PermissionLevelCanManage}, - } - for _, tt := range tests { - t.Run(tt.job, func(t *testing.T) { - got, err := experimentPermissionLevel(tt.job) - require.NoError(t, err) - assert.Equal(t, tt.experiment, got) - }) - } -} - -func TestGrantWorkloadPermissionsRejectsUnsupportedLevelBeforeRequests(t *testing.T) { - server := testserver.New(t) - t.Cleanup(server.Close) - - w, err := databricks.NewWorkspaceClient(&databricks.Config{Host: server.URL, Token: "token"}) - require.NoError(t, err) - server.RequestCallback = func(req *testserver.Request) { - t.Errorf("unexpected permission request: %s %s", req.Method, req.URL.Path) - } - err = grantWorkloadPermissions(t.Context(), w, "123", "exp-456", []permission{ - {GroupName: new("data-team"), Level: "CAN_USE"}, - }) - require.EqualError(t, err, `unsupported AIR permission level "CAN_USE"`) -} - -func TestGrantWorkloadPermissionsUpdatesJobAndExperiment(t *testing.T) { - server := testserver.New(t) - t.Cleanup(server.Close) - - requests := make(map[string]string) - for _, objectPath := range []string{"jobs/123", "experiments/exp-456"} { - server.Handle("PATCH", "/api/2.0/permissions/"+objectPath, func(req testserver.Request) any { - requests[objectPath] = string(req.Body) - return map[string]any{} - }) - } w, err := databricks.NewWorkspaceClient(&databricks.Config{Host: server.URL, Token: "token"}) require.NoError(t, err) - err = grantWorkloadPermissions(t.Context(), w, "123", "exp-456", []permission{ + err = grantJobPermissions(t.Context(), w, "123", []permission{ {UserName: new("alice@example.com"), Level: "CAN_MANAGE"}, {GroupName: new("data-team"), Level: "CAN_VIEW"}, {ServicePrincipalName: new("training-sp"), Level: "CAN_MANAGE_RUN"}, @@ -102,12 +34,5 @@ func TestGrantWorkloadPermissionsUpdatesJobAndExperiment(t *testing.T) { {"group_name": "data-team", "permission_level": "CAN_VIEW"}, {"service_principal_name": "training-sp", "permission_level": "CAN_MANAGE_RUN"} ] - }`, requests["jobs/123"]) - assert.JSONEq(t, `{ - "access_control_list": [ - {"user_name": "alice@example.com", "permission_level": "CAN_MANAGE"}, - {"group_name": "data-team", "permission_level": "CAN_READ"}, - {"service_principal_name": "training-sp", "permission_level": "CAN_EDIT"} - ] - }`, requests["experiments/exp-456"]) + }`, requestBody) } diff --git a/experimental/air/cmd/runsubmit.go b/experimental/air/cmd/runsubmit.go index 2474ed8a0ac..59fc184f00b 100644 --- a/experimental/air/cmd/runsubmit.go +++ b/experimental/air/cmd/runsubmit.go @@ -363,7 +363,6 @@ func submitWorkload(ctx context.Context, w *databricks.WorkspaceClient, cfg *run runtimeVersion, _ := cfg.runtimeVersion() payload := buildSubmitPayload(cfg, commandPath, dlRuntimeImage(ctx, runtimeVersion), usagePolicyID, snap, deps) payload.IdempotencyToken = token - experimentID := preparePermissionExperiment(ctx, w, cfg) provisionedCapacityID := "" if cfg.Compute.ProvisionedCapacityID != nil { @@ -378,7 +377,7 @@ func submitWorkload(ctx context.Context, w *databricks.WorkspaceClient, cfg *run if err != nil { return 0, "", err } - applySubmittedPermissions(ctx, w, runID, experimentID, cfg.Permissions) + applySubmittedPermissions(ctx, w, runID, cfg.Permissions) dashboardURL := strings.TrimRight(w.Config.Host, "/") + "/jobs/runs/" + strconv.FormatInt(runID, 10) return runID, dashboardURL, nil