diff --git a/acceptance/experimental/air/run-submit/output.txt b/acceptance/experimental/air/run-submit/output.txt index 5b54d70cc4d..c20c8467f1e 100644 --- a/acceptance/experimental/air/run-submit/output.txt +++ b/acceptance/experimental/air/run-submit/output.txt @@ -10,8 +10,35 @@ 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 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": "GET", + "path": "/api/2.2/jobs/runs/get", + "q": { + "run_id": "555" + } +} +{ + "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.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..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 carries the code_source_path" -trace print_requests.py //api/2.2/jobs/runs/submit +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 2e641379092..8a55a09605c 100644 --- a/acceptance/experimental/air/run-submit/test.toml +++ b/acceptance/experimental/air/run-submit/test.toml @@ -21,6 +21,16 @@ Response.Body = ''' {"run_id": 555} ''' +[[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 = '{}' + # 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..61ac4a0bc68 --- /dev/null +++ b/experimental/air/cmd/runpermissions.go @@ -0,0 +1,64 @@ +package aircmd + +import ( + "context" + "fmt" + "strconv" + + "github.com/databricks/cli/libs/log" + "github.com/databricks/databricks-sdk-go" + "github.com/databricks/databricks-sdk-go/service/iam" + "github.com/databricks/databricks-sdk-go/service/jobs" +) + +// permissionAccessControl builds an ACL entry for a validated permission grant. +func permissionAccessControl(p permission) iam.AccessControlRequest { + acl := iam.AccessControlRequest{PermissionLevel: iam.PermissionLevel(p.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 +} + +// 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)) + for _, p := range permissions { + jobACL = append(jobACL, permissionAccessControl(p)) + } + + _, 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) + } + return nil +} + +// applySubmittedPermissions resolves the submitted job and adds its configured ACLs. +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 = grantJobPermissions(ctx, w, strconv.FormatInt(run.JobId, 10), 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..a32568a2788 --- /dev/null +++ b/experimental/air/cmd/runpermissions_test.go @@ -0,0 +1,38 @@ +package aircmd + +import ( + "testing" + + "github.com/databricks/cli/libs/testserver" + "github.com/databricks/databricks-sdk-go" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func TestGrantJobPermissions(t *testing.T) { + server := testserver.New(t) + t.Cleanup(server.Close) + + 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{} + }) + + w, err := databricks.NewWorkspaceClient(&databricks.Config{Host: server.URL, Token: "token"}) + require.NoError(t, err) + 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"}, + }) + 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"} + ] + }`, requestBody) +} diff --git a/experimental/air/cmd/runsubmit.go b/experimental/air/cmd/runsubmit.go index c580fec8c47..59fc184f00b 100644 --- a/experimental/air/cmd/runsubmit.go +++ b/experimental/air/cmd/runsubmit.go @@ -377,6 +377,7 @@ func submitWorkload(ctx context.Context, w *databricks.WorkspaceClient, cfg *run if err != nil { return 0, "", err } + applySubmittedPermissions(ctx, w, runID, cfg.Permissions) dashboardURL := strings.TrimRight(w.Config.Host, "/") + "/jobs/runs/" + strconv.FormatInt(runID, 10) return runID, dashboardURL, nil