From e5e2694d6284c96f2a510be8acb88cc61c89efb9 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] 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 6412af3cafb..d697ccba651 100644 --- a/experimental/air/cmd/runsubmit.go +++ b/experimental/air/cmd/runsubmit.go @@ -343,6 +343,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 { @@ -357,6 +358,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