diff --git a/.nextchanges/cli/u2m-client-id.md b/.nextchanges/cli/u2m-client-id.md new file mode 100644 index 00000000000..7221e2a850c --- /dev/null +++ b/.nextchanges/cli/u2m-client-id.md @@ -0,0 +1 @@ +* Allow OAuth U2M logins to override the CLI client ID with `--client-id`, profile `client_id`, or `DATABRICKS_CLIENT_ID`. ([#6594](https://github.com/databricks/cli/pull/6594)) diff --git a/acceptance/bin/discovery_browser.py b/acceptance/bin/discovery_browser.py index 42099fa06dd..96978edafc5 100755 --- a/acceptance/bin/discovery_browser.py +++ b/acceptance/bin/discovery_browser.py @@ -31,6 +31,11 @@ dest_parsed = urllib.parse.urlparse(destination_url) dest_params = urllib.parse.parse_qs(dest_parsed.query) +expected_client_id = os.environ.get("DATABRICKS_TEST_CLIENT_ID") +if expected_client_id and dest_params.get("client_id") != [expected_client_id]: + sys.stderr.write(f"Expected client_id {expected_client_id!r}, got {dest_params.get('client_id')!r}\n") + sys.exit(1) + redirect_uri = dest_params.get("redirect_uri", [None])[0] state = dest_params.get("state", [None])[0] diff --git a/acceptance/cmd/auth/login/discovery/out.databrickscfg b/acceptance/cmd/auth/login/discovery/out.databrickscfg index 56763df71ca..c1769c4cba9 100644 --- a/acceptance/cmd/auth/login/discovery/out.databrickscfg +++ b/acceptance/cmd/auth/login/discovery/out.databrickscfg @@ -5,6 +5,7 @@ host = [DATABRICKS_URL] account_id = test-account-123 workspace_id = [NUMID] +client_id = discovery-client-id auth_type = databricks-cli [__settings__] diff --git a/acceptance/cmd/auth/login/discovery/output.txt b/acceptance/cmd/auth/login/discovery/output.txt index c687b07fd5e..42e065b2b5b 100644 --- a/acceptance/cmd/auth/login/discovery/output.txt +++ b/acceptance/cmd/auth/login/discovery/output.txt @@ -1,5 +1,5 @@ ->>> [CLI] auth login --profile discovery-test +>>> [CLI] auth login --profile discovery-test --client-id discovery-client-id Opening login.databricks.com in your browser... Profile discovery-test was successfully saved diff --git a/acceptance/cmd/auth/login/discovery/script b/acceptance/cmd/auth/login/discovery/script index 4ae4c682d9e..43d3b92bc0a 100644 --- a/acceptance/cmd/auth/login/discovery/script +++ b/acceptance/cmd/auth/login/discovery/script @@ -2,8 +2,9 @@ sethome "./home" # Use the discovery browser script that simulates login.databricks.com export BROWSER="discovery_browser.py" +export DATABRICKS_TEST_CLIENT_ID="discovery-client-id" -trace $CLI auth login --profile discovery-test +trace $CLI auth login --profile discovery-test --client-id discovery-client-id trace $CLI auth profiles diff --git a/acceptance/cmd/auth/login/host-from-profile/home/.databrickscfg.tmpl b/acceptance/cmd/auth/login/host-from-profile/home/.databrickscfg.tmpl index 283a01f6d97..e3c1b1e0eae 100644 --- a/acceptance/cmd/auth/login/host-from-profile/home/.databrickscfg.tmpl +++ b/acceptance/cmd/auth/login/host-from-profile/home/.databrickscfg.tmpl @@ -1,2 +1,4 @@ [existing-profile] host = $DATABRICKS_HOST +auth_type = databricks-cli +client_id = custom-client-id diff --git a/acceptance/cmd/auth/login/host-from-profile/out.databrickscfg b/acceptance/cmd/auth/login/host-from-profile/out.databrickscfg index 0c13bde2570..e39250c3133 100644 --- a/acceptance/cmd/auth/login/host-from-profile/out.databrickscfg +++ b/acceptance/cmd/auth/login/host-from-profile/out.databrickscfg @@ -3,5 +3,6 @@ [existing-profile] host = [DATABRICKS_URL] -workspace_id = [NUMID] auth_type = databricks-cli +workspace_id = [NUMID] +client_id = flag-client-id diff --git a/acceptance/cmd/auth/login/host-from-profile/out.test.toml b/acceptance/cmd/auth/login/host-from-profile/out.test.toml index 98ea5040486..0938e678987 100644 --- a/acceptance/cmd/auth/login/host-from-profile/out.test.toml +++ b/acceptance/cmd/auth/login/host-from-profile/out.test.toml @@ -1,2 +1,2 @@ Cloud = false -EnvMatrix.DATABRICKS_BUNDLE_ENGINE = ["terraform", "direct"] +EnvMatrix.DATABRICKS_BUNDLE_ENGINE = ["direct"] diff --git a/acceptance/cmd/auth/login/host-from-profile/output.txt b/acceptance/cmd/auth/login/host-from-profile/output.txt index 6faae38ae5c..841ec08ade0 100644 --- a/acceptance/cmd/auth/login/host-from-profile/output.txt +++ b/acceptance/cmd/auth/login/host-from-profile/output.txt @@ -2,16 +2,34 @@ === Initial profile [existing-profile] host = [DATABRICKS_URL] +auth_type = databricks-cli +client_id = custom-client-id === Login with existing profile (no host argument) >>> [CLI] auth login --profile existing-profile Profile existing-profile was successfully saved +=== OAuth client ID + +>>> print_requests.py //oidc/v1/authorize --get +json.q.client_id = "custom-client-id"; + +=== Login with a client ID override + +>>> [CLI] auth login --profile existing-profile --client-id flag-client-id +Profile existing-profile was successfully saved + +=== Overridden OAuth client ID + +>>> print_requests.py //oidc/v1/authorize --get +json.q.client_id = "flag-client-id"; + === Profile after login ; The profile defined in the DEFAULT section is to be used as a fallback when no profile is explicitly specified. [DEFAULT] [existing-profile] host = [DATABRICKS_URL] -workspace_id = [NUMID] auth_type = databricks-cli +workspace_id = [NUMID] +client_id = flag-client-id diff --git a/acceptance/cmd/auth/login/host-from-profile/script b/acceptance/cmd/auth/login/host-from-profile/script index e9e0cefeeba..c800d378fea 100644 --- a/acceptance/cmd/auth/login/host-from-profile/script +++ b/acceptance/cmd/auth/login/host-from-profile/script @@ -14,6 +14,15 @@ export BROWSER="browser.py" title "Login with existing profile (no host argument)" trace $CLI auth login --profile existing-profile +title "OAuth client ID\n" +trace print_requests.py //oidc/v1/authorize --get | gron.py | grep client_id + +title "Login with a client ID override\n" +trace $CLI auth login --profile existing-profile --client-id flag-client-id + +title "Overridden OAuth client ID\n" +trace print_requests.py //oidc/v1/authorize --get | gron.py | grep client_id + title "Profile after login\n" cat "./home/.databrickscfg" diff --git a/acceptance/cmd/auth/login/host-from-profile/test.toml b/acceptance/cmd/auth/login/host-from-profile/test.toml index 36c0e7e237b..eda89653a0e 100644 --- a/acceptance/cmd/auth/login/host-from-profile/test.toml +++ b/acceptance/cmd/auth/login/host-from-profile/test.toml @@ -1,3 +1,5 @@ Ignore = [ "home" ] +RecordRequests = true +EnvMatrix.DATABRICKS_BUNDLE_ENGINE = ["direct"] diff --git a/cmd/auth/login.go b/cmd/auth/login.go index 5835297b335..3587dfafbdd 100644 --- a/cmd/auth/login.go +++ b/cmd/auth/login.go @@ -132,6 +132,7 @@ a new profile is created. var configureServerless bool var skipWorkspace bool var scopes string + var clientID string cmd.Flags().DurationVar(&loginTimeout, "timeout", defaultTimeout, "Timeout for completing login challenge in the browser") cmd.Flags().BoolVar(&configureCluster, "configure-cluster", false, @@ -142,6 +143,8 @@ a new profile is created. "Skip workspace selection for account-level access") cmd.Flags().StringVar(&scopes, "scopes", "", "Comma-separated list of OAuth scopes to request (defaults to 'all-apis')") + cmd.Flags().StringVar(&clientID, "client-id", "", + "OAuth client ID to use for U2M authentication") cmd.PreRunE = profileHostConflictCheck @@ -256,6 +259,9 @@ a new profile is created. if err != nil { return err } + if clientID == "" { + clientID = u2mClientIDFromProfile(existingProfile) + } // If no host is available from any source, use the discovery flow // via login.databricks.com. @@ -268,6 +274,7 @@ a new profile is created. profileName: profileName, timeout: loginTimeout, scopes: scopes, + clientID: clientID, existingProfile: existingProfile, browserFunc: getBrowserFunc(cmd), tokenStore: tokenStore, @@ -302,6 +309,9 @@ a new profile is created. u2m.WithBrowser(getBrowserFunc(cmd)), u2m.WithTokenCache(storage.WrapForOAuthArgument(ctx, tokenStore, mode, oauthArgument)), } + if clientID != "" { + persistentAuthOpts = append(persistentAuthOpts, u2m.WithClientID(clientID)) + } if len(scopesList) > 0 { persistentAuthOpts = append(persistentAuthOpts, u2m.WithScopes(scopesList)) } @@ -393,6 +403,7 @@ a new profile is created. ConfigFile: env.Get(ctx, "DATABRICKS_CONFIG_FILE"), ServerlessComputeID: serverlessComputeID, Scopes: scopesList, + ClientID: clientID, }, clearKeys...) if err != nil { return err @@ -634,6 +645,7 @@ type discoveryLoginInputs struct { profileName string timeout time.Duration scopes string + clientID string existingProfile *profile.Profile browserFunc func(string) error tokenStore storage.Store @@ -660,6 +672,9 @@ func discoveryLogin(ctx context.Context, in discoveryLoginInputs) error { u2m.WithDiscoveryLogin(), u2m.WithTokenCache(storage.WrapForOAuthArgument(ctx, in.tokenStore, in.mode, arg)), } + if in.clientID != "" { + opts = append(opts, u2m.WithClientID(in.clientID)) + } if len(scopesList) > 0 { opts = append(opts, u2m.WithScopes(scopesList)) } @@ -750,6 +765,7 @@ func discoveryLogin(ctx context.Context, in discoveryLoginInputs) error { WorkspaceID: workspaceID, Scopes: scopesList, ConfigFile: configFile, + ClientID: in.clientID, }, clearKeys...) if err != nil { if configFile != "" { @@ -785,6 +801,14 @@ func oauthLoginClearKeys() []string { return databrickscfg.AuthCredentialKeys() } +// u2mClientIDFromProfile excludes client IDs belonging to other auth types. +func u2mClientIDFromProfile(p *profile.Profile) string { + if p == nil || p.AuthType != authTypeDatabricksCLI { + return "" + } + return p.ClientID +} + // promptForWorkspaceSelection lists workspaces for a SPOG account and lets the // user pick one. Returns the selected workspace ID or empty string if skipped. // This is best-effort: errors are returned to the caller for logging, not shown diff --git a/cmd/auth/login_test.go b/cmd/auth/login_test.go index cd9e20bc234..0cde4addb2b 100644 --- a/cmd/auth/login_test.go +++ b/cmd/auth/login_test.go @@ -458,6 +458,35 @@ func TestSplitScopes(t *testing.T) { } } +func TestU2MClientIDFromProfile(t *testing.T) { + tests := []struct { + name string + profile *profile.Profile + want string + }{ + {name: "no profile"}, + { + name: "implicit auth type", + profile: &profile.Profile{ClientID: "custom-client-id"}, + }, + { + name: "M2M auth type", + profile: &profile.Profile{AuthType: "oauth-m2m", ClientID: "custom-client-id"}, + }, + { + name: "U2M auth type", + profile: &profile.Profile{AuthType: authTypeDatabricksCLI, ClientID: "custom-client-id"}, + want: "custom-client-id", + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + assert.Equal(t, tt.want, u2mClientIDFromProfile(tt.profile)) + }) + } +} + func TestRunHostDiscovery_NoHost(t *testing.T) { ctx := t.Context() args := &auth.AuthArguments{} @@ -952,9 +981,11 @@ func TestDiscoveryLogin_ReloginPreservesExistingProfileScopes(t *testing.T) { } existingProfile := &profile.Profile{ - Name: "DISCOVERY", - Host: "https://old-workspace.example.com", - Scopes: "sql,clusters", + Name: "DISCOVERY", + Host: "https://old-workspace.example.com", + Scopes: "sql,clusters", + AuthType: authTypeDatabricksCLI, + ClientID: "custom-client-id", } // No --scopes flag (empty string), should fall back to existing profile scopes. @@ -963,6 +994,7 @@ func TestDiscoveryLogin_ReloginPreservesExistingProfileScopes(t *testing.T) { dc: dc, profileName: "DISCOVERY", timeout: time.Second, + clientID: existingProfile.ClientID, existingProfile: existingProfile, browserFunc: func(string) error { return nil }, tokenStore: newTestStore(), @@ -974,6 +1006,51 @@ func TestDiscoveryLogin_ReloginPreservesExistingProfileScopes(t *testing.T) { require.NotNil(t, savedProfile) assert.Equal(t, "https://workspace.example.com", savedProfile.Host) assert.Equal(t, "sql,clusters", savedProfile.Scopes) + assert.Equal(t, "custom-client-id", savedProfile.ClientID) +} + +func TestDiscoveryLogin_ExplicitClientIDOverridesExistingProfile(t *testing.T) { + tmpDir := t.TempDir() + configPath := filepath.Join(tmpDir, ".databrickscfg") + err := os.WriteFile(configPath, []byte(""), 0o600) + require.NoError(t, err) + t.Setenv("DATABRICKS_CONFIG_FILE", configPath) + + oauthArg, err := u2m.NewBasicDiscoveryOAuthArgument("DISCOVERY") + require.NoError(t, err) + oauthArg.SetDiscoveredHost("https://workspace.example.com") + + dc := &fakeDiscoveryClient{ + oauthArg: oauthArg, + persistentAuth: &fakeDiscoveryPersistentAuth{ + token: &oauth2.Token{AccessToken: "test-token"}, + }, + introspectionErr: errors.New("introspection failed"), + } + + existingProfile := &profile.Profile{ + Name: "DISCOVERY", + Host: "https://old-workspace.example.com", + AuthType: authTypeDatabricksCLI, + ClientID: "profile-client-id", + } + + ctx, _ := cmdio.NewTestContextWithStdout(t.Context()) + err = discoveryLogin(ctx, discoveryLoginInputs{ + dc: dc, + profileName: "DISCOVERY", + timeout: time.Second, + clientID: "flag-client-id", + existingProfile: existingProfile, + browserFunc: func(string) error { return nil }, + tokenStore: newTestStore(), + }) + require.NoError(t, err) + + savedProfile, err := loadProfileByName(ctx, "DISCOVERY", profile.DefaultProfiler) + require.NoError(t, err) + require.NotNil(t, savedProfile) + assert.Equal(t, "flag-client-id", savedProfile.ClientID) } func TestDiscoveryLogin_ExplicitScopesOverrideExistingProfile(t *testing.T) { diff --git a/cmd/auth/token.go b/cmd/auth/token.go index dc8f54ecfc3..d6e9376caff 100644 --- a/cmd/auth/token.go +++ b/cmd/auth/token.go @@ -265,6 +265,9 @@ func loadToken(ctx context.Context, args loadTokenArgs) (*oauth2.Token, error) { return nil, err } allArgs := append([]u2m.PersistentAuthOption{u2m.WithTokenCache(storage.OAuthTokenCache(ctx, args.tokenStore, args.mode))}, args.persistentAuthOpts...) + if clientID := u2mClientIDFromProfile(existingProfile); clientID != "" { + allArgs = append(allArgs, u2m.WithClientID(clientID)) + } allArgs = append(allArgs, u2m.WithOAuthArgument(oauthArgument)) persistentAuth, err := u2m.NewPersistentAuth(ctx, allArgs...) if err != nil { @@ -430,6 +433,9 @@ func runInlineLogin(ctx context.Context, profiler profile.Profiler, tokenStore s u2m.WithBrowser(func(url string) error { return browser.Open(ctx, url) }), u2m.WithTokenCache(storage.WrapForOAuthArgument(ctx, tokenStore, mode, oauthArgument)), } + if clientID := u2mClientIDFromProfile(existingProfile); clientID != "" { + persistentAuthOpts = append(persistentAuthOpts, u2m.WithClientID(clientID)) + } if len(scopesList) > 0 { persistentAuthOpts = append(persistentAuthOpts, u2m.WithScopes(scopesList)) } @@ -458,6 +464,7 @@ func runInlineLogin(ctx context.Context, profiler profile.Profiler, tokenStore s WorkspaceID: loginArgs.WorkspaceID, ConfigFile: env.Get(ctx, "DATABRICKS_CONFIG_FILE"), Scopes: scopesList, + ClientID: u2mClientIDFromProfile(existingProfile), }, clearKeys...) if err != nil { return "", nil, err diff --git a/cmd/auth/token_test.go b/cmd/auth/token_test.go index 11a5a937bd9..d08bdd815fc 100644 --- a/cmd/auth/token_test.go +++ b/cmd/auth/token_test.go @@ -6,6 +6,7 @@ import ( "encoding/json" "errors" "net/http" + "net/url" "testing" "time" @@ -167,6 +168,13 @@ func TestToken_loadToken(t *testing.T) { Name: "valid-token", Host: "https://valid-token.cloud.databricks.com", }, + { + Name: "custom-client", + Host: "https://accounts.cloud.databricks.com", + AccountID: "custom-client", + AuthType: authTypeDatabricksCLI, + ClientID: "custom-client-id", + }, }, } tokenStore := &inMemoryStore{ @@ -218,6 +226,9 @@ func TestToken_loadToken(t *testing.T) { RefreshToken: "valid-token", Expiry: time.Now().Add(1 * time.Hour), }, + "custom-client": { + RefreshToken: "custom-refresh-token", + }, }, } validateToken := func(got *oauth2.Token) { @@ -324,6 +335,35 @@ func TestToken_loadToken(t *testing.T) { }, validateToken: validateToken, }, + { + name: "profile client ID is used for refresh", + args: loadTokenArgs{ + authArguments: &auth.AuthArguments{}, + profileName: "custom-client", + args: []string{}, + tokenTimeout: time.Hour, + profiler: profiler, + tokenStore: tokenStore, + persistentAuthOpts: []u2m.PersistentAuthOption{ + u2m.WithTokenCache(storage.ToU2MTokenCache(tokenStore)), + u2m.WithOAuthEndpointSupplier(&MockApiClient{}), + u2m.WithHttpClient(&http.Client{Transport: fixtures.SliceTransport{{ + MatchAny: true, + ExpectedRequest: url.Values{ + "client_id": {"custom-client-id"}, + "grant_type": {"refresh_token"}, + "refresh_token": {"custom-refresh-token"}, + }, + Response: map[string]string{ + "access_token": "new-access-token", + "token_type": "Bearer", + "expires_in": "3600", + }, + }}}), + }, + }, + validateToken: validateToken, + }, { name: "succeeds with host", args: loadTokenArgs{ diff --git a/libs/auth/credentials.go b/libs/auth/credentials.go index 10736a19bd1..75803593d6a 100644 --- a/libs/auth/credentials.go +++ b/libs/auth/credentials.go @@ -107,10 +107,14 @@ func (c CLICredentials) Configure(ctx context.Context, cfg *config.Config) (cred if err != nil { return nil, err } - ts, err := c.persistentAuth(ctx, + opts := []u2m.PersistentAuthOption{ u2m.WithOAuthArgument(oauthArg), u2m.WithTokenCache(storage.OAuthTokenCache(ctx, tokenStore, mode)), - ) + } + if cfg.AuthType == c.Name() && cfg.ClientID != "" { + opts = append(opts, u2m.WithClientID(cfg.ClientID)) + } + ts, err := c.persistentAuth(ctx, opts...) if err != nil { return nil, err } diff --git a/libs/auth/credentials_test.go b/libs/auth/credentials_test.go index 1c6fceac2d7..2717d7c3e48 100644 --- a/libs/auth/credentials_test.go +++ b/libs/auth/credentials_test.go @@ -232,6 +232,71 @@ func TestCLICredentialsConfigure_ThreadsResolvedTokenCache(t *testing.T) { assert.Len(t, receivedOpts, 2) } +func TestCLICredentialsConfigure_ClientID(t *testing.T) { + tests := []struct { + name string + authType string + source config.SourceType + wantOpts int + }{ + { + name: "U2M config file client ID", + authType: "databricks-cli", + source: config.SourceFile, + wantOpts: 3, + }, + { + name: "U2M environment client ID", + authType: "databricks-cli", + source: config.SourceEnv, + wantOpts: 3, + }, + { + name: "U2M dynamic client ID", + authType: "databricks-cli", + source: config.SourceDynamicConfig, + wantOpts: 3, + }, + { + name: "non-U2M config file client ID", + authType: "oauth-m2m", + source: config.SourceFile, + wantOpts: 2, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + hermeticAuthStorage(t) + cfg := &config.Config{ + Host: "https://workspace.test", + AuthType: tt.authType, + ClientID: "custom-client-id", + } + for i := range config.ConfigAttributes { + if config.ConfigAttributes[i].Name == "client_id" { + cfg.SetAttrSource(&config.ConfigAttributes[i], config.Source{Type: tt.source}) + break + } + } + + var receivedOpts []u2m.PersistentAuthOption + c := CLICredentials{ + persistentAuthFn: func(_ context.Context, opts ...u2m.PersistentAuthOption) (auth.TokenSource, error) { + receivedOpts = opts + return auth.TokenSourceFn(func(_ context.Context) (*oauth2.Token, error) { + return &oauth2.Token{AccessToken: "tok"}, nil + }), nil + }, + } + + _, err := c.Configure(t.Context(), cfg) + require.NoError(t, err) + assert.Len(t, receivedOpts, tt.wantOpts) + }) + } +} + // TestCLICredentialsConfigure_PropagatesStorageResolutionError confirms // Configure surfaces invalid DATABRICKS_AUTH_STORAGE values instead of // silently falling back to the file cache. If Configure ever stops calling diff --git a/libs/auth/u2m/discovery_token_source.go b/libs/auth/u2m/discovery_token_source.go index daa0ee118e4..a4342c053e8 100644 --- a/libs/auth/u2m/discovery_token_source.go +++ b/libs/auth/u2m/discovery_token_source.go @@ -59,7 +59,7 @@ const discoveryTargetAccount = "ACCOUNT" // the discovery OAuth flow. The OIDC authorize path with all OAuth query params // is URL-encoded as the destination_url parameter. func BuildDiscoveryAuthorizeURL(redirectAddr, state string, pkce PKCEParams, scopes []string) string { - return buildDiscoveryAuthorizeURL(defaultLoginDatabricksHost, redirectAddr, state, pkce, scopes, "") + return buildDiscoveryAuthorizeURL(defaultLoginDatabricksHost, redirectAddr, state, pkce, scopes, appClientID, "") } // buildDiscoveryAuthorizeURL builds the discovery authorize URL against the @@ -68,10 +68,10 @@ func BuildDiscoveryAuthorizeURL(redirectAddr, state string, pkce PKCEParams, sco // non-empty it is set as the top-level `target` query parameter, which // login.databricks.com uses to route the user to a specific selector page // (e.g. "ACCOUNT" for the account selector). -func buildDiscoveryAuthorizeURL(host, redirectAddr, state string, pkce PKCEParams, scopes []string, target string) string { +func buildDiscoveryAuthorizeURL(host, redirectAddr, state string, pkce PKCEParams, scopes []string, clientID, target string) string { // Build the nested OIDC authorize path with query parameters. authParams := url.Values{} - authParams.Set("client_id", appClientID) + authParams.Set("client_id", clientID) authParams.Set("redirect_uri", "http://"+redirectAddr) authParams.Set("response_type", "code") authParams.Set("scope", strings.Join(scopes, " ")) @@ -139,7 +139,7 @@ func (d *discoveryTokenSource) challenge() error { if host == "" { host = defaultLoginDatabricksHost } - authorizeURL := buildDiscoveryAuthorizeURL(host, d.pa.redirectAddr, state, pkce, scopes, d.target) + authorizeURL := buildDiscoveryAuthorizeURL(host, d.pa.redirectAddr, state, pkce, scopes, d.pa.clientID, d.target) code, returnedState, issuer, err := cb.handlerWithIssuer(authorizeURL) if err != nil { @@ -164,7 +164,7 @@ func (d *discoveryTokenSource) challenge() error { // Exchange authorization code for tokens. cfg := &oauth2.Config{ - ClientID: appClientID, + ClientID: d.pa.clientID, Endpoint: oauth2.Endpoint{ TokenURL: tokenEndpoint, AuthStyle: oauth2.AuthStyleInParams, diff --git a/libs/auth/u2m/discovery_token_source_test.go b/libs/auth/u2m/discovery_token_source_test.go index 8941f909e82..1ee8d3f6ff0 100644 --- a/libs/auth/u2m/discovery_token_source_test.go +++ b/libs/auth/u2m/discovery_token_source_test.go @@ -186,7 +186,7 @@ func TestBuildDiscoveryAuthorizeURL_HostOverride(t *testing.T) { } for _, tc := range tests { t.Run(tc.name, func(t *testing.T) { - got := buildDiscoveryAuthorizeURL(tc.host, "localhost:8020", "s", pkce, scopes, "") + got := buildDiscoveryAuthorizeURL(tc.host, "localhost:8020", "s", pkce, scopes, appClientID, "") u, err := url.Parse(got) if err != nil { t.Fatalf("parsing URL: %v", err) @@ -215,7 +215,7 @@ func TestBuildDiscoveryAuthorizeURL_Target(t *testing.T) { } for _, tc := range tests { t.Run(tc.name, func(t *testing.T) { - got := buildDiscoveryAuthorizeURL(defaultLoginDatabricksHost, "localhost:8020", "s", pkce, scopes, tc.target) + got := buildDiscoveryAuthorizeURL(defaultLoginDatabricksHost, "localhost:8020", "s", pkce, scopes, appClientID, tc.target) u, err := url.Parse(got) if err != nil { t.Fatalf("parsing URL: %v", err) @@ -273,6 +273,12 @@ func TestDiscoveryTokenSource_Challenge(t *testing.T) { if r.URL.Path != "/oidc/v1/token" { t.Errorf("token server: want path /oidc/v1/token, got %s", r.URL.Path) } + if err := r.ParseForm(); err != nil { + t.Fatalf("token server: parsing form: %v", err) + } + if got := r.Form.Get("client_id"); got != "custom-client-id" { + t.Errorf("token server: client_id = %q, want %q", got, "custom-client-id") + } w.Header().Set("Content-Type", "application/json") fmt.Fprintf(w, `{"access_token":"test-access-token","refresh_token":"test-refresh-token","token_type":"Bearer","expires_in":3600}`) })) @@ -308,6 +314,7 @@ func TestDiscoveryTokenSource_Challenge(t *testing.T) { WithOAuthEndpointSupplier(MockOAuthEndpointSupplier{}), WithOAuthArgument(arg), WithDiscoveryLogin(), + WithClientID("custom-client-id"), ) if err != nil { t.Fatalf("NewPersistentAuth(): %v", err) @@ -340,6 +347,9 @@ func TestDiscoveryTokenSource_Challenge(t *testing.T) { if err != nil { t.Fatalf("parsing destination_url: %v", err) } + if got := dest.Query().Get("client_id"); got != "custom-client-id" { + t.Errorf("authorize URL: client_id = %q, want %q", got, "custom-client-id") + } state = dest.Query().Get("state") if state == "" { t.Fatal("state is empty in authorize URL") diff --git a/libs/auth/u2m/persistent_auth.go b/libs/auth/u2m/persistent_auth.go index 8a6756ee342..7cc09afe039 100644 --- a/libs/auth/u2m/persistent_auth.go +++ b/libs/auth/u2m/persistent_auth.go @@ -69,6 +69,8 @@ var ( // The PersistentAuth is safe for concurrent use. The token cache is locked // during token retrieval, refresh and storage. type PersistentAuth struct { + clientID string + // cache is the token cache to store and lookup tokens. cache cache.TokenCache @@ -162,6 +164,13 @@ func WithOAuthArgument(arg OAuthArgument) PersistentAuthOption { } } +// WithClientID sets the OAuth client ID for the PersistentAuth. +func WithClientID(clientID string) PersistentAuthOption { + return func(a *PersistentAuth) { + a.clientID = clientID + } +} + // WithBrowser sets the browser function for the PersistentAuth. func WithBrowser(b func(url string) error) PersistentAuthOption { return func(a *PersistentAuth) { @@ -234,7 +243,9 @@ func WithDiscoveryAccountTarget() PersistentAuthOption { // NewPersistentAuth creates a new PersistentAuth with the provided options. func NewPersistentAuth(ctx context.Context, opts ...PersistentAuthOption) (*PersistentAuth, error) { - p := &PersistentAuth{} + p := &PersistentAuth{ + clientID: appClientID, // defaults to databricks-cli + } for _, opt := range opts { opt(p) } @@ -603,7 +614,8 @@ func (a *PersistentAuth) oauth2Config() (*oauth2.Config, error) { endpoints, err = a.endpointSupplier.GetWorkspaceOAuthEndpoints(a.ctx, argg.GetWorkspaceHost()) case AccountOAuthArgument: endpoints, err = a.endpointSupplier.GetAccountOAuthEndpoints( - a.ctx, argg.GetAccountHost(), argg.GetAccountId()) + a.ctx, argg.GetAccountHost(), argg.GetAccountId(), + ) case UnifiedOAuthArgument: endpoints, err = a.endpointSupplier.GetUnifiedOAuthEndpoints(a.ctx, argg.GetHost(), argg.GetAccountId()) case DiscoveryOAuthArgument: @@ -615,7 +627,7 @@ func (a *PersistentAuth) oauth2Config() (*oauth2.Config, error) { return nil, fmt.Errorf("fetching OAuth endpoints: %w", err) } return &oauth2.Config{ - ClientID: appClientID, + ClientID: a.clientID, Endpoint: oauth2.Endpoint{ AuthURL: endpoints.AuthorizationEndpoint, TokenURL: endpoints.TokenEndpoint, diff --git a/libs/auth/u2m/persistent_auth_test.go b/libs/auth/u2m/persistent_auth_test.go index 31c95913869..50d825a9d78 100644 --- a/libs/auth/u2m/persistent_auth_test.go +++ b/libs/auth/u2m/persistent_auth_test.go @@ -123,6 +123,48 @@ func (m MockOAuthEndpointSupplier) GetEndpointsFromURL(_ context.Context, _ stri return nil, ErrOAuthNotSupported } +func TestPersistentAuthClientID(t *testing.T) { + tests := []struct { + name string + opts []PersistentAuthOption + want string + }{ + { + name: "default", + want: appClientID, + }, + { + name: "custom", + opts: []PersistentAuthOption{WithClientID("custom-client-id")}, + want: "custom-client-id", + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + arg, err := NewBasicWorkspaceOAuthArgument("https://workspace.test") + if err != nil { + t.Fatalf("NewBasicWorkspaceOAuthArgument(): %v", err) + } + opts := append([]PersistentAuthOption{ + WithOAuthArgument(arg), + WithOAuthEndpointSupplier(MockOAuthEndpointSupplier{}), + }, tt.opts...) + p, err := NewPersistentAuth(t.Context(), opts...) + if err != nil { + t.Fatalf("NewPersistentAuth(): %v", err) + } + cfg, err := p.oauth2Config() + if err != nil { + t.Fatalf("oauth2Config(): %v", err) + } + if cfg.ClientID != tt.want { + t.Errorf("client ID = %q, want %q", cfg.ClientID, tt.want) + } + }) + } +} + func TestToken_RefreshesExpiredAccessToken(t *testing.T) { ctx := t.Context() expectedKey := "https://accounts.cloud.databricks.test/oidc/accounts/xyz" @@ -157,9 +199,10 @@ func TestToken_RefreshesExpiredAccessToken(t *testing.T) { WithHttpClient(&http.Client{ Transport: fixtures.SliceTransport{ { - Method: "POST", - Resource: "/oidc/accounts/xyz/v1/token", - Response: `access_token=refreshed&refresh_token=def`, + Method: "POST", + Resource: "/oidc/accounts/xyz/v1/token", + ExpectedRequest: url.Values{"client_id": {"custom-client-id"}, "grant_type": {"refresh_token"}, "refresh_token": {"cde"}}, + Response: `access_token=refreshed&refresh_token=def`, ResponseHeaders: map[string][]string{ "Content-Type": {"application/x-www-form-urlencoded"}, }, @@ -168,6 +211,7 @@ func TestToken_RefreshesExpiredAccessToken(t *testing.T) { }), WithOAuthEndpointSupplier(MockOAuthEndpointSupplier{}), WithOAuthArgument(arg), + WithClientID("custom-client-id"), ) if err != nil { t.Errorf("NewPersistentAuth(): want no error, got %v", err) @@ -967,8 +1011,10 @@ func TestChallenge(t *testing.T) { if u.Path != "/oidc/accounts/xyz/v1/authorize" { t.Fatalf("browser(): want path '/oidc/accounts/xyz/v1/authorize', got %s", u.Path) } - // for now we're ignoring asserting the fields of the redirect query := u.Query() + if query.Get("client_id") != "custom-client-id" { + t.Fatalf("browser(): client_id = %q, want %q", query.Get("client_id"), "custom-client-id") + } browserOpened <- query.Get("state") return nil } @@ -1009,6 +1055,7 @@ func TestChallenge(t *testing.T) { }), WithOAuthEndpointSupplier(MockOAuthEndpointSupplier{}), WithOAuthArgument(arg), + WithClientID("custom-client-id"), ) if err != nil { t.Errorf("NewPersistentAuth(): want no error, got %v", err) diff --git a/libs/databrickscfg/profile/file.go b/libs/databrickscfg/profile/file.go index b7f6074c811..529b2a0c34b 100644 --- a/libs/databrickscfg/profile/file.go +++ b/libs/databrickscfg/profile/file.go @@ -86,6 +86,7 @@ func (f FileProfilerImpl) LoadProfiles(ctx context.Context, fn ProfileMatchFunct ClusterID: all["cluster_id"], ServerlessComputeID: all["serverless_compute_id"], HasClientCredentials: all["client_id"] != "" && all["client_secret"] != "", + ClientID: all["client_id"], Scopes: all["scopes"], AuthType: all["auth_type"], } diff --git a/libs/databrickscfg/profile/file_test.go b/libs/databrickscfg/profile/file_test.go index 8f6c5ad790c..a59b2202693 100644 --- a/libs/databrickscfg/profile/file_test.go +++ b/libs/databrickscfg/profile/file_test.go @@ -1,6 +1,7 @@ package profile import ( + "os" "path/filepath" "testing" @@ -64,6 +65,23 @@ func TestLoadProfilesNoConfiguration(t *testing.T) { require.ErrorIs(t, err, ErrNoConfiguration) } +func TestLoadProfilesClientID(t *testing.T) { + configPath := filepath.Join(t.TempDir(), ".databrickscfg") + err := os.WriteFile(configPath, []byte(`[u2m] +host = https://workspace.test +auth_type = databricks-cli +client_id = custom-client-id +`), 0o600) + require.NoError(t, err) + + ctx := env.Set(t.Context(), "DATABRICKS_CONFIG_FILE", configPath) + profiles, err := (FileProfilerImpl{}).LoadProfiles(ctx, MatchAllProfiles) + require.NoError(t, err) + require.Len(t, profiles, 1) + assert.Equal(t, "custom-client-id", profiles[0].ClientID) + assert.False(t, profiles[0].HasClientCredentials) +} + func TestLoadProfilesMatchWorkspace(t *testing.T) { ctx := t.Context() ctx = env.Set(ctx, "DATABRICKS_CONFIG_FILE", "./testdata/databrickscfg") diff --git a/libs/databrickscfg/profile/profile.go b/libs/databrickscfg/profile/profile.go index efd358cd4e5..6a2ff7b5e1e 100644 --- a/libs/databrickscfg/profile/profile.go +++ b/libs/databrickscfg/profile/profile.go @@ -17,6 +17,7 @@ type Profile struct { ClusterID string ServerlessComputeID string HasClientCredentials bool + ClientID string Scopes string AuthType string }