diff --git a/.nextchanges/cli/spog-provisioned-url.md b/.nextchanges/cli/spog-provisioned-url.md new file mode 100644 index 00000000000..2a054142548 --- /dev/null +++ b/.nextchanges/cli/spog-provisioned-url.md @@ -0,0 +1 @@ +* `databricks auth login` saves the account's primary (SPOG) URL in the browser-based flow: picking an account saves it directly, and picking a workspace whose account has one runs a second sign-in through it so the profile holds an account-level SPOG token. `--resource` on a SPOG host with a workspace (`?o=` or `--workspace-id`) now logs in at the workspace's own URL, since resource indicators need a workspace-level login. ([#6854](https://github.com/databricks/cli/pull/6854)) diff --git a/cmd/auth/login.go b/cmd/auth/login.go index 2c60557a809..ba99aab19b2 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, resources) if err != nil { return err } @@ -556,6 +561,48 @@ func setHostAndAccountId(ctx context.Context, existingProfile *profile.Profile, return nil } +// setLoginHostAndAccountId is setHostAndAccountId for auth login. RFC 8707 +// resource indicators are only honored on workspace-level authorize, so a +// --resource login for a workspace on a SPOG host logs in at the workspace's +// own host rather than the SPOG host's account-level OIDC. +func setLoginHostAndAccountId(ctx context.Context, existingProfile *profile.Profile, authArguments *auth.AuthArguments, args, resources []string) error { + if err := setHostAndAccountId(ctx, existingProfile, authArguments, args); err != nil { + return err + } + host := (&config.Config{Host: authArguments.Host}).CanonicalHostName() + if len(resources) == 0 || !auth.HasUnifiedHostSignal(authArguments.DiscoveryURL) || auth.IsClassicAccountHost(host) { + return nil + } + if authArguments.WorkspaceID == "" || authArguments.WorkspaceID == auth.WorkspaceIDNone { + return errors.New("--resource on a unified host needs a workspace: add ?o= to --host or pass --workspace-id") + } + workspaceHost, err := lookupWorkspaceHost(ctx, host, authArguments.WorkspaceID) + if err != nil { + return fmt.Errorf("finding the host of workspace %s for --resource: %w. Pass the workspace URL with --host instead", authArguments.WorkspaceID, err) + } + cmdio.LogString(ctx, fmt.Sprintf("Logging in at %s, the workspace's own URL, because --resource needs a workspace-level login.", workspaceHost)) + authArguments.Host = workspaceHost + authArguments.DiscoveryURL = "" + return nil +} + +// lookupWorkspaceHost returns the canonical host of workspaceID, read from the +// token endpoint of the workspace OAuth metadata the SPOG host serves for it. +func lookupWorkspaceHost(ctx context.Context, spogHost, workspaceID string) (string, error) { + discoveryURL := spogHost + "/oidc/.well-known/oauth-authorization-server?" + url.Values{"o": {workspaceID}}.Encode() + lookupCtx, cancel := context.WithTimeout(ctx, provisionedURLTimeout) + defer cancel() + tokenEndpoint, err := auth.LookupOAuthTokenEndpoint(lookupCtx, discoveryURL, nil) + if err != nil { + return "", err + } + u, err := url.Parse(tokenEndpoint) + if err != nil || u.Scheme == "" || u.Host == "" { + return "", fmt.Errorf("invalid token endpoint %q", tokenEndpoint) + } + return u.Scheme + "://" + u.Host, nil +} + // 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 +618,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 +731,95 @@ 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) +} + +// switchToWorkspacePrimaryURL replaces a workspace host with the primary +// (SPOG) URL of its owning account, keeping account_id and workspace_id so the +// profile targets the same workspace. The host is only switched once +// discovery on the primary URL confirms it is a unified host. Best-effort: +// failures are logged and leave the host unchanged. +func switchToWorkspacePrimaryURL(ctx context.Context, authArguments *auth.AuthArguments) { + if !shouldResolveWorkspacePrimaryURL(authArguments) { + return + } + lookupCtx, cancel := context.WithTimeout(ctx, provisionedURLTimeout) + defer cancel() + resp, err := auth.LookupWorkspacePrimaryURL(lookupCtx, authArguments.Host, nil) + if err != nil { + log.Debugf(ctx, "Workspace primary URL lookup failed: %v", err) + return + } + primaryURL := strings.TrimSuffix(resp.PrimaryURL, "/") + if primaryURL == "" || + (&config.Config{Host: primaryURL}).CanonicalHostName() == (&config.Config{Host: authArguments.Host}).CanonicalHostName() { + return + } + + // 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, authArguments.Host) + return + } + authArguments.Host = primaryURL + authArguments.WorkspaceID = workspaceID + authArguments.DiscoveryURL = spogArgs.DiscoveryURL +} + // discoveryLoginInputs groups the dependencies of discoveryLogin. // See https://google.github.io/styleguide/go/best-practices#option-structure. type discoveryLoginInputs struct { @@ -688,11 +832,18 @@ 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. +// When a workspace whose account has a primary URL is selected, the CLI logs in +// again through that URL, so the profile holds an account-level SPOG token. func discoveryLogin(ctx context.Context, in discoveryLoginInputs) error { arg, err := in.dc.NewOAuthArgument(in.profileName) if err != nil { @@ -774,13 +925,29 @@ 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) + } + + var tokenArg u2m.OAuthArgument = arg + if spogHost, spogWorkspaceID, spogArg, spogTok, ok := loginAtWorkspacePrimaryURL(ctx, in, discoveredHost, accountID, workspaceID, scopesList); ok { + discoveredHost, workspaceID, tokenArg, tok = spogHost, spogWorkspaceID, spogArg, spogTok + } + 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", @@ -805,7 +972,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 } @@ -813,6 +980,56 @@ func discoveryLogin(ctx context.Context, in discoveryLoginInputs) error { return nil } +// loginAtWorkspacePrimaryURL logs in again through the primary (SPOG) URL of +// the account that owns the workspace selected in the discovery flow. A second +// login is needed because the discovery token was issued by the workspace's +// OIDC, which the SPOG host's account-level OIDC won't refresh. Returns the +// SPOG host, workspace ID, and the OAuth argument and token to save, or +// ok=false to keep the workspace login. Best-effort: failures are logged and +// never block login. +func loginAtWorkspacePrimaryURL(ctx context.Context, in discoveryLoginInputs, host, accountID, workspaceID string, scopesList []string) (string, string, u2m.OAuthArgument, *oauth2.Token, bool) { + authArguments := &auth.AuthArguments{ + Host: host, + AccountID: accountID, + WorkspaceID: workspaceID, + Profile: in.profileName, + } + switchToWorkspacePrimaryURL(ctx, authArguments) + if authArguments.Host == host { + return "", "", nil, nil, false + } + + oauthArg, err := authArguments.ToOAuthArgument() + if err != nil { + log.Warnf(ctx, "Setting up login at primary URL %s failed, keeping %s: %v", authArguments.Host, host, err) + return "", "", nil, nil, false + } + opts := []u2m.PersistentAuthOption{ + u2m.WithOAuthArgument(oauthArg), + u2m.WithBrowser(in.browserFunc), + } + if in.clientID != "" { + opts = append(opts, u2m.WithClientID(in.clientID)) + } + if len(scopesList) > 0 { + opts = append(opts, u2m.WithScopes(scopesList)) + } + persistentAuth, err := in.dc.NewPersistentAuth(ctx, opts...) + if err != nil { + log.Warnf(ctx, "Setting up login at primary URL %s failed, keeping %s: %v", authArguments.Host, host, err) + return "", "", nil, nil, false + } + defer persistentAuth.Close() + + cmdio.LogString(ctx, fmt.Sprintf("Opening %s, your account's primary URL, in your browser...", authArguments.Host)) + tok, err := persistentAuth.Challenge() + if err != nil { + log.Warnf(ctx, "Login at primary URL %s failed, keeping %s: %v", authArguments.Host, host, err) + return "", "", nil, nil, false + } + return authArguments.Host, authArguments.WorkspaceID, oauthArg, tok, true +} + // splitScopes splits a comma-separated scopes string into a trimmed slice. func splitScopes(scopes string) []string { var result []string diff --git a/cmd/auth/login_test.go b/cmd/auth/login_test.go index 1c4428596dd..a37245b640b 100644 --- a/cmd/auth/login_test.go +++ b/cmd/auth/login_test.go @@ -8,8 +8,10 @@ import ( "log/slog" "net/http" "net/http/httptest" + "net/url" "os" "path/filepath" + "strings" "sync" "testing" "time" @@ -113,11 +115,18 @@ type fakeDiscoveryClient struct { oauthArgErr error persistentAuth discoveryPersistentAuth persistentAuthErr error - introspection *auth.IntrospectionResult - introspectionErr error + // followUpAuths, when set, serve NewPersistentAuth calls after the first + // (e.g. the second login at a workspace's primary URL), in order. + followUpAuths []discoveryPersistentAuth + // newFollowUpAuth, when set, serves NewPersistentAuth calls after the + // first and receives their options. + newFollowUpAuth func(ctx context.Context, opts ...u2m.PersistentAuthOption) (discoveryPersistentAuth, error) + 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 +140,15 @@ func (f *fakeDiscoveryClient) NewPersistentAuth(ctx context.Context, opts ...u2m if f.persistentAuthErr != nil { return nil, f.persistentAuthErr } + f.newPersistentAuthCall++ + if f.newPersistentAuthCall > 1 && f.newFollowUpAuth != nil { + return f.newFollowUpAuth(ctx, opts...) + } + if f.newPersistentAuthCall > 1 && len(f.followUpAuths) > 0 { + next := f.followUpAuths[0] + f.followUpAuths = f.followUpAuths[1:] + return next, nil + } return f.persistentAuth, nil } @@ -1247,6 +1265,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 +1652,472 @@ 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}, + {"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 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 TestDiscoveryLogin_WorkspaceSelectionLogsInAtPrimaryURL(t *testing.T) { + spog := newUnifiedHostServer(t) + workspace := newDiscoveryServer(t, map[string]any{ + "account_id": "spog-account", + "workspace_id": "12345", + "primary_url": spog.URL, + }) + + 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"}}, + followUpAuths: []discoveryPersistentAuth{&fakeDiscoveryPersistentAuth{token: &oauth2.Token{AccessToken: "spog-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) + + assert.Equal(t, 2, dc.newPersistentAuthCall, "expected a second login at the primary URL") + 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) + require.Contains(t, store.Tokens, "DISCOVERY") + assert.Equal(t, "spog-token", store.Tokens["DISCOVERY"].AccessToken) +} + +func TestDiscoveryLogin_PrimaryURLLoginFailureKeepsWorkspace(t *testing.T) { + spog := newUnifiedHostServer(t) + workspace := newDiscoveryServer(t, map[string]any{ + "account_id": "spog-account", + "workspace_id": "12345", + "primary_url": spog.URL, + }) + + 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"}}, + followUpAuths: []discoveryPersistentAuth{&fakeDiscoveryPersistentAuth{challengeErr: errors.New("browser closed")}}, + 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, workspace.URL, savedProfile.Host) + assert.Equal(t, "12345", savedProfile.WorkspaceID) + require.Contains(t, store.Tokens, "DISCOVERY") + assert.Equal(t, "workspace-token", store.Tokens["DISCOVERY"].AccessToken) +} + +type unifiedPathEndpointSupplier struct { + MockApiClient +} + +func (*unifiedPathEndpointSupplier) GetUnifiedOAuthEndpoints(_ context.Context, host, accountID string) (*u2m.OAuthAuthorizationServer, error) { + return &u2m.OAuthAuthorizationServer{ + AuthorizationEndpoint: host + "/oidc/accounts/" + accountID + "/v1/authorize", + TokenEndpoint: host + "/oidc/accounts/" + accountID + "/v1/token", + }, nil +} + +func declineInBrowser(t *testing.T, authorizeURL *string) func(string) error { + return func(rawURL string) error { + *authorizeURL = rawURL + u, err := url.Parse(rawURL) + require.NoError(t, err) + q := u.Query() + callback := q.Get("redirect_uri") + "?" + url.Values{ + "error": {"access_denied"}, + "error_description": {"declined in test"}, + "state": {q.Get("state")}, + }.Encode() + go func() { + resp, err := http.Get(callback) + if err == nil { + resp.Body.Close() + } + }() + return nil + } +} + +func hostOfURL(t *testing.T, rawURL string) string { + t.Helper() + u, err := url.Parse(rawURL) + require.NoError(t, err) + return u.Host +} + +func TestDiscoveryLogin_PrimaryURLLoginUsesAccountLevelOAuthAndLoginOptions(t *testing.T) { + spog := newUnifiedHostServer(t) + workspace := newDiscoveryServer(t, map[string]any{ + "account_id": "spog-account", + "workspace_id": "12345", + "primary_url": spog.URL, + }) + + 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) + + // The second login runs a real PersistentAuth with the options discoveryLogin + // passes, so the authorize URL shows which OAuth endpoints, client ID and + // scopes it would use. + var authorizeURL string + dc := &fakeDiscoveryClient{ + oauthArg: oauthArg, + persistentAuth: &fakeDiscoveryPersistentAuth{token: &oauth2.Token{AccessToken: "workspace-token"}}, + newFollowUpAuth: func(ctx context.Context, opts ...u2m.PersistentAuthOption) (discoveryPersistentAuth, error) { + return u2m.NewPersistentAuth(ctx, append(opts, u2m.WithOAuthEndpointSupplier(&unifiedPathEndpointSupplier{}))...) + }, + introspection: &auth.IntrospectionResult{}, + } + + ctx, _ := cmdio.NewTestContextWithStdout(t.Context()) + err = discoveryLogin(ctx, discoveryLoginInputs{ + dc: dc, + profileName: "DISCOVERY", + timeout: 10 * time.Second, + scopes: "sql,jobs", + clientID: "custom-client", + browserFunc: declineInBrowser(t, &authorizeURL), + tokenStore: newTestStore(), + }) + require.NoError(t, err) + + require.NotEmpty(t, authorizeURL, "the second login should open the browser") + u, err := url.Parse(authorizeURL) + require.NoError(t, err) + assert.Equal(t, hostOfURL(t, spog.URL), u.Host) + assert.Equal(t, "/oidc/accounts/spog-account/v1/authorize", u.Path) + assert.Equal(t, "custom-client", u.Query().Get("client_id")) + scopes := strings.Fields(u.Query().Get("scope")) + assert.Contains(t, scopes, "sql") + assert.Contains(t, scopes, "jobs") + + // Declining the second login keeps the workspace profile. + savedProfile, err := loadProfileByName(ctx, "DISCOVERY", profile.DefaultProfiler) + require.NoError(t, err) + require.NotNil(t, savedProfile) + assert.Equal(t, workspace.URL, savedProfile.Host) +} + +func TestDiscoveryLogin_PrimaryURLLoginSetupFailureKeepsWorkspace(t *testing.T) { + tests := []struct { + name string + // primaryURL maps the SPOG server's URL to the primary_url the workspace reports. + primaryURL func(spogURL string) string + followUpErr error + wantFollowUpRun bool + }{ + { + name: "second login can't be set up", + primaryURL: func(spogURL string) string { return spogURL }, + followUpErr: errors.New("setup failed"), + wantFollowUpRun: true, + }, + { + // OAuth arguments only accept http for 127.0.0.1, so a localhost + // primary URL passes discovery but fails ToOAuthArgument. + name: "oauth argument for the primary URL can't be built", + primaryURL: func(spogURL string) string { return strings.Replace(spogURL, "127.0.0.1", "localhost", 1) }, + }, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + spog := newUnifiedHostServer(t) + workspace := newDiscoveryServer(t, map[string]any{ + "account_id": "spog-account", + "workspace_id": "12345", + "primary_url": tt.primaryURL(spog.URL), + }) + + 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) + + followUpRun := false + dc := &fakeDiscoveryClient{ + oauthArg: oauthArg, + persistentAuth: &fakeDiscoveryPersistentAuth{token: &oauth2.Token{AccessToken: "workspace-token"}}, + newFollowUpAuth: func(context.Context, ...u2m.PersistentAuthOption) (discoveryPersistentAuth, error) { + followUpRun = true + return nil, tt.followUpErr + }, + 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) + + assert.Equal(t, tt.wantFollowUpRun, followUpRun) + savedProfile, err := loadProfileByName(ctx, "DISCOVERY", profile.DefaultProfiler) + require.NoError(t, err) + require.NotNil(t, savedProfile) + assert.Equal(t, workspace.URL, savedProfile.Host) + require.Contains(t, store.Tokens, "DISCOVERY") + assert.Equal(t, "workspace-token", store.Tokens["DISCOVERY"].AccessToken) + }) + } +} + +func TestSwitchToWorkspacePrimaryURL(t *testing.T) { + spog := newUnifiedHostServer(t) + notUnified := newDiscoveryServer(t, map[string]any{"workspace_id": "999"}) + lookupFails := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.WriteHeader(http.StatusInternalServerError) + })) + t.Cleanup(lookupFails.Close) + + tests := []struct { + name string + host string + workspaceID string + wantSwitched bool + wantWorkspaceID string + }{ + { + name: "switches to the primary url", + host: newDiscoveryServer(t, map[string]any{"account_id": "spog-account", "workspace_id": "12345", "primary_url": spog.URL + "/"}).URL, + workspaceID: "12345", + wantSwitched: true, + wantWorkspaceID: "12345", + }, + { + // A workspace_id inherited from an existing profile can belong to + // another workspace; on SPOG it decides routing. + name: "uses the host's workspace id over an inherited one", + host: newDiscoveryServer(t, map[string]any{"account_id": "spog-account", "workspace_id": "12345", "primary_url": spog.URL}).URL, + workspaceID: "67890", + wantSwitched: true, + wantWorkspaceID: "12345", + }, + { + name: "no primary url", + host: newDiscoveryServer(t, map[string]any{"account_id": "spog-account", "workspace_id": "12345"}).URL, + workspaceID: "12345", + wantWorkspaceID: "12345", + }, + { + name: "primary url is not a unified host", + host: newDiscoveryServer(t, map[string]any{"account_id": "spog-account", "workspace_id": "12345", "primary_url": notUnified.URL}).URL, + workspaceID: "12345", + wantWorkspaceID: "12345", + }, + { + name: "lookup fails", + host: lookupFails.URL, + workspaceID: "12345", + wantWorkspaceID: "12345", + }, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + args := &auth.AuthArguments{Host: tt.host, AccountID: "spog-account", WorkspaceID: tt.workspaceID} + switchToWorkspacePrimaryURL(t.Context(), args) + + if tt.wantSwitched { + assert.Equal(t, spog.URL, args.Host) + assert.True(t, auth.HasUnifiedHostSignal(args.DiscoveryURL), "discovery URL %q", args.DiscoveryURL) + } else { + assert.Equal(t, tt.host, args.Host) + } + assert.Equal(t, "spog-account", args.AccountID) + assert.Equal(t, tt.wantWorkspaceID, args.WorkspaceID) + }) + } +} + +func TestSetLoginHostAndAccountId_Resources(t *testing.T) { + workspace, spog := newSpogWorkspacePair(t) + // A unified host that serves no workspace-level OAuth metadata. + spogWithoutWorkspaceOAuth := newUnifiedHostServer(t) + resources := []string{workspace.URL + "/ai-gateway/mcp-services/system.ai.github"} + + tests := []struct { + name string + host string + workspaceID string + resources []string + wantHost string + wantErr string + }{ + {name: "spog host with ?o= logs in at the workspace host", host: spog.URL + "?o=12345", resources: resources, wantHost: workspace.URL}, + {name: "spog host with --workspace-id logs in at the workspace host", host: spog.URL, workspaceID: "12345", resources: resources, wantHost: workspace.URL}, + {name: "spog host without a workspace", host: spog.URL, resources: resources, wantErr: "--resource on a unified host needs a workspace"}, + {name: "spog host that doesn't serve the workspace", host: spogWithoutWorkspaceOAuth.URL + "?o=12345", resources: resources, wantErr: "Pass the workspace URL with --host instead"}, + {name: "workspace host is kept", host: workspace.URL, resources: resources, wantHost: workspace.URL}, + {name: "spog host without resources is kept", host: spog.URL + "?o=12345", wantHost: spog.URL}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + args := &auth.AuthArguments{Host: tt.host, WorkspaceID: tt.workspaceID} + err := setLoginHostAndAccountId(cmdio.MockDiscard(t.Context()), nil, args, []string{}, tt.resources) + if tt.wantErr != "" { + assert.ErrorContains(t, err, tt.wantErr) + return + } + require.NoError(t, err) + assert.Equal(t, tt.wantHost, args.Host) + }) + } +} diff --git a/libs/auth/provisioned_url.go b/libs/auth/provisioned_url.go new file mode 100644 index 00000000000..4128e302f68 --- /dev/null +++ b/libs/auth/provisioned_url.go @@ -0,0 +1,118 @@ +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 +} + +// 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/provisioned_url_test.go b/libs/auth/provisioned_url_test.go new file mode 100644 index 00000000000..c0e996596a4 --- /dev/null +++ b/libs/auth/provisioned_url_test.go @@ -0,0 +1,131 @@ +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") +} + +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(), server.URL+"/oidc/.well-known/oauth-authorization-server?o=123", nil) + require.NoError(t, err) + assert.Equal(t, "https://dbc-123.cloud.databricks.test/oidc/v1/token", endpoint) +}