diff --git a/.nextchanges/cli/spog-provisioned-url.md b/.nextchanges/cli/spog-provisioned-url.md new file mode 100644 index 00000000000..892aed8ac9f --- /dev/null +++ b/.nextchanges/cli/spog-provisioned-url.md @@ -0,0 +1 @@ +* `databricks auth login` now uses your account's primary (SPOG) URL. With `--host ` and a workspace (`?o=` or `--workspace-id`), the profile keeps a workspace-scoped token, the same token a login to the workspace's own host gets, so `--resource` works there too. In the browser-based flow, picking an account or a workspace saves the profile with the primary URL. Workspace-scoped SPOG profiles record the workspace's OAuth metadata URL in `discovery_url` and need this CLI version or newer: older versions can't refresh their token and ask you to run `databricks auth login` again. ([#6854](https://github.com/databricks/cli/pull/6854)) diff --git a/cmd/auth/login.go b/cmd/auth/login.go index 2c60557a809..2516ffcb983 100644 --- a/cmd/auth/login.go +++ b/cmd/auth/login.go @@ -4,6 +4,8 @@ import ( "context" "errors" "fmt" + "net/http" + "net/url" "runtime" "strconv" "strings" @@ -43,8 +45,11 @@ func promptForProfile(ctx context.Context, defaultValue string) (string, error) const ( minimalDbConnectVersion = "13.1" defaultTimeout = 1 * time.Hour - authTypeDatabricksCLI = "databricks-cli" - discoveryFallbackTip = "\n\nTip: you can specify a workspace directly with: databricks auth login --host " + // provisionedURLTimeout bounds the best-effort SPOG host lookup so a + // hung endpoint can't stall login for the full login timeout. + provisionedURLTimeout = 30 * time.Second + authTypeDatabricksCLI = "databricks-cli" + discoveryFallbackTip = "\n\nTip: you can specify a workspace directly with: databricks auth login --host " // discoveryHostEnvVar overrides the default https://login.databricks.com // host used by the discovery login flow. Intended for testing and // development against non-production environments. @@ -305,7 +310,7 @@ a new profile is created. }) } - err = setHostAndAccountId(ctx, existingProfile, authArguments, args) + err = setLoginHostAndAccountId(ctx, existingProfile, authArguments, args) if err != nil { return err } @@ -382,6 +387,9 @@ a new profile is created. // from .well-known discovery, so stale values would be misleading). clearKeys := oauthLoginClearKeys() clearKeys = append(clearKeys, databrickscfg.ExperimentalIsUnifiedHostKey) + if existingProfile != nil && auth.IsSpogWorkspaceDiscoveryURL(existingProfile.DiscoveryURL) { + clearKeys = append(clearKeys, databrickscfg.DiscoveryURLKey) + } switch { case configureCluster: @@ -427,6 +435,7 @@ a new profile is created. ServerlessComputeID: serverlessComputeID, Scopes: scopesList, ClientID: clientID, + DiscoveryURL: spogWorkspaceDiscoveryURLToSave(authArguments.DiscoveryURL), }, clearKeys...) if err != nil { return err @@ -553,9 +562,54 @@ func setHostAndAccountId(ctx context.Context, existingProfile *profile.Profile, } } + useProfileSpogWorkspaceOAuth(authArguments, existingProfile) + + return nil +} + +// setLoginHostAndAccountId is setHostAndAccountId for auth login. A workspace +// named for this login (--workspace-id, or ?o= on the host flag or argument) +// on a SPOG host also gets a workspace-scoped token. Other commands only follow +// the token type of the existing profile, since its stored token decides which +// OAuth endpoints can refresh it. +func setLoginHostAndAccountId(ctx context.Context, existingProfile *profile.Profile, authArguments *auth.AuthArguments, args []string) error { + workspaceID := workspaceNamedForLogin(authArguments, args) + if err := setHostAndAccountId(ctx, existingProfile, authArguments, args); err != nil { + return err + } + useSpogWorkspaceOAuth(ctx, authArguments, workspaceID) return nil } +// workspaceNamedForLogin returns the workspace ID passed to this login with +// --workspace-id or as ?o= on the host flag or argument. A workspace_id +// inherited from an existing profile doesn't count, so re-login keeps the +// profile's token type. +func workspaceNamedForLogin(authArguments *auth.AuthArguments, args []string) string { + if authArguments.WorkspaceID != "" { + return authArguments.WorkspaceID + } + host := authArguments.Host + if host == "" && len(args) > 0 { + host = args[0] + } + return auth.ExtractHostQueryParams(host).WorkspaceID +} + +// useProfileSpogWorkspaceOAuth keeps an existing profile's workspace-scoped +// SPOG token type while authArguments still target the profile's host and +// workspace. +func useProfileSpogWorkspaceOAuth(authArguments *auth.AuthArguments, existingProfile *profile.Profile) { + if existingProfile == nil || !auth.IsSpogWorkspaceDiscoveryURL(existingProfile.DiscoveryURL) || + existingProfile.WorkspaceID != authArguments.WorkspaceID { + return + } + if (&config.Config{Host: existingProfile.Host}).CanonicalHostName() != (&config.Config{Host: authArguments.Host}).CanonicalHostName() { + return + } + authArguments.DiscoveryURL = existingProfile.DiscoveryURL +} + // needsAccountIDPrompt reports whether the target host requires an account ID // for OAuth URL construction. True for classic account hosts (accounts.*) and // for unified hosts detected via account-scoped DiscoveryURL. @@ -571,15 +625,23 @@ func needsAccountIDPrompt(host, discoveryURL string) bool { // .well-known/databricks-config from the host. Populates account_id and // workspace_id from discovery if not already set. func runHostDiscovery(ctx context.Context, authArguments *auth.AuthArguments) { + runHostDiscoveryWithRetryTimeout(ctx, authArguments, 0) +} + +// runHostDiscoveryWithRetryTimeout is runHostDiscovery with the SDK's total +// retry budget capped at retryTimeoutSeconds (0 keeps the SDK default). +// EnsureResolved doesn't take a context, so this is the only way to bound it. +func runHostDiscoveryWithRetryTimeout(ctx context.Context, authArguments *auth.AuthArguments, retryTimeoutSeconds int) { if authArguments.Host == "" { return } cfg := &config.Config{ - Host: authArguments.Host, - AccountID: authArguments.AccountID, - WorkspaceID: authArguments.WorkspaceID, - HTTPTimeoutSeconds: 5, + Host: authArguments.Host, + AccountID: authArguments.AccountID, + WorkspaceID: authArguments.WorkspaceID, + HTTPTimeoutSeconds: 5, + RetryTimeoutSeconds: retryTimeoutSeconds, // Use only ConfigAttributes (env vars + struct tags), skip config file // loading to avoid interference from existing profiles. Loaders: []config.Loader{config.ConfigAttributes}, @@ -676,6 +738,164 @@ func validateDiscoveryFlagCompatibility(cmd *cobra.Command) error { return nil } +// shouldResolveProvisionedURL reports whether to look up the account's primary +// provisioned (SPOG) URL for the given host and account. It is true only for a +// classic account host with an account ID: an account ID can also be present on +// a concrete workspace host (via --account-id, ?a=, or token introspection), +// and rewriting that host to the account SPOG URL would discard the user's +// targeted workspace. +func shouldResolveProvisionedURL(host, accountID string) bool { + return accountID != "" && auth.IsClassicAccountHost((&config.Config{Host: host}).CanonicalHostName()) +} + +// resolvePrimaryProvisionedURL returns the account's primary provisioned URL +// (its SPOG host) for the given account, or host unchanged when the account has +// no provisioned URL or the lookup fails. The lookup is bounded by +// provisionedURLTimeout so a hung endpoint can't stall login, and is +// best-effort: failures are logged and never block login. +func resolvePrimaryProvisionedURL(ctx context.Context, host, accountID, accessToken string, httpClient *http.Client) string { + lookupCtx, cancel := context.WithTimeout(ctx, provisionedURLTimeout) + defer cancel() + spogURL, err := auth.LookupPrimaryProvisionedURL(lookupCtx, host, accountID, accessToken, httpClient) + if err != nil { + log.Warnf(ctx, "Primary provisioned URL lookup failed: %v", err) + return host + } + if spogURL == "" { + return host + } + return strings.TrimSuffix(spogURL, "/") +} + +// shouldResolveWorkspacePrimaryURL reports whether to look up the primary +// (SPOG) URL of the account that owns the workspace host. Both IDs are +// required so the profile can still target the same workspace through the +// SPOG host; classic account hosts and hosts that are already unified are +// skipped. +func shouldResolveWorkspacePrimaryURL(authArguments *auth.AuthArguments) bool { + if authArguments.Host == "" || authArguments.AccountID == "" || + authArguments.WorkspaceID == "" || authArguments.WorkspaceID == auth.WorkspaceIDNone { + return false + } + if auth.IsClassicAccountHost((&config.Config{Host: authArguments.Host}).CanonicalHostName()) { + return false + } + return !auth.HasUnifiedHostSignal(authArguments.DiscoveryURL) && !auth.IsSpogWorkspaceDiscoveryURL(authArguments.DiscoveryURL) +} + +// spogWorkspaceForHost returns the primary (SPOG) URL of the account that +// owns the workspace host, the workspace's ID, and the workspace discovery URL +// to save with them. The switch only happens once the SPOG host is confirmed +// to be a unified host that serves this workspace's OAuth from the workspace +// host itself, so the workspace-scoped token already minted keeps refreshing +// where it was issued. Best-effort: ok=false keeps the workspace host. +func spogWorkspaceForHost(ctx context.Context, authArguments *auth.AuthArguments) (spogHost, workspaceID, discoveryURL string, ok bool) { + if !shouldResolveWorkspacePrimaryURL(authArguments) { + return "", "", "", false + } + workspaceHost := (&config.Config{Host: authArguments.Host}).CanonicalHostName() + lookupCtx, cancel := context.WithTimeout(ctx, provisionedURLTimeout) + defer cancel() + resp, err := auth.LookupWorkspacePrimaryURL(lookupCtx, workspaceHost, nil) + if err != nil { + log.Debugf(ctx, "Workspace primary URL lookup failed: %v", err) + return "", "", "", false + } + if resp.PrimaryURL == "" { + return "", "", "", false + } + primaryURL := (&config.Config{Host: resp.PrimaryURL}).CanonicalHostName() + if primaryURL == workspaceHost { + return "", "", "", false + } + + // On the SPOG host workspace_id decides routing, so use the ID of the + // workspace at this host over one inherited from an existing profile. + workspaceID = authArguments.WorkspaceID + if resp.WorkspaceID != "" { + workspaceID = resp.WorkspaceID + } + + spogArgs := &auth.AuthArguments{ + Host: primaryURL, + AccountID: authArguments.AccountID, + WorkspaceID: workspaceID, + } + runHostDiscoveryWithRetryTimeout(ctx, spogArgs, int(provisionedURLTimeout.Seconds())) + if !auth.HasUnifiedHostSignal(spogArgs.DiscoveryURL) { + log.Warnf(ctx, "Workspace primary URL %s is not a unified host; keeping %s", primaryURL, workspaceHost) + return "", "", "", false + } + + discoveryURL = auth.SpogWorkspaceDiscoveryURL(primaryURL, workspaceID) + tokenHost, ok := lookupTokenEndpointHost(ctx, discoveryURL) + if !ok || tokenHost != hostOf(workspaceHost) { + log.Warnf(ctx, "Primary URL %s does not serve OAuth for workspace %s; keeping %s", primaryURL, workspaceID, workspaceHost) + return "", "", "", false + } + return primaryURL, workspaceID, discoveryURL, true +} + +// useSpogWorkspaceOAuth points a SPOG-host login for a workspace named before +// login at that workspace's OAuth endpoints, so the profile gets the same +// workspace-scoped token a login to the workspace's own host would. Without a +// workspace, or if the SPOG host doesn't serve the workspace's OAuth metadata, +// the login stays account-level. +func useSpogWorkspaceOAuth(ctx context.Context, authArguments *auth.AuthArguments, workspaceID string) { + if !canUseSpogWorkspaceOAuth(authArguments, workspaceID) { + return + } + discoveryURL := auth.SpogWorkspaceDiscoveryURL((&config.Config{Host: authArguments.Host}).CanonicalHostName(), workspaceID) + if _, ok := lookupTokenEndpointHost(ctx, discoveryURL); !ok { + log.Warnf(ctx, "%s does not serve OAuth for workspace %s; logging in at the account level", authArguments.Host, workspaceID) + return + } + authArguments.DiscoveryURL = discoveryURL +} + +// canUseSpogWorkspaceOAuth reports whether a login for workspaceID targets a +// SPOG host. Classic accounts.* hosts also have an account-scoped discovery +// URL, but only serve account-level OAuth. +func canUseSpogWorkspaceOAuth(authArguments *auth.AuthArguments, workspaceID string) bool { + if workspaceID == "" || workspaceID == auth.WorkspaceIDNone || !auth.HasUnifiedHostSignal(authArguments.DiscoveryURL) { + return false + } + return !auth.IsClassicAccountHost((&config.Config{Host: authArguments.Host}).CanonicalHostName()) +} + +// lookupTokenEndpointHost returns the host of the token endpoint served at +// discoveryURL. Failures are logged and reported as ok=false. +func lookupTokenEndpointHost(ctx context.Context, discoveryURL string) (string, bool) { + lookupCtx, cancel := context.WithTimeout(ctx, provisionedURLTimeout) + defer cancel() + tokenEndpoint, err := auth.LookupOAuthTokenEndpoint(lookupCtx, discoveryURL, nil) + if err != nil { + log.Debugf(ctx, "OAuth metadata lookup at %s failed: %v", discoveryURL, err) + return "", false + } + host := hostOf(tokenEndpoint) + return host, host != "" +} + +// hostOf returns the host[:port] of rawURL, or an empty string if it can't +// be parsed. +func hostOf(rawURL string) string { + u, err := url.Parse(rawURL) + if err != nil { + return "" + } + return u.Host +} + +// spogWorkspaceDiscoveryURLToSave returns discoveryURL when it marks a +// workspace-scoped SPOG profile, the only discovery_url login writes. +func spogWorkspaceDiscoveryURLToSave(discoveryURL string) string { + if auth.IsSpogWorkspaceDiscoveryURL(discoveryURL) { + return discoveryURL + } + return "" +} + // discoveryLoginInputs groups the dependencies of discoveryLogin. // See https://google.github.io/styleguide/go/best-practices#option-structure. type discoveryLoginInputs struct { @@ -688,11 +908,16 @@ type discoveryLoginInputs struct { browserFunc func(string) error tokenStore storage.Store mode storage.StorageMode + // httpClient overrides the client used for the primary provisioned URL + // lookup. Nil in production (uses http.DefaultClient); set in tests. + httpClient *http.Client } // discoveryLogin runs the login.databricks.com discovery flow. The user -// authenticates in the browser, selects a workspace, and the CLI receives -// the workspace host from the OAuth callback's iss parameter. +// authenticates in the browser and selects a workspace or an account; the CLI +// receives the resulting host from the OAuth callback's iss parameter. When an +// account is selected (a classic account host), the profile is switched to the +// account's primary provisioned (SPOG) URL, matching the --account-id path. func discoveryLogin(ctx context.Context, in discoveryLoginInputs) error { arg, err := in.dc.NewOAuthArgument(in.profileName) if err != nil { @@ -774,13 +999,40 @@ func discoveryLogin(ctx context.Context, in discoveryLoginInputs) error { } } + // If the user selected an account (rather than a specific workspace), switch + // to the account's primary provisioned (SPOG) URL so the saved profile + // targets the unified host, matching the --account-id login path. A classic + // account host means an account was selected; a workspace selection yields a + // workspace host that introspection still backfills accountID for, so gate on + // the host type to avoid rewriting a concrete workspace host. + if shouldResolveProvisionedURL(discoveredHost, accountID) { + discoveredHost = resolvePrimaryProvisionedURL(ctx, discoveredHost, accountID, tok.AccessToken, in.httpClient) + } + + // A workspace whose account has a primary (SPOG) URL is saved on that URL + // with its workspace-scoped token, matching the --host workspace path. + var tokenArg u2m.OAuthArgument = arg + var discoveryURL string + if spogHost, spogWorkspaceID, spogDiscoveryURL, ok := spogWorkspaceForHost(ctx, &auth.AuthArguments{ + Host: discoveredHost, + AccountID: accountID, + WorkspaceID: workspaceID, + }); ok { + spogArg, err := u2m.NewProfileWorkspaceOAuthArgumentWithDiscoveryURL(spogHost, spogDiscoveryURL, in.profileName) + if err != nil { + return err + } + discoveredHost, workspaceID, discoveryURL, tokenArg = spogHost, spogWorkspaceID, spogDiscoveryURL, spogArg + } + configFile := env.Get(ctx, "DATABRICKS_CONFIG_FILE") clearKeys := oauthLoginClearKeys() - // Discovery login always produces a workspace-level profile pointing at the - // discovered host. Any previous routing metadata (is_unified_host, - // cluster_id, serverless_compute_id) from a prior login to a different host - // type must be cleared so they don't leak into the new profile. account_id - // and workspace_id are re-added from discovery/introspection results. + // Discovery login produces a profile pointing at the discovered host (or the + // account's primary provisioned URL when an account was selected). Any + // previous routing metadata (is_unified_host, cluster_id, + // serverless_compute_id) from a prior login to a different host type must be + // cleared so they don't leak into the new profile. account_id and + // workspace_id are re-added from discovery/introspection results. clearKeys = append( clearKeys, "account_id", @@ -789,15 +1041,19 @@ func discoveryLogin(ctx context.Context, in discoveryLoginInputs) error { "cluster_id", "serverless_compute_id", ) + if in.existingProfile != nil && auth.IsSpogWorkspaceDiscoveryURL(in.existingProfile.DiscoveryURL) { + clearKeys = append(clearKeys, databrickscfg.DiscoveryURLKey) + } err = databrickscfg.SaveToProfile(ctx, &config.Config{ - Profile: in.profileName, - Host: discoveredHost, - AuthType: authTypeDatabricksCLI, - AccountID: accountID, - WorkspaceID: workspaceID, - Scopes: scopesList, - ConfigFile: configFile, - ClientID: in.clientID, + Profile: in.profileName, + Host: discoveredHost, + AuthType: authTypeDatabricksCLI, + AccountID: accountID, + WorkspaceID: workspaceID, + Scopes: scopesList, + ConfigFile: configFile, + ClientID: in.clientID, + DiscoveryURL: discoveryURL, }, clearKeys...) if err != nil { if configFile != "" { @@ -805,7 +1061,7 @@ func discoveryLogin(ctx context.Context, in discoveryLoginInputs) error { } return fmt.Errorf("saving profile %q: %w", in.profileName, err) } - if err := storeLoginToken(ctx, in.tokenStore, in.mode, arg, tok); err != nil { + if err := storeLoginToken(ctx, in.tokenStore, in.mode, tokenArg, tok); err != nil { return err } diff --git a/cmd/auth/login_test.go b/cmd/auth/login_test.go index 1c4428596dd..2f43bd70f12 100644 --- a/cmd/auth/login_test.go +++ b/cmd/auth/login_test.go @@ -8,6 +8,7 @@ import ( "log/slog" "net/http" "net/http/httptest" + "net/url" "os" "path/filepath" "sync" @@ -116,8 +117,9 @@ type fakeDiscoveryClient struct { introspection *auth.IntrospectionResult introspectionErr error // For assertions - introspectHost string - introspectToken string + introspectHost string + introspectToken string + newPersistentAuthCall int } func (f *fakeDiscoveryClient) NewOAuthArgument(profileName string) (*u2m.BasicDiscoveryOAuthArgument, error) { @@ -131,6 +133,7 @@ func (f *fakeDiscoveryClient) NewPersistentAuth(ctx context.Context, opts ...u2m if f.persistentAuthErr != nil { return nil, f.persistentAuthErr } + f.newPersistentAuthCall++ return f.persistentAuth, nil } @@ -1247,6 +1250,181 @@ func TestDiscoveryLogin_SPOGHostPopulatesAccountIDFromDiscovery(t *testing.T) { assert.Equal(t, "discovered-ws", savedProfile.WorkspaceID, "workspace_id should come from host discovery") } +func TestShouldResolveProvisionedURL(t *testing.T) { + tests := []struct { + name string + host string + account string + expected bool + }{ + {"classic account host with account id", "https://accounts.cloud.databricks.com", "abc-123", true}, + {"account host without scheme", "accounts.cloud.databricks.com", "abc-123", true}, + {"classic account host without account id", "https://accounts.cloud.databricks.com", "", false}, + {"workspace host with account id", "https://dbc-abc.cloud.databricks.com", "abc-123", false}, + {"unified host with account id", "https://mycompany.databricks.com", "abc-123", false}, + {"empty host with account id", "", "abc-123", false}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + assert.Equal(t, tt.expected, shouldResolveProvisionedURL(tt.host, tt.account)) + }) + } +} + +// rewriteHostTransport routes every request to target (a test server), keeping +// the request's path and query, so a lookup addressed to a classic account host +// can be served locally. +type rewriteHostTransport struct { + target string +} + +func (rt rewriteHostTransport) RoundTrip(req *http.Request) (*http.Response, error) { + u, err := url.Parse(rt.target) + if err != nil { + return nil, err + } + req.URL.Scheme = u.Scheme + req.URL.Host = u.Host + return http.DefaultTransport.RoundTrip(req) +} + +// assertNoRequestTransport fails the test if any HTTP request is made through it. +type assertNoRequestTransport struct { + t *testing.T +} + +func (rt assertNoRequestTransport) RoundTrip(req *http.Request) (*http.Response, error) { + rt.t.Errorf("unexpected provisioned-URL lookup to %s", req.URL) + return nil, errors.New("unexpected request") +} + +func TestDiscoveryLogin_AccountSelectionResolvesProvisionedURL(t *testing.T) { + // The provisioned-urls endpoint returns the account's primary SPOG host. + spogServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + assert.Equal(t, "/api/2.0/accounts/introspection-account/provisioned-urls/primary", r.URL.Path) + w.WriteHeader(http.StatusOK) + _, _ = w.Write([]byte(`{"url": "https://dbc-spog.cloud.databricks.com"}`)) + })) + defer spogServer.Close() + + tmpDir := t.TempDir() + configPath := filepath.Join(tmpDir, ".databrickscfg") + require.NoError(t, os.WriteFile(configPath, []byte(""), 0o600)) + t.Setenv("DATABRICKS_CONFIG_FILE", configPath) + + // A classic account host triggers the provisioned-URL lookup. The reserved + // .invalid TLD keeps host metadata discovery from making a real network call + // (it fast-fails, so account_id falls back to introspection). + oauthArg, err := u2m.NewBasicDiscoveryOAuthArgument("DISCOVERY") + require.NoError(t, err) + oauthArg.SetDiscoveredHost("https://accounts.invalid") + + dc := &fakeDiscoveryClient{ + oauthArg: oauthArg, + persistentAuth: &fakeDiscoveryPersistentAuth{token: &oauth2.Token{AccessToken: "test-token"}}, + introspection: &auth.IntrospectionResult{AccountID: "introspection-account"}, + } + + ctx, _ := cmdio.NewTestContextWithStdout(t.Context()) + err = discoveryLogin(ctx, discoveryLoginInputs{ + dc: dc, + profileName: "DISCOVERY", + timeout: 5 * time.Second, + browserFunc: func(string) error { return nil }, + tokenStore: newTestStore(), + httpClient: &http.Client{Transport: rewriteHostTransport{target: spogServer.URL}}, + }) + require.NoError(t, err) + + savedProfile, err := loadProfileByName(ctx, "DISCOVERY", profile.DefaultProfiler) + require.NoError(t, err) + require.NotNil(t, savedProfile) + assert.Equal(t, "https://dbc-spog.cloud.databricks.com", savedProfile.Host, "host should be switched to the account's primary provisioned URL") + assert.Equal(t, "introspection-account", savedProfile.AccountID) +} + +func TestDiscoveryLogin_WorkspaceSelectionKeepsDiscoveredHost(t *testing.T) { + // A workspace host is not a classic account host, so even though token + // introspection backfills an account_id, the provisioned-URL lookup must not + // run and the discovered workspace host must be preserved. + server := newDiscoveryServer(t, map[string]any{ + "workspace_id": "discovered-ws", + }) + + tmpDir := t.TempDir() + configPath := filepath.Join(tmpDir, ".databrickscfg") + require.NoError(t, os.WriteFile(configPath, []byte(""), 0o600)) + t.Setenv("DATABRICKS_CONFIG_FILE", configPath) + + oauthArg, err := u2m.NewBasicDiscoveryOAuthArgument("DISCOVERY") + require.NoError(t, err) + oauthArg.SetDiscoveredHost(server.URL) + + dc := &fakeDiscoveryClient{ + oauthArg: oauthArg, + persistentAuth: &fakeDiscoveryPersistentAuth{token: &oauth2.Token{AccessToken: "test-token"}}, + introspection: &auth.IntrospectionResult{AccountID: "introspection-account"}, + } + + ctx, _ := cmdio.NewTestContextWithStdout(t.Context()) + err = discoveryLogin(ctx, discoveryLoginInputs{ + dc: dc, + profileName: "DISCOVERY", + timeout: 5 * time.Second, + browserFunc: func(string) error { return nil }, + tokenStore: newTestStore(), + // Any provisioned-URL lookup here would be a bug: fail the test if attempted. + httpClient: &http.Client{Transport: assertNoRequestTransport{t: t}}, + }) + require.NoError(t, err) + + savedProfile, err := loadProfileByName(ctx, "DISCOVERY", profile.DefaultProfiler) + require.NoError(t, err) + require.NotNil(t, savedProfile) + assert.Equal(t, server.URL, savedProfile.Host, "workspace host must be preserved, not rewritten to a provisioned URL") + assert.Equal(t, "introspection-account", savedProfile.AccountID) +} + +func TestDiscoveryLogin_AccountSelectionLookupFailureKeepsHost(t *testing.T) { + // When the provisioned-URL lookup fails, login still succeeds and the profile + // keeps the discovered account host (best-effort enrichment). + failServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.WriteHeader(http.StatusInternalServerError) + })) + defer failServer.Close() + + tmpDir := t.TempDir() + configPath := filepath.Join(tmpDir, ".databrickscfg") + require.NoError(t, os.WriteFile(configPath, []byte(""), 0o600)) + t.Setenv("DATABRICKS_CONFIG_FILE", configPath) + + oauthArg, err := u2m.NewBasicDiscoveryOAuthArgument("DISCOVERY") + require.NoError(t, err) + oauthArg.SetDiscoveredHost("https://accounts.invalid") + + dc := &fakeDiscoveryClient{ + oauthArg: oauthArg, + persistentAuth: &fakeDiscoveryPersistentAuth{token: &oauth2.Token{AccessToken: "test-token"}}, + introspection: &auth.IntrospectionResult{AccountID: "introspection-account"}, + } + + ctx, _ := cmdio.NewTestContextWithStdout(t.Context()) + err = discoveryLogin(ctx, discoveryLoginInputs{ + dc: dc, + profileName: "DISCOVERY", + timeout: 5 * time.Second, + browserFunc: func(string) error { return nil }, + tokenStore: newTestStore(), + httpClient: &http.Client{Transport: rewriteHostTransport{target: failServer.URL}}, + }) + require.NoError(t, err) + + savedProfile, err := loadProfileByName(ctx, "DISCOVERY", profile.DefaultProfiler) + require.NoError(t, err) + require.NotNil(t, savedProfile) + assert.Equal(t, "https://accounts.invalid", savedProfile.Host, "host stays the discovered account host when the lookup fails") +} + func TestDiscoveryLogin_IntrospectionFallsBackWhenDiscoveryFails(t *testing.T) { tmpDir := t.TempDir() configPath := filepath.Join(tmpDir, ".databrickscfg") @@ -1459,3 +1637,378 @@ func TestLoginRejectsPositionalArgWithProfileFlag(t *testing.T) { err := cmd.Execute() assert.ErrorContains(t, err, `argument "https://example.com" cannot be combined with --host or --profile`) } + +func TestShouldResolveWorkspacePrimaryURL(t *testing.T) { + tests := []struct { + name string + args auth.AuthArguments + expected bool + }{ + {"workspace host with account and workspace id", auth.AuthArguments{Host: "https://dbc-abc.cloud.databricks.com", AccountID: "acc", WorkspaceID: "123"}, true}, + {"workspace host without workspace id", auth.AuthArguments{Host: "https://dbc-abc.cloud.databricks.com", AccountID: "acc"}, false}, + {"workspace host with none workspace id", auth.AuthArguments{Host: "https://dbc-abc.cloud.databricks.com", AccountID: "acc", WorkspaceID: auth.WorkspaceIDNone}, false}, + {"workspace host without account id", auth.AuthArguments{Host: "https://dbc-abc.cloud.databricks.com", WorkspaceID: "123"}, false}, + {"classic account host", auth.AuthArguments{Host: "https://accounts.cloud.databricks.com", AccountID: "acc", WorkspaceID: "123"}, false}, + {"unified host", auth.AuthArguments{Host: "https://acme.databricks.com", AccountID: "acc", WorkspaceID: "123", DiscoveryURL: "https://acme.databricks.com/oidc/accounts/acc/.well-known/oauth-authorization-server"}, false}, + {"spog workspace host", auth.AuthArguments{Host: "https://acme.databricks.com", AccountID: "acc", WorkspaceID: "123", DiscoveryURL: auth.SpogWorkspaceDiscoveryURL("https://acme.databricks.com", "123")}, false}, + {"empty host", auth.AuthArguments{AccountID: "acc", WorkspaceID: "123"}, false}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + assert.Equal(t, tt.expected, shouldResolveWorkspacePrimaryURL(&tt.args)) + }) + } +} + +// newUnifiedHostServer serves /.well-known/databricks-config for a unified +// (SPOG) host with an account-scoped OIDC endpoint. +func newUnifiedHostServer(t *testing.T) *httptest.Server { + t.Helper() + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.URL.Path == "/.well-known/databricks-config" { + w.Header().Set("Content-Type", "application/json") + _ = json.NewEncoder(w).Encode(map[string]any{ + "account_id": "spog-account", + "oidc_endpoint": "http://" + r.Host + "/oidc/accounts/{account_id}", + "host_type": "UNIFIED_HOST", + }) + return + } + w.WriteHeader(http.StatusNotFound) + })) + t.Cleanup(server.Close) + return server +} + +// newSpogServer serves /.well-known/databricks-config for a unified (SPOG) +// host with an account-scoped OIDC endpoint, and workspace-level OAuth +// metadata (selected with ?o=) whose endpoints are on the host returned by +// workspaceHost, as the SPOG host serves them for a workspace. +func newSpogServer(t *testing.T, workspaceHost func() string) *httptest.Server { + t.Helper() + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.Header().Set("Content-Type", "application/json") + switch r.URL.Path { + case "/.well-known/databricks-config": + _ = json.NewEncoder(w).Encode(map[string]any{ + "account_id": "spog-account", + "oidc_endpoint": "http://" + r.Host + "/oidc/accounts/{account_id}", + "host_type": "UNIFIED_HOST", + }) + case "/oidc/.well-known/oauth-authorization-server": + if r.URL.Query().Get("o") == "" { + w.WriteHeader(http.StatusNotFound) + return + } + _ = json.NewEncoder(w).Encode(map[string]any{ + "authorization_endpoint": workspaceHost() + "/oidc/v1/authorize", + "token_endpoint": workspaceHost() + "/oidc/v1/token", + }) + default: + w.WriteHeader(http.StatusNotFound) + } + })) + t.Cleanup(server.Close) + return server +} + +// newSpogWorkspacePair returns a canonical workspace server whose account's +// primary URL is a SPOG server that serves the workspace's OAuth from the +// workspace server. +func newSpogWorkspacePair(t *testing.T) (workspace, spog *httptest.Server) { + t.Helper() + var workspaceURL string + spog = newSpogServer(t, func() string { return workspaceURL }) + workspace = newDiscoveryServer(t, map[string]any{ + "account_id": "spog-account", + "workspace_id": "12345", + "primary_url": spog.URL, + }) + workspaceURL = workspace.URL + return workspace, spog +} + +func TestSpogWorkspaceForHost_SwitchesToPrimaryURL(t *testing.T) { + workspace, spog := newSpogWorkspacePair(t) + + tests := []struct { + name string + workspaceID string + }{ + {"workspace id from discovery", "12345"}, + // A re-login with --host can inherit another + // workspace's ID from the existing profile. + {"workspace id inherited from another workspace", "67890"}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + spogHost, workspaceID, discoveryURL, ok := spogWorkspaceForHost(t.Context(), &auth.AuthArguments{ + Host: workspace.URL, + AccountID: "spog-account", + WorkspaceID: tt.workspaceID, + }) + require.True(t, ok) + assert.Equal(t, spog.URL, spogHost) + assert.Equal(t, "12345", workspaceID) + assert.Equal(t, auth.SpogWorkspaceDiscoveryURL(spog.URL, "12345"), discoveryURL) + }) + } +} + +func TestSpogWorkspaceForHost_KeepsWorkspaceHost(t *testing.T) { + tests := []struct { + name string + workspace func(t *testing.T) string + }{ + { + name: "no primary url", + workspace: func(t *testing.T) string { + return newDiscoveryServer(t, map[string]any{"account_id": "spog-account", "workspace_id": "12345"}).URL + }, + }, + { + name: "primary url is not a unified host", + workspace: func(t *testing.T) string { + notUnified := newDiscoveryServer(t, map[string]any{"workspace_id": "999"}) + return newDiscoveryServer(t, map[string]any{"account_id": "spog-account", "workspace_id": "12345", "primary_url": notUnified.URL}).URL + }, + }, + { + name: "primary url lookup fails", + workspace: func(t *testing.T) string { + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.WriteHeader(http.StatusInternalServerError) + })) + t.Cleanup(server.Close) + return server.URL + }, + }, + { + name: "primary url serves the workspace's tokens from another host", + workspace: func(t *testing.T) string { + spog := newSpogServer(t, func() string { return "https://other-workspace.test" }) + return newDiscoveryServer(t, map[string]any{"account_id": "spog-account", "workspace_id": "12345", "primary_url": spog.URL}).URL + }, + }, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + _, _, _, ok := spogWorkspaceForHost(t.Context(), &auth.AuthArguments{ + Host: tt.workspace(t), + AccountID: "spog-account", + WorkspaceID: "12345", + }) + assert.False(t, ok) + }) + } +} + +func TestSetLoginHostAndAccountId_SpogWorkspaceOAuth(t *testing.T) { + workspace, spog := newSpogWorkspacePair(t) + spogWorkspaceDiscoveryURL := auth.SpogWorkspaceDiscoveryURL(spog.URL, "12345") + + tests := []struct { + name string + host string + workspaceID string + existingProfile *profile.Profile + wantWorkspaceScoped bool + wantHost string + wantWorkspaceIDSaved string + }{ + { + name: "workspace from ?o= gets workspace oauth", + host: spog.URL + "?o=12345", + wantWorkspaceScoped: true, + wantHost: spog.URL, + wantWorkspaceIDSaved: "12345", + }, + { + name: "workspace from --workspace-id gets workspace oauth", + host: spog.URL, + workspaceID: "12345", + wantWorkspaceScoped: true, + wantHost: spog.URL, + wantWorkspaceIDSaved: "12345", + }, + { + name: "no workspace stays account level", + host: spog.URL, + wantHost: spog.URL, + }, + { + name: "workspace inherited from an account-level profile stays account level", + host: spog.URL, + existingProfile: &profile.Profile{Name: "P", Host: spog.URL, AccountID: "spog-account", WorkspaceID: "12345"}, + wantHost: spog.URL, + wantWorkspaceIDSaved: "12345", + }, + { + name: "workspace inherited from a workspace-scoped profile keeps workspace oauth", + host: spog.URL, + existingProfile: &profile.Profile{Name: "P", Host: spog.URL, AccountID: "spog-account", WorkspaceID: "12345", DiscoveryURL: spogWorkspaceDiscoveryURL}, + wantWorkspaceScoped: true, + wantHost: spog.URL, + wantWorkspaceIDSaved: "12345", + }, + { + name: "?o= on an account-level profile's host stays account level", + host: "", + existingProfile: &profile.Profile{Name: "P", Host: spog.URL + "?o=12345", AccountID: "spog-account"}, + wantHost: spog.URL, + wantWorkspaceIDSaved: "12345", + }, + { + name: "canonical workspace host logs in as before", + host: workspace.URL, + wantHost: workspace.URL, + wantWorkspaceIDSaved: "12345", + }, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + args := &auth.AuthArguments{Host: tt.host, WorkspaceID: tt.workspaceID} + err := setLoginHostAndAccountId(t.Context(), tt.existingProfile, args, []string{}) + require.NoError(t, err) + + assert.Equal(t, tt.wantHost, args.Host) + assert.Equal(t, tt.wantWorkspaceIDSaved, args.WorkspaceID) + if tt.wantWorkspaceScoped { + assert.Equal(t, spogWorkspaceDiscoveryURL, args.DiscoveryURL) + } else { + assert.False(t, auth.IsSpogWorkspaceDiscoveryURL(args.DiscoveryURL), "discovery URL %q", args.DiscoveryURL) + } + }) + } +} + +func TestDiscoveryLogin_WorkspaceSelectionSavesSpogWorkspaceProfile(t *testing.T) { + workspace, spog := newSpogWorkspacePair(t) + + tmpDir := t.TempDir() + configPath := filepath.Join(tmpDir, ".databrickscfg") + require.NoError(t, os.WriteFile(configPath, []byte(""), 0o600)) + t.Setenv("DATABRICKS_CONFIG_FILE", configPath) + + oauthArg, err := u2m.NewBasicDiscoveryOAuthArgument("DISCOVERY") + require.NoError(t, err) + oauthArg.SetDiscoveredHost(workspace.URL) + + dc := &fakeDiscoveryClient{ + oauthArg: oauthArg, + persistentAuth: &fakeDiscoveryPersistentAuth{token: &oauth2.Token{AccessToken: "workspace-token"}}, + introspection: &auth.IntrospectionResult{}, + } + store := &inMemoryStore{Tokens: map[string]*oauth2.Token{}} + + ctx, _ := cmdio.NewTestContextWithStdout(t.Context()) + err = discoveryLogin(ctx, discoveryLoginInputs{ + dc: dc, + profileName: "DISCOVERY", + timeout: 5 * time.Second, + browserFunc: func(string) error { return nil }, + tokenStore: store, + }) + require.NoError(t, err) + + savedProfile, err := loadProfileByName(ctx, "DISCOVERY", profile.DefaultProfiler) + require.NoError(t, err) + require.NotNil(t, savedProfile) + assert.Equal(t, spog.URL, savedProfile.Host) + assert.Equal(t, "spog-account", savedProfile.AccountID) + assert.Equal(t, "12345", savedProfile.WorkspaceID) + assert.Equal(t, auth.SpogWorkspaceDiscoveryURL(spog.URL, "12345"), savedProfile.DiscoveryURL) + require.Contains(t, store.Tokens, "DISCOVERY") + assert.Equal(t, "workspace-token", store.Tokens["DISCOVERY"].AccessToken) + assert.Equal(t, 1, dc.newPersistentAuthCall, "the discovery login is the only login") +} + +func TestSetHostAndAccountId_FollowsProfileTokenType(t *testing.T) { + _, spog := newSpogWorkspacePair(t) + spogWorkspaceDiscoveryURL := auth.SpogWorkspaceDiscoveryURL(spog.URL, "12345") + + tests := []struct { + name string + host string + workspaceID string + existingProfile *profile.Profile + wantWorkspaceScoped bool + }{ + { + name: "?o= on the host stays account level", + host: spog.URL + "?o=12345", + }, + { + name: "--workspace-id stays account level", + host: spog.URL, + workspaceID: "12345", + }, + { + name: "account-level profile with ?o= in its host stays account level", + existingProfile: &profile.Profile{Name: "P", Host: spog.URL + "?o=12345", AccountID: "spog-account"}, + }, + { + name: "workspace-scoped profile keeps workspace oauth", + existingProfile: &profile.Profile{Name: "P", Host: spog.URL, AccountID: "spog-account", WorkspaceID: "12345", DiscoveryURL: spogWorkspaceDiscoveryURL}, + wantWorkspaceScoped: true, + }, + { + name: "workspace-scoped profile targeted at another workspace stays account level", + workspaceID: "67890", + existingProfile: &profile.Profile{Name: "P", Host: spog.URL, AccountID: "spog-account", WorkspaceID: "12345", DiscoveryURL: spogWorkspaceDiscoveryURL}, + }, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + args := &auth.AuthArguments{Host: tt.host, WorkspaceID: tt.workspaceID} + err := setHostAndAccountId(t.Context(), tt.existingProfile, args, []string{}) + require.NoError(t, err) + + if tt.wantWorkspaceScoped { + assert.Equal(t, spogWorkspaceDiscoveryURL, args.DiscoveryURL) + } else { + assert.True(t, auth.HasUnifiedHostSignal(args.DiscoveryURL), "discovery URL %q", args.DiscoveryURL) + } + }) + } +} + +func TestSetHostAndAccountId_DoesNotSwitchToPrimaryURL(t *testing.T) { + // setHostAndAccountId also serves auth token, which must keep the profile's + // host so the cached token is found and refreshed where it was issued. + spog := newUnifiedHostServer(t) + workspace := newDiscoveryServer(t, map[string]any{ + "account_id": "spog-account", + "workspace_id": "12345", + "primary_url": spog.URL, + }) + + args := &auth.AuthArguments{Host: workspace.URL} + err := setHostAndAccountId(t.Context(), nil, args, []string{}) + require.NoError(t, err) + + assert.Equal(t, workspace.URL, args.Host) + assert.False(t, auth.HasUnifiedHostSignal(args.DiscoveryURL), "discovery URL %q", args.DiscoveryURL) +} + +func TestCanUseSpogWorkspaceOAuth(t *testing.T) { + const spogDiscoveryURL = "https://acme.databricks.test/oidc/accounts/abc/.well-known/oauth-authorization-server" + const accountsDiscoveryURL = "https://accounts.cloud.databricks.com/oidc/accounts/abc/.well-known/oauth-authorization-server" + tests := []struct { + name string + args auth.AuthArguments + workspaceID string + want bool + }{ + {"spog host with a workspace", auth.AuthArguments{Host: "https://acme.databricks.test", DiscoveryURL: spogDiscoveryURL}, "123", true}, + {"spog host without a workspace", auth.AuthArguments{Host: "https://acme.databricks.test", DiscoveryURL: spogDiscoveryURL}, "", false}, + {"spog host with the none workspace", auth.AuthArguments{Host: "https://acme.databricks.test", DiscoveryURL: spogDiscoveryURL}, auth.WorkspaceIDNone, false}, + {"classic account host with a workspace", auth.AuthArguments{Host: "https://accounts.cloud.databricks.com", DiscoveryURL: accountsDiscoveryURL}, "123", false}, + {"workspace host", auth.AuthArguments{Host: "https://dbc-123.cloud.databricks.test"}, "123", false}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + assert.Equal(t, tt.want, canUseSpogWorkspaceOAuth(&tt.args, tt.workspaceID)) + }) + } +} diff --git a/cmd/auth/logout.go b/cmd/auth/logout.go index d1d559b0a60..2e6e991f0a5 100644 --- a/cmd/auth/logout.go +++ b/cmd/auth/logout.go @@ -274,6 +274,12 @@ func clearTokenStore(ctx context.Context, p profile.Profile, profiler profile.Pr return fmt.Errorf("failed to delete profile-keyed token for profile %q: %w", p.Name, err) } + // Workspace-scoped SPOG tokens have no host-keyed copy, and the host's + // key belongs to account-level logins. + if auth.IsSpogWorkspaceDiscoveryURL(p.DiscoveryURL) { + return nil + } + hostCacheKey, matchFn := hostCacheKeyAndMatchFn(p) if hostCacheKey == "" { return fmt.Errorf("failed to get host-based cache key for profile %q", p.Name) diff --git a/cmd/auth/logout_test.go b/cmd/auth/logout_test.go index 366b2cc4e4e..10d3316dbca 100644 --- a/cmd/auth/logout_test.go +++ b/cmd/auth/logout_test.go @@ -380,6 +380,43 @@ auth_type = databricks-cli assert.Nil(t, tokenStore.Tokens[hostKey]) } +func TestLogoutSpogWorkspaceScopedProfileKeepsHostToken(t *testing.T) { + spogServer := newWellKnownServer(t, true, "spog-acct") + + ctx := cmdio.MockDiscard(t.Context()) + configPath := writeTempConfig(t, `[DEFAULT] +[spog-ws] +host = `+spogServer.URL+` +account_id = spog-acct +workspace_id = 123 +auth_type = databricks-cli +discovery_url = `+spogServer.URL+`/oidc/.well-known/oauth-authorization-server?o=123 +`) + t.Setenv("DATABRICKS_CONFIG_FILE", configPath) + + // The host key holds an account-level login's token, which logging out + // of a workspace-scoped profile must not remove. + hostKey := spogServer.URL + "/oidc/accounts/spog-acct" + tokenStore := &inMemoryStore{ + Tokens: map[string]*oauth2.Token{ + "spog-ws": {AccessToken: "workspace-token"}, + hostKey: {AccessToken: "account-token"}, + }, + } + + err := runLogout(ctx, logoutArgs{ + profileName: "spog-ws", + autoApprove: true, + profiler: profile.DefaultProfiler, + tokenStore: tokenStore, + configFilePath: configPath, + }) + require.NoError(t, err) + + assert.Nil(t, tokenStore.Tokens["spog-ws"]) + assert.NotNil(t, tokenStore.Tokens[hostKey]) +} + func TestHostCacheKeyAndMatchFn(t *testing.T) { wsServer := newWellKnownServer(t, false, "ws-account") spogServer := newWellKnownServer(t, true, "spog-account") diff --git a/cmd/auth/token.go b/cmd/auth/token.go index e28d71083a0..6611a171d7b 100644 --- a/cmd/auth/token.go +++ b/cmd/auth/token.go @@ -231,6 +231,9 @@ func loadToken(ctx context.Context, args loadTokenArgs) (*oauth2.Token, error) { } else { matchFn = profile.WithHost(args.authArguments.Host) } + if args.authArguments.WorkspaceID == "" { + matchFn = withoutSpogWorkspaceToken(matchFn) + } matchingProfiles, err := args.profiler.LoadProfiles(ctx, matchFn) if err != nil && !errors.Is(err, profile.ErrNoConfiguration) { @@ -279,6 +282,10 @@ func loadToken(ctx context.Context, args loadTokenArgs) (*oauth2.Token, error) { ) } + // The profile may have been resolved from --host above, after + // setHostAndAccountId ran. + useProfileSpogWorkspaceOAuth(args.authArguments, existingProfile) + args.authArguments.Profile = args.profileName ctx, cancel := context.WithTimeout(ctx, args.tokenTimeout) @@ -519,3 +526,12 @@ func runInlineLogin(ctx context.Context, profiler profile.Profiler, tokenStore s } return profileName, p, nil } + +// withoutSpogWorkspaceToken narrows matchFn to profiles that don't hold a +// workspace-scoped SPOG token. A host lookup without a workspace ID asks for +// an account-level token, which such a profile can't refresh. +func withoutSpogWorkspaceToken(matchFn profile.ProfileMatchFunction) profile.ProfileMatchFunction { + return func(p profile.Profile) bool { + return matchFn(p) && !auth.IsSpogWorkspaceDiscoveryURL(p.DiscoveryURL) + } +} diff --git a/cmd/auth/token_test.go b/cmd/auth/token_test.go index e37e7cdc01b..78d3c098526 100644 --- a/cmd/auth/token_test.go +++ b/cmd/auth/token_test.go @@ -18,6 +18,7 @@ import ( "github.com/databricks/cli/libs/env" "github.com/databricks/databricks-sdk-go/httpclient/fixtures" "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" "golang.org/x/oauth2" ) @@ -965,3 +966,146 @@ func TestWriteTokenErrorOutput(t *testing.T) { assert.Equal(t, unauthenticatedErrorCode, got.ErrorCode) assert.Equal(t, "refresh token is invalid", got.Message) } + +// recordingEndpointSupplier records which kind of OAuth endpoints a refresh +// was resolved from. +type recordingEndpointSupplier struct { + MockApiClient + used *[]string +} + +func (r *recordingEndpointSupplier) GetUnifiedOAuthEndpoints(ctx context.Context, host, accountId string) (*u2m.OAuthAuthorizationServer, error) { + *r.used = append(*r.used, "unified") + return r.MockApiClient.GetUnifiedOAuthEndpoints(ctx, host, accountId) +} + +func (r *recordingEndpointSupplier) GetEndpointsFromURL(_ context.Context, rawURL string) (*u2m.OAuthAuthorizationServer, error) { + *r.used = append(*r.used, rawURL) + return &u2m.OAuthAuthorizationServer{ + TokenEndpoint: "https://workspace.test/oidc/v1/token", + AuthorizationEndpoint: "https://workspace.test/oidc/v1/authorize", + }, nil +} + +func TestToken_loadTokenSpogProfileTokenType(t *testing.T) { + _, spog := newSpogWorkspacePair(t) + spogWorkspaceDiscoveryURL := auth.SpogWorkspaceDiscoveryURL(spog.URL, "12345") + + profiler := profile.InMemoryProfiler{ + Profiles: profile.Profiles{ + {Name: "account-level-o", Host: spog.URL + "?o=12345", AccountID: "spog-account"}, + {Name: "workspace-scoped", Host: spog.URL, AccountID: "spog-account", WorkspaceID: "12345", DiscoveryURL: spogWorkspaceDiscoveryURL}, + }, + } + + tests := []struct { + name string + authArguments *auth.AuthArguments + profileName string + wantEndpoints string + }{ + { + name: "account-level profile with ?o= in its host refreshes at the account endpoint", + authArguments: &auth.AuthArguments{}, + profileName: "account-level-o", + wantEndpoints: "unified", + }, + { + name: "--workspace-id does not change an account-level profile", + authArguments: &auth.AuthArguments{WorkspaceID: "12345"}, + profileName: "account-level-o", + wantEndpoints: "unified", + }, + { + name: "workspace-scoped profile refreshes at the workspace endpoints", + authArguments: &auth.AuthArguments{}, + profileName: "workspace-scoped", + wantEndpoints: spogWorkspaceDiscoveryURL, + }, + { + name: "workspace-scoped profile matched by --host refreshes at the workspace endpoints", + authArguments: &auth.AuthArguments{Host: spog.URL, WorkspaceID: "12345"}, + wantEndpoints: spogWorkspaceDiscoveryURL, + }, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + tokenStore := &inMemoryStore{Tokens: map[string]*oauth2.Token{ + "account-level-o": {RefreshToken: "account-level-o"}, + "workspace-scoped": {RefreshToken: "workspace-scoped"}, + }} + var used []string + _, err := loadToken(cmdio.MockDiscard(t.Context()), loadTokenArgs{ + authArguments: tt.authArguments, + profileName: tt.profileName, + args: []string{}, + tokenTimeout: time.Minute, + profiler: profiler, + tokenStore: tokenStore, + persistentAuthOpts: []u2m.PersistentAuthOption{ + u2m.WithTokenStore(tokenStore), + u2m.WithOAuthEndpointSupplier(&recordingEndpointSupplier{used: &used}), + u2m.WithHttpClient(&http.Client{Transport: fixtures.SliceTransport{refreshSuccessTokenResponse}}), + }, + }) + require.NoError(t, err) + assert.Equal(t, []string{tt.wantEndpoints}, used) + }) + } +} + +func TestToken_loadTokenHostWithoutWorkspaceIDSkipsSpogWorkspaceProfile(t *testing.T) { + _, spog := newSpogWorkspacePair(t) + workspaceScoped := profile.Profile{ + Name: "workspace-scoped", + Host: spog.URL, + AccountID: "spog-account", + WorkspaceID: "12345", + DiscoveryURL: auth.SpogWorkspaceDiscoveryURL(spog.URL, "12345"), + } + accountLevel := profile.Profile{Name: "account-level", Host: spog.URL, AccountID: "spog-account", WorkspaceID: "12345"} + + tests := []struct { + name string + profiles profile.Profiles + wantErr bool + }{ + { + name: "uses the account-level profile", + profiles: profile.Profiles{workspaceScoped, accountLevel}, + }, + { + name: "fails without an account-level profile", + profiles: profile.Profiles{workspaceScoped}, + wantErr: true, + }, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + tokenStore := &inMemoryStore{Tokens: map[string]*oauth2.Token{ + "workspace-scoped": {RefreshToken: "workspace-scoped"}, + "account-level": {RefreshToken: "account-level"}, + }} + var used []string + _, err := loadToken(cmdio.MockDiscard(t.Context()), loadTokenArgs{ + authArguments: &auth.AuthArguments{Host: spog.URL}, + args: []string{}, + tokenTimeout: time.Minute, + profiler: profile.InMemoryProfiler{Profiles: tt.profiles}, + tokenStore: tokenStore, + persistentAuthOpts: []u2m.PersistentAuthOption{ + u2m.WithTokenStore(tokenStore), + u2m.WithOAuthEndpointSupplier(&recordingEndpointSupplier{used: &used}), + u2m.WithHttpClient(&http.Client{Transport: fixtures.SliceTransport{refreshSuccessTokenResponse}}), + }, + }) + if tt.wantErr { + require.Error(t, err) + assert.Empty(t, used, "the workspace-scoped token must not be refreshed") + return + } + require.NoError(t, err) + assert.Equal(t, []string{"unified"}, used) + }) + } +} diff --git a/libs/auth/arguments.go b/libs/auth/arguments.go index deac0b5b1cc..12c21df837e 100644 --- a/libs/auth/arguments.go +++ b/libs/auth/arguments.go @@ -69,5 +69,11 @@ func (a AuthArguments) ToOAuthArgument() (u2m.OAuthArgument, error) { return u2m.NewProfileUnifiedOAuthArgument(host, cfg.AccountID, a.Profile) } + // A SPOG host that targets one workspace with a workspace-scoped token: + // OAuth runs against the endpoints served by the workspace discovery URL. + if IsSpogWorkspaceDiscoveryURL(cfg.DiscoveryURL) { + return u2m.NewProfileWorkspaceOAuthArgumentWithDiscoveryURL(host, cfg.DiscoveryURL, a.Profile) + } + return u2m.NewProfileWorkspaceOAuthArgument(host, a.Profile) } diff --git a/libs/auth/arguments_test.go b/libs/auth/arguments_test.go index 873ce557dfe..150ba64513f 100644 --- a/libs/auth/arguments_test.go +++ b/libs/auth/arguments_test.go @@ -257,3 +257,22 @@ func TestToOAuthArgument_NoAccountIDSkipsUnifiedRouting(t *testing.T) { _, ok := got.(u2m.WorkspaceOAuthArgument) assert.True(t, ok, "expected WorkspaceOAuthArgument when no caller AccountID, got %T", got) } + +func TestToOAuthArgument_SpogWorkspaceDiscoveryURLRoutesToWorkspace(t *testing.T) { + discoveryURL := SpogWorkspaceDiscoveryURL("https://acme.databricks.test", "123") + args := AuthArguments{ + Host: "https://acme.databricks.test", + AccountID: "spog-account", + WorkspaceID: "123", + DiscoveryURL: discoveryURL, + Profile: "my-profile", + } + got, err := args.ToOAuthArgument() + require.NoError(t, err) + + ws, ok := got.(u2m.BasicWorkspaceOAuthArgument) + require.True(t, ok, "expected BasicWorkspaceOAuthArgument, got %T", got) + assert.Equal(t, "https://acme.databricks.test", ws.GetWorkspaceHost()) + assert.Equal(t, discoveryURL, ws.GetDiscoveryURL()) + assert.Equal(t, "my-profile", ws.GetCacheKey()) +} diff --git a/libs/auth/config_type.go b/libs/auth/config_type.go index 0d93b1bf075..de6f0f5e8c2 100644 --- a/libs/auth/config_type.go +++ b/libs/auth/config_type.go @@ -50,6 +50,11 @@ func ResolveConfigType(cfg *config.Config) config.ConfigType { return configType } + // A workspace-scoped SPOG token only reaches its workspace. + if IsSpogWorkspaceDiscoveryURL(cfg.DiscoveryURL) { + return config.WorkspaceConfig + } + if !IsSPOG(cfg, cfg.AccountID) { return configType } diff --git a/libs/auth/config_type_test.go b/libs/auth/config_type_test.go index 8ebe8ff7d68..3444b172492 100644 --- a/libs/auth/config_type_test.go +++ b/libs/auth/config_type_test.go @@ -1,10 +1,12 @@ package auth import ( + "context" "testing" "github.com/databricks/databricks-sdk-go/config" "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" ) func TestHasUnifiedHostSignal(t *testing.T) { @@ -67,6 +69,16 @@ func TestResolveConfigType(t *testing.T) { }, want: config.AccountConfig, }, + { + name: "SPOG workspace-scoped OIDC routes to WorkspaceConfig", + cfg: &config.Config{ + Host: "https://spog.databricks.com", + AccountID: "acct-123", + WorkspaceID: "ws-456", + DiscoveryURL: "https://spog.databricks.com/oidc/.well-known/oauth-authorization-server?o=ws-456", + }, + want: config.WorkspaceConfig, + }, { name: "workspace-scoped OIDC with account_id stays WorkspaceConfig", cfg: &config.Config{ @@ -100,3 +112,34 @@ func TestResolveConfigType(t *testing.T) { }) } } + +func TestResolveConfigType_UnifiedHostMetadata(t *testing.T) { + cases := []struct { + name string + discoveryURL string + }{ + {"workspace-scoped SPOG profile", "https://spog.databricks.com/oidc/.well-known/oauth-authorization-server?o=ws-456"}, + {"account-level SPOG profile with a workspace", "https://spog.databricks.com/oidc/accounts/acct-123/.well-known/oauth-authorization-server"}, + } + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + cfg := &config.Config{ + Host: "https://spog.databricks.com", + AccountID: "acct-123", + WorkspaceID: "ws-456", + Token: "token", + DiscoveryURL: tc.discoveryURL, + Loaders: []config.Loader{config.ConfigAttributes}, + HostMetadataResolver: func(context.Context, string) (*config.HostMetadata, error) { + return &config.HostMetadata{ + AccountID: "acct-123", + HostType: config.UnifiedHost, + OIDCEndpoint: "https://spog.databricks.com/oidc/accounts/{account_id}", + }, nil + }, + } + require.NoError(t, cfg.EnsureResolved()) + assert.Equal(t, config.WorkspaceConfig, ResolveConfigType(cfg)) + }) + } +} diff --git a/libs/auth/error.go b/libs/auth/error.go index fdb5f256f43..56ee05229ae 100644 --- a/libs/auth/error.go +++ b/libs/auth/error.go @@ -212,6 +212,13 @@ func BuildLoginCommand(ctx context.Context, profile string, arg u2m.OAuthArgumen cmd = append(cmd, "--host", arg.GetAccountHost(), "--account-id", arg.GetAccountId()) case u2m.WorkspaceOAuthArgument: cmd = append(cmd, "--host", arg.GetWorkspaceHost()) + // A workspace-scoped token on a SPOG host needs the workspace named + // at login; --host alone would log in at the account level. + if d, ok := arg.(u2m.DiscoveryURLProvider); ok { + if workspaceID := SpogWorkspaceIDFromDiscoveryURL(d.GetDiscoveryURL()); workspaceID != "" { + cmd = append(cmd, "--workspace-id", workspaceID) + } + } } } return strings.Join(cmd, " ") diff --git a/libs/auth/provisioned_url.go b/libs/auth/provisioned_url.go new file mode 100644 index 00000000000..fb77a71c945 --- /dev/null +++ b/libs/auth/provisioned_url.go @@ -0,0 +1,102 @@ +package auth + +import ( + "context" + "encoding/json" + "errors" + "fmt" + "io" + "net/http" + "net/url" + "strings" +) + +// ProvisionedURLResponse represents the response from the account "primary +// provisioned URL" endpoint at +// /api/2.0/accounts/{account_id}/provisioned-urls/primary. The primary +// provisioned URL is the account's SPOG (Single Pane of Glass) host. +type ProvisionedURLResponse struct { + URL string `json:"url"` +} + +// WorkspacePrimaryURLResponse is the subset of a workspace host's +// /.well-known/databricks-config response that carries the owning account's +// primary (SPOG) URL. PrimaryURL is only returned when the request sets +// include_primary_url=true. +type WorkspacePrimaryURLResponse struct { + PrimaryURL string `json:"primary_url"` + WorkspaceID string `json:"workspace_id"` +} + +// LookupPrimaryProvisionedURL looks up the account's primary provisioned URL +// (its SPOG host) by account ID. It calls +// /api/2.0/accounts/{account_id}/provisioned-urls/primary on the given host +// using the supplied access token. Returns an error if the request fails or +// the response cannot be parsed. Callers should treat errors as non-fatal +// (best-effort profile enrichment). +func LookupPrimaryProvisionedURL(ctx context.Context, host, accountID, accessToken string, httpClient *http.Client) (string, error) { + endpoint := strings.TrimSuffix(host, "/") + "/api/2.0/accounts/" + url.PathEscape(accountID) + "/provisioned-urls/primary" + var provisioned ProvisionedURLResponse + err := getJSON(ctx, httpClient, endpoint, accessToken, "provisioned-urls", &provisioned) + // Accounts without a primary URL return 404. + if errors.Is(err, errNotFound) { + return "", nil + } + if err != nil { + return "", err + } + return provisioned.URL, nil +} + +// LookupWorkspacePrimaryURL looks up the primary (SPOG) URL of the account +// that owns the workspace at host, along with that workspace's ID. It calls the +// unauthenticated /.well-known/databricks-config?include_primary_url=true +// endpoint. PrimaryURL is empty when the account has no primary URL. Callers +// should treat errors as non-fatal (best-effort profile enrichment). +func LookupWorkspacePrimaryURL(ctx context.Context, host string, httpClient *http.Client) (*WorkspacePrimaryURLResponse, error) { + endpoint := strings.TrimSuffix(host, "/") + "/.well-known/databricks-config?include_primary_url=true" + var discovery WorkspacePrimaryURLResponse + if err := getJSON(ctx, httpClient, endpoint, "", "databricks-config", &discovery); err != nil { + return nil, err + } + return &discovery, nil +} + +// getJSON issues a GET to endpoint and decodes the JSON response into out. An +// empty accessToken sends the request unauthenticated. name identifies the +// endpoint in error messages. +// errNotFound is returned by getJSON for a 404 response. +var errNotFound = errors.New("not found") + +func getJSON(ctx context.Context, httpClient *http.Client, endpoint, accessToken, name string, out any) error { + if httpClient == nil { + httpClient = http.DefaultClient + } + req, err := http.NewRequestWithContext(ctx, http.MethodGet, endpoint, nil) + if err != nil { + return fmt.Errorf("creating %s request: %w", name, err) + } + if accessToken != "" { + req.Header.Set("Authorization", "Bearer "+accessToken) + } + + resp, err := httpClient.Do(req) + if err != nil { + return fmt.Errorf("calling %s endpoint: %w", name, err) + } + defer resp.Body.Close() + + if resp.StatusCode != http.StatusOK { + // Drain the body so the underlying TCP connection can be reused. + _, _ = io.Copy(io.Discard, resp.Body) + if resp.StatusCode == http.StatusNotFound { + return fmt.Errorf("%s endpoint returned status %d: %w", name, resp.StatusCode, errNotFound) + } + return fmt.Errorf("%s endpoint returned status %d", name, resp.StatusCode) + } + + if err := json.NewDecoder(resp.Body).Decode(out); err != nil { + return fmt.Errorf("decoding %s response: %w", name, err) + } + return nil +} diff --git a/libs/auth/provisioned_url_test.go b/libs/auth/provisioned_url_test.go new file mode 100644 index 00000000000..4b334d8d9b0 --- /dev/null +++ b/libs/auth/provisioned_url_test.go @@ -0,0 +1,118 @@ +package auth + +import ( + "net/http" + "net/http/httptest" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func TestLookupPrimaryProvisionedURL_Success(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.WriteHeader(http.StatusOK) + _, _ = w.Write([]byte(`{"url": "https://dbc-abc123.cloud.databricks.com"}`)) + })) + defer server.Close() + + url, err := LookupPrimaryProvisionedURL(t.Context(), server.URL, "abc-123", "test-token", nil) + require.NoError(t, err) + assert.Equal(t, "https://dbc-abc123.cloud.databricks.com", url) +} + +func TestLookupPrimaryProvisionedURL_HTTPError(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.WriteHeader(http.StatusInternalServerError) + })) + defer server.Close() + + _, err := LookupPrimaryProvisionedURL(t.Context(), server.URL, "abc-123", "test-token", nil) + assert.ErrorContains(t, err, "status 500") +} + +func TestLookupPrimaryProvisionedURL_MalformedJSON(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.WriteHeader(http.StatusOK) + _, _ = w.Write([]byte(`not json`)) + })) + defer server.Close() + + _, err := LookupPrimaryProvisionedURL(t.Context(), server.URL, "abc-123", "test-token", nil) + assert.ErrorContains(t, err, "decoding provisioned-urls response") +} + +func TestLookupPrimaryProvisionedURL_VerifyRequestDetails(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + assert.Equal(t, "/api/2.0/accounts/abc-123/provisioned-urls/primary", r.URL.Path) + assert.Equal(t, "Bearer my-secret-token", r.Header.Get("Authorization")) + w.WriteHeader(http.StatusOK) + _, _ = w.Write([]byte(`{"url": "https://dbc-abc123.cloud.databricks.com"}`)) + })) + defer server.Close() + + _, err := LookupPrimaryProvisionedURL(t.Context(), server.URL, "abc-123", "my-secret-token", nil) + require.NoError(t, err) +} + +func TestLookupPrimaryProvisionedURL_EscapesAccountID(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + assert.Equal(t, "/api/2.0/accounts/a%2Fb%3Fc/provisioned-urls/primary", r.URL.EscapedPath()) + assert.Empty(t, r.URL.RawQuery) + _, _ = w.Write([]byte(`{"url": "https://acme.databricks.com"}`)) + })) + defer server.Close() + + _, err := LookupPrimaryProvisionedURL(t.Context(), server.URL, "a/b?c", "token", nil) + require.NoError(t, err) +} + +func TestLookupPrimaryProvisionedURL_NotFoundMeansNoPrimaryURL(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.WriteHeader(http.StatusNotFound) + _, _ = w.Write([]byte(`{"error_code":"NOT_FOUND","message":"Primary provisioned URL does not exist"}`)) + })) + defer server.Close() + + url, err := LookupPrimaryProvisionedURL(t.Context(), server.URL, "abc", "token", nil) + require.NoError(t, err) + assert.Empty(t, url) +} + +func TestLookupWorkspacePrimaryURL_Success(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + assert.Equal(t, "/.well-known/databricks-config", r.URL.Path) + assert.Equal(t, "true", r.URL.Query().Get("include_primary_url")) + assert.Empty(t, r.Header.Get("Authorization")) + w.WriteHeader(http.StatusOK) + _, _ = w.Write([]byte(`{"workspace_id": "123", "primary_url": "https://acme.databricks.com"}`)) + })) + defer server.Close() + + resp, err := LookupWorkspacePrimaryURL(t.Context(), server.URL, nil) + require.NoError(t, err) + assert.Equal(t, "https://acme.databricks.com", resp.PrimaryURL) + assert.Equal(t, "123", resp.WorkspaceID) +} + +func TestLookupWorkspacePrimaryURL_NoPrimaryURL(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.WriteHeader(http.StatusOK) + _, _ = w.Write([]byte(`{"workspace_id": "123"}`)) + })) + defer server.Close() + + resp, err := LookupWorkspacePrimaryURL(t.Context(), server.URL, nil) + require.NoError(t, err) + assert.Empty(t, resp.PrimaryURL) +} + +func TestLookupWorkspacePrimaryURL_HTTPError(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.WriteHeader(http.StatusNotFound) + })) + defer server.Close() + + _, err := LookupWorkspacePrimaryURL(t.Context(), server.URL, nil) + assert.ErrorContains(t, err, "databricks-config endpoint returned status 404") +} diff --git a/libs/auth/spog_workspace.go b/libs/auth/spog_workspace.go new file mode 100644 index 00000000000..78b71d5bebb --- /dev/null +++ b/libs/auth/spog_workspace.go @@ -0,0 +1,59 @@ +package auth + +import ( + "context" + "net/http" + "net/url" + "strings" +) + +// spogWorkspaceDiscoveryPath is the authorization server metadata path for +// workspace-level OAuth. On a SPOG host the workspace is selected with the +// o= query parameter, and the served endpoints are those of the +// workspace's canonical host. +const spogWorkspaceDiscoveryPath = "/oidc/.well-known/oauth-authorization-server" + +// SpogWorkspaceDiscoveryURL returns the workspace-level OAuth metadata URL for +// workspaceID on the SPOG host. Saved as a profile's discovery_url, it marks +// the profile as holding a workspace-scoped token: OAuth runs against the +// workspace's own endpoints while API calls go to the SPOG host. +func SpogWorkspaceDiscoveryURL(spogHost, workspaceID string) string { + return strings.TrimSuffix(spogHost, "/") + spogWorkspaceDiscoveryPath + "?" + url.Values{"o": {workspaceID}}.Encode() +} + +// IsSpogWorkspaceDiscoveryURL reports whether discoveryURL selects one +// workspace on a SPOG host, as produced by [SpogWorkspaceDiscoveryURL]. +// Canonical workspace discovery URLs never carry the o= parameter. +func IsSpogWorkspaceDiscoveryURL(discoveryURL string) bool { + u, err := url.Parse(discoveryURL) + if err != nil { + return false + } + return u.Path == spogWorkspaceDiscoveryPath && u.Query().Get("o") != "" +} + +// SpogWorkspaceIDFromDiscoveryURL returns the workspace ID selected by a +// discovery URL produced by [SpogWorkspaceDiscoveryURL], or "" for any other URL. +func SpogWorkspaceIDFromDiscoveryURL(discoveryURL string) string { + if !IsSpogWorkspaceDiscoveryURL(discoveryURL) { + return "" + } + u, _ := url.Parse(discoveryURL) + return u.Query().Get("o") +} + +// oauthServerMetadata is the subset of an OAuth authorization server metadata +// document needed to check which host serves a workspace's tokens. +type oauthServerMetadata struct { + TokenEndpoint string `json:"token_endpoint"` +} + +// LookupOAuthTokenEndpoint returns the token_endpoint served by the OAuth +// authorization server metadata document at discoveryURL. +func LookupOAuthTokenEndpoint(ctx context.Context, discoveryURL string, httpClient *http.Client) (string, error) { + var metadata oauthServerMetadata + if err := getJSON(ctx, httpClient, discoveryURL, "", "oauth-authorization-server", &metadata); err != nil { + return "", err + } + return metadata.TokenEndpoint, nil +} diff --git a/libs/auth/spog_workspace_test.go b/libs/auth/spog_workspace_test.go new file mode 100644 index 00000000000..7ec40730a1d --- /dev/null +++ b/libs/auth/spog_workspace_test.go @@ -0,0 +1,65 @@ +package auth + +import ( + "github.com/databricks/cli/libs/auth/u2m" + "net/http" + "net/http/httptest" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func TestSpogWorkspaceDiscoveryURL(t *testing.T) { + assert.Equal(t, + "https://acme.databricks.test/oidc/.well-known/oauth-authorization-server?o=123", + SpogWorkspaceDiscoveryURL("https://acme.databricks.test/", "123")) +} + +func TestIsSpogWorkspaceDiscoveryURL(t *testing.T) { + tests := []struct { + name string + discoveryURL string + want bool + }{ + {"spog workspace", "https://acme.databricks.test/oidc/.well-known/oauth-authorization-server?o=123", true}, + {"canonical workspace", "https://dbc-123.cloud.databricks.test/oidc/.well-known/oauth-authorization-server", false}, + {"account scoped", "https://acme.databricks.test/oidc/accounts/abc/.well-known/oauth-authorization-server", false}, + {"other path with o", "https://acme.databricks.test/oidc/accounts/abc/.well-known/oauth-authorization-server?o=123", false}, + {"empty", "", false}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + assert.Equal(t, tt.want, IsSpogWorkspaceDiscoveryURL(tt.discoveryURL)) + }) + } +} + +func TestLookupOAuthTokenEndpoint(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + assert.Equal(t, "/oidc/.well-known/oauth-authorization-server", r.URL.Path) + assert.Equal(t, "123", r.URL.Query().Get("o")) + _, _ = w.Write([]byte(`{"token_endpoint": "https://dbc-123.cloud.databricks.test/oidc/v1/token"}`)) + })) + defer server.Close() + + endpoint, err := LookupOAuthTokenEndpoint(t.Context(), SpogWorkspaceDiscoveryURL(server.URL, "123"), nil) + require.NoError(t, err) + assert.Equal(t, "https://dbc-123.cloud.databricks.test/oidc/v1/token", endpoint) +} + +func TestSpogWorkspaceIDFromDiscoveryURL(t *testing.T) { + assert.Equal(t, "123", SpogWorkspaceIDFromDiscoveryURL(SpogWorkspaceDiscoveryURL("https://acme.databricks.test", "123"))) + assert.Empty(t, SpogWorkspaceIDFromDiscoveryURL("https://acme.databricks.test/oidc/accounts/abc/.well-known/oauth-authorization-server")) + assert.Empty(t, SpogWorkspaceIDFromDiscoveryURL("")) +} + +func TestBuildLoginCommand_SpogWorkspaceArgumentNamesWorkspace(t *testing.T) { + arg, err := u2m.NewProfileWorkspaceOAuthArgumentWithDiscoveryURL("https://acme.databricks.test", SpogWorkspaceDiscoveryURL("https://acme.databricks.test", "123"), "") + require.NoError(t, err) + assert.Equal(t, "databricks auth login --host https://acme.databricks.test --workspace-id 123", BuildLoginCommand(t.Context(), "", arg)) + + plain, err := u2m.NewProfileWorkspaceOAuthArgument("https://dbc-123.cloud.databricks.test", "") + require.NoError(t, err) + assert.Equal(t, "databricks auth login --host https://dbc-123.cloud.databricks.test", BuildLoginCommand(t.Context(), "", plain)) +} diff --git a/libs/auth/u2m/persistent_auth.go b/libs/auth/u2m/persistent_auth.go index a87e37783f4..7219264a8ec 100644 --- a/libs/auth/u2m/persistent_auth.go +++ b/libs/auth/u2m/persistent_auth.go @@ -625,7 +625,11 @@ func (a *PersistentAuth) oauth2Config() (*oauth2.Config, error) { var err error switch argg := a.oAuthArgument.(type) { case WorkspaceOAuthArgument: - endpoints, err = a.endpointSupplier.GetWorkspaceOAuthEndpoints(a.ctx, argg.GetWorkspaceHost()) + if d, ok := argg.(DiscoveryURLProvider); ok && d.GetDiscoveryURL() != "" { + endpoints, err = a.endpointSupplier.GetEndpointsFromURL(a.ctx, d.GetDiscoveryURL()) + } else { + endpoints, err = a.endpointSupplier.GetWorkspaceOAuthEndpoints(a.ctx, argg.GetWorkspaceHost()) + } case AccountOAuthArgument: endpoints, err = a.endpointSupplier.GetAccountOAuthEndpoints( a.ctx, argg.GetAccountHost(), argg.GetAccountId(), diff --git a/libs/auth/u2m/persistent_auth_test.go b/libs/auth/u2m/persistent_auth_test.go index 831fac8402c..c6744e9c935 100644 --- a/libs/auth/u2m/persistent_auth_test.go +++ b/libs/auth/u2m/persistent_auth_test.go @@ -1583,3 +1583,81 @@ func TestChallenge_Discovery(t *testing.T) { t.Errorf("refresh token = %q, want %q", returnedToken.RefreshToken, "discovery-refresh-token") } } + +// discoveryURLEndpointSupplier serves the endpoints of workspaceHost for one +// discovery URL and records the URLs it was asked for. +type discoveryURLEndpointSupplier struct { + MockOAuthEndpointSupplier + discoveryURL string + workspaceHost string + requested *[]string +} + +func (s discoveryURLEndpointSupplier) GetEndpointsFromURL(_ context.Context, rawURL string) (*OAuthAuthorizationServer, error) { + *s.requested = append(*s.requested, rawURL) + if rawURL != s.discoveryURL { + return nil, ErrOAuthNotSupported + } + return &OAuthAuthorizationServer{ + AuthorizationEndpoint: s.workspaceHost + "/oidc/v1/authorize", + TokenEndpoint: s.workspaceHost + "/oidc/v1/token", + }, nil +} + +func TestToken_WorkspaceArgumentWithDiscoveryURLRefreshesAtDiscoveredEndpoint(t *testing.T) { + const discoveryURL = "https://acme.databricks.test/oidc/.well-known/oauth-authorization-server?o=123" + cache := &tokenStoreMock{ + lookup: func(key string) (*oauth2.Token, error) { + return &oauth2.Token{ + AccessToken: "expired", + RefreshToken: "cde", + Expiry: time.Now().Add(-1 * time.Minute), + }, nil + }, + store: func(key string, tok *oauth2.Token) error { return nil }, + } + arg, err := NewProfileWorkspaceOAuthArgumentWithDiscoveryURL("https://acme.databricks.test", discoveryURL, "my-profile") + if err != nil { + t.Fatalf("NewProfileWorkspaceOAuthArgumentWithDiscoveryURL(): want no error, got %v", err) + } + var requested []string + p, err := NewPersistentAuth( + t.Context(), + WithTokenStore(cache), + WithHttpClient(&http.Client{ + Transport: fixtures.SliceTransport{ + { + Method: "POST", + Resource: "/oidc/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"}, + }, + }, + }, + }), + WithOAuthEndpointSupplier(discoveryURLEndpointSupplier{ + discoveryURL: discoveryURL, + workspaceHost: "https://dbc-123.cloud.databricks.test", + requested: &requested, + }), + WithOAuthArgument(arg), + WithClientID("custom-client-id"), + ) + if err != nil { + t.Fatalf("NewPersistentAuth(): want no error, got %v", err) + } + defer p.Close() + + tok, err := p.Token() + if err != nil { + t.Fatalf("p.Token(): want no error, got %v", err) + } + if tok.AccessToken != "refreshed" { + t.Errorf("p.Token(): want access token 'refreshed', got %s", tok.AccessToken) + } + if len(requested) != 1 || requested[0] != discoveryURL { + t.Errorf("endpoints requested: want [%s], got %v", discoveryURL, requested) + } +} diff --git a/libs/auth/u2m/workspace_oauth_argument.go b/libs/auth/u2m/workspace_oauth_argument.go index 66daba94ebe..6b37fc257a9 100644 --- a/libs/auth/u2m/workspace_oauth_argument.go +++ b/libs/auth/u2m/workspace_oauth_argument.go @@ -14,6 +14,13 @@ type WorkspaceOAuthArgument interface { GetWorkspaceHost() string } +// DiscoveryURLProvider is implemented by OAuth arguments whose OAuth endpoints +// come from an explicit authorization server metadata URL rather than being +// derived from the host. An empty URL means "derive from the host". +type DiscoveryURLProvider interface { + GetDiscoveryURL() string +} + // BasicWorkspaceOAuthArgument is a basic implementation of the WorkspaceOAuthArgument // interface that links each host with exactly one OAuth token. type BasicWorkspaceOAuthArgument struct { @@ -24,6 +31,11 @@ type BasicWorkspaceOAuthArgument struct { // profile is the optional profile name. When set, GetCacheKey() returns // the profile name instead of the host-based key. profile string + + // discoveryURL is the optional authorization server metadata URL. It is + // set when host is a SPOG URL that targets one workspace but OAuth runs + // against that workspace's own endpoints. + discoveryURL string } func validateHost(host string) error { @@ -56,11 +68,29 @@ func NewProfileWorkspaceOAuthArgument(host, profile string) (BasicWorkspaceOAuth return BasicWorkspaceOAuthArgument{host: host, profile: profile}, nil } +// NewProfileWorkspaceOAuthArgumentWithDiscoveryURL creates a +// BasicWorkspaceOAuthArgument whose OAuth endpoints are fetched from +// discoveryURL instead of being derived from host. +func NewProfileWorkspaceOAuthArgumentWithDiscoveryURL(host, discoveryURL, profile string) (BasicWorkspaceOAuthArgument, error) { + arg, err := NewProfileWorkspaceOAuthArgument(host, profile) + if err != nil { + return BasicWorkspaceOAuthArgument{}, err + } + arg.discoveryURL = discoveryURL + return arg, nil +} + // GetWorkspaceHost returns the host of the workspace to authenticate to. func (a BasicWorkspaceOAuthArgument) GetWorkspaceHost() string { return a.host } +// GetDiscoveryURL returns the explicit authorization server metadata URL, or +// an empty string when endpoints are derived from the host. +func (a BasicWorkspaceOAuthArgument) GetDiscoveryURL() string { + return a.discoveryURL +} + // GetCacheKey returns a unique key for caching the OAuth token for the workspace. // If a profile is set, the profile name is returned as the cache key. // Otherwise, the key is in the format "". @@ -68,12 +98,22 @@ func (a BasicWorkspaceOAuthArgument) GetCacheKey() string { if a.profile != "" { return a.profile } - return a.GetHostCacheKey() + return a.hostKey() } // GetHostCacheKey returns the host-based cache key regardless of whether a -// profile is set. The key is in the format "". +// profile is set. The key is in the format "". It is empty when a +// discovery URL is set: the host-keyed mirror only serves old SDKs that look +// tokens up by host, and they would derive this host's account-level OAuth +// endpoints rather than the workspace endpoints the token came from. func (a BasicWorkspaceOAuthArgument) GetHostCacheKey() string { + if a.discoveryURL != "" { + return "" + } + return a.hostKey() +} + +func (a BasicWorkspaceOAuthArgument) hostKey() string { host := strings.TrimSuffix(a.host, "/") if !strings.HasPrefix(host, "http") { host = "https://" + host @@ -84,4 +124,5 @@ func (a BasicWorkspaceOAuthArgument) GetHostCacheKey() string { var ( _ WorkspaceOAuthArgument = BasicWorkspaceOAuthArgument{} _ HostCacheKeyProvider = BasicWorkspaceOAuthArgument{} + _ DiscoveryURLProvider = BasicWorkspaceOAuthArgument{} ) diff --git a/libs/auth/u2m/workspace_oauth_argument_test.go b/libs/auth/u2m/workspace_oauth_argument_test.go index 652f06206e1..495defa810a 100644 --- a/libs/auth/u2m/workspace_oauth_argument_test.go +++ b/libs/auth/u2m/workspace_oauth_argument_test.go @@ -4,6 +4,7 @@ import ( "testing" "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" ) func TestBasicWorkspaceOAuthArgument_GetCacheKey(t *testing.T) { @@ -117,3 +118,18 @@ func TestValidateHost(t *testing.T) { } } } + +func TestBasicWorkspaceOAuthArgument_DiscoveryURL(t *testing.T) { + const discoveryURL = "https://acme.databricks.test/oidc/.well-known/oauth-authorization-server?o=123" + arg, err := NewProfileWorkspaceOAuthArgumentWithDiscoveryURL("https://acme.databricks.test", discoveryURL, "my-profile") + require.NoError(t, err) + + assert.Equal(t, discoveryURL, arg.GetDiscoveryURL()) + assert.Equal(t, "my-profile", arg.GetCacheKey()) + assert.Empty(t, arg.GetHostCacheKey(), "no host-keyed mirror for a workspace token on a SPOG host") + + plain, err := NewProfileWorkspaceOAuthArgument("https://acme.databricks.test", "my-profile") + require.NoError(t, err) + assert.Empty(t, plain.GetDiscoveryURL()) + assert.Equal(t, "https://acme.databricks.test", plain.GetHostCacheKey()) +} diff --git a/libs/databrickscfg/ops.go b/libs/databrickscfg/ops.go index bb066a89edd..1a164f0604b 100644 --- a/libs/databrickscfg/ops.go +++ b/libs/databrickscfg/ops.go @@ -370,6 +370,10 @@ func matchOrCreateSection(ctx context.Context, configFile *config.File, cfg *con // (never read or written) so stale values don't influence routing. const ExperimentalIsUnifiedHostKey = "experimental_is_unified_host" +// DiscoveryURLKey is the INI key for a profile's OAuth authorization server +// metadata URL. +const DiscoveryURLKey = "discovery_url" + // AuthCredentialKeys returns the config file key names for all auth credential // fields from the SDK's ConfigAttributes. These are fields annotated with an // auth type (e.g. pat, basic, oauth, azure, google). Use this to clear stale diff --git a/libs/databrickscfg/profile/file.go b/libs/databrickscfg/profile/file.go index 0844794df3a..3b4f8dba31f 100644 --- a/libs/databrickscfg/profile/file.go +++ b/libs/databrickscfg/profile/file.go @@ -90,6 +90,7 @@ func (f FileProfilerImpl) LoadProfiles(ctx context.Context, fn ProfileMatchFunct Scopes: all["scopes"], Resources: all["resources"], AuthType: all["auth_type"], + DiscoveryURL: all["discovery_url"], } if fn(profile) { profiles = append(profiles, profile) diff --git a/libs/databrickscfg/profile/profile.go b/libs/databrickscfg/profile/profile.go index 3d32989e13d..95de474280f 100644 --- a/libs/databrickscfg/profile/profile.go +++ b/libs/databrickscfg/profile/profile.go @@ -21,6 +21,7 @@ type Profile struct { Scopes string Resources string AuthType string + DiscoveryURL string } func (p Profile) Cloud() string {