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