diff --git a/cmd/auth/docker_profile.go b/cmd/auth/docker_profile.go new file mode 100644 index 00000000000..8c3a7a0798b --- /dev/null +++ b/cmd/auth/docker_profile.go @@ -0,0 +1,35 @@ +package auth + +import ( + "fmt" + + authlib "github.com/databricks/cli/libs/auth" + "github.com/databricks/cli/libs/databrickscfg/profile" + "github.com/databricks/databricks-sdk-go/config" +) + +// validateDockerCredentialProfile requires metadata eligible for workspace-scoped U2M authentication. +func validateDockerCredentialProfile(p profile.Profile) error { + if p.HasClientCredentials { + return fmt.Errorf("profile %q uses client credentials. Docker credential helper requires a profile created by databricks auth login", p.Name) + } + if p.AuthType != authTypeDatabricksCLI { + return fmt.Errorf("profile %q uses auth_type %q. Docker credential helper requires a profile created by databricks auth login", p.Name, p.AuthType) + } + if isDockerCredentialAccountOnlyProfile(p) { + return fmt.Errorf("profile %q does not target a workspace. Run databricks auth login --host and retry with that profile", p.Name) + } + return nil +} + +// isDockerCredentialAccountOnlyProfile treats classic account hosts and unrouted account profiles as unsafe for workspace requests. +func isDockerCredentialAccountOnlyProfile(p profile.Profile) bool { + if p.Host == "" { + return true + } + cfg := &config.Config{Host: p.Host, AccountID: p.AccountID, WorkspaceID: p.WorkspaceID} + if authlib.IsClassicAccountHost(cfg.CanonicalHostName()) { + return true + } + return p.AccountID != "" && (p.WorkspaceID == "" || p.WorkspaceID == authlib.WorkspaceIDNone) +} diff --git a/cmd/auth/token.go b/cmd/auth/token.go index d5e88e64d72..782c9329e31 100644 --- a/cmd/auth/token.go +++ b/cmd/auth/token.go @@ -16,6 +16,7 @@ import ( "github.com/databricks/cli/libs/cmdio" "github.com/databricks/cli/libs/databrickscfg" "github.com/databricks/cli/libs/databrickscfg/profile" + "github.com/databricks/cli/libs/dockercredentials" "github.com/databricks/cli/libs/env" "github.com/databricks/cli/libs/flags" "github.com/databricks/cli/libs/log" @@ -31,7 +32,22 @@ func helpfulError(ctx context.Context, profile string, persistentAuth u2m.OAuthA return fmt.Sprintf("Try logging in again with `%s` before retrying. If this fails, please report this issue to the Databricks CLI maintainers at https://github.com/databricks/cli/issues/new", loginMsg) } +type ( + tokenLoader func(context.Context, loadTokenArgs) (*oauth2.Token, error) + registryHostResolver func(string, string, string) (string, error) +) + func newTokenCommand(authArguments *auth.AuthArguments) *cobra.Command { + return newTokenCommandWithLoader(authArguments, loadToken) +} + +// newTokenCommandWithLoader isolates cache-backed token acquisition from command parsing in tests. +func newTokenCommandWithLoader(authArguments *auth.AuthArguments, load tokenLoader) *cobra.Command { + return newTokenCommandWithRegistryHost(authArguments, load, dockercredentials.RegistryHost) +} + +// newTokenCommandWithRegistryHost isolates registry reconstruction so tests can model cloud and environment matches. +func newTokenCommandWithRegistryHost(authArguments *auth.AuthArguments, load tokenLoader, registryHost registryHostResolver) *cobra.Command { cmd := &cobra.Command{ Use: "token [PROFILE]", Short: "Get authentication token", @@ -50,18 +66,31 @@ and secret is not supported.`, cmd.Flags().BoolVar(&forceRefresh, "force-refresh", false, "Force a token refresh even if the cached token is still valid.") - cmd.PreRunE = profileHostConflictCheck + // Docker format is an internal credential-helper contract, not a user-facing output mode. + var format string + cmd.Flags().StringVar(&format, "format", "", "Hidden output format") + _ = cmd.Flags().MarkHidden("format") + + cmd.PreRunE = func(cmd *cobra.Command, args []string) error { + if format == "docker" { + return validateDockerTokenRequest(cmd, args) + } + return profileHostConflictCheck(cmd, args) + } cmd.RunE = func(cmd *cobra.Command, args []string) error { ctx := cmd.Context() profileName := cmd.Flag("profile").Value.String() + if format != "" && format != "docker" { + return fmt.Errorf("unsupported token format %q", format) + } tokenStore, mode, err := storage.ResolveStore(ctx, "") if err != nil { return err } - t, err := loadToken(ctx, loadTokenArgs{ + loadArgs := loadTokenArgs{ authArguments: authArguments, profileName: profileName, args: args, @@ -71,7 +100,13 @@ and secret is not supported.`, tokenStore: tokenStore, mode: mode, persistentAuthOpts: nil, - }) + } + + if format == "docker" { + return writeDockerTokenOutput(ctx, cmd, loadArgs, load, registryHost) + } + + t, err := load(ctx, loadArgs) if err != nil { return err } @@ -85,6 +120,108 @@ and secret is not supported.`, return cmd } +type dockerGetResponse struct { + Username string `json:"Username"` + Secret string `json:"Secret"` +} + +// writeDockerTokenOutput makes Docker's registry address the sole profile selector for its get response. +// See https://docs.docker.com/reference/cli/docker/login/#credential-helper-protocol. +func writeDockerTokenOutput(ctx context.Context, cmd *cobra.Command, args loadTokenArgs, load tokenLoader, registryHost registryHostResolver) error { + rawServer, err := io.ReadAll(cmd.InOrStdin()) + if err != nil { + return fmt.Errorf("read Docker credential request: %w", err) + } + registry, err := dockercredentials.ParseRegistryHost(string(rawServer)) + if err != nil { + return err + } + + profileName, err := dockerTokenProfileName(ctx, registry, args.profiler, registryHost) + if err != nil { + return err + } + + args.authArguments = &auth.AuthArguments{} + args.profileName = profileName + args.args = nil + + t, err := load(ctx, args) + if err != nil { + return err + } + + return json.NewEncoder(cmd.OutOrStdout()).Encode(dockerGetResponse{ + Username: dockercredentials.OAuthTokenUsername, + Secret: t.AccessToken, + }) +} + +// validateDockerTokenRequest makes the registry address the sole profile selector for Docker-format requests. +func validateDockerTokenRequest(cmd *cobra.Command, args []string) error { + if len(args) > 0 { + return errors.New("--format=docker does not accept positional arguments") + } + for _, name := range []string{"profile", "host", "account-id", "workspace-id"} { + flag := cmd.Flag(name) + if flag != nil && flag.Changed { + return fmt.Errorf("--format=docker does not support --%s", name) + } + } + + return nil +} + +// dockerTokenProfileName finds one compatible profile, using the registry DNS zone to disambiguate repeated workspace IDs. +func dockerTokenProfileName(ctx context.Context, registry dockercredentials.Registry, profiler profile.Profiler, registryHost registryHostResolver) (string, error) { + workspaceProfiles, err := profiler.LoadProfiles(ctx, func(p profile.Profile) bool { + return p.WorkspaceID == registry.WorkspaceID + }) + if err != nil { + return "", err + } + if len(workspaceProfiles) == 0 { + return "", fmt.Errorf("no Databricks profile found for workspace ID %s from registry host %s. Run databricks auth login --host and set workspace_id for that profile", registry.WorkspaceID, registry.Host) + } + if len(workspaceProfiles) == 1 { + return validateDockerTokenProfile(registry, workspaceProfiles[0], registryHost) + } + + var matchingProfiles profile.Profiles + for _, p := range workspaceProfiles { + if validateDockerCredentialProfile(p) != nil { + continue + } + expectedHost, err := registryHost(registry.WorkspaceID, registry.Region, p.Host) + if err == nil && expectedHost == registry.Host { + matchingProfiles = append(matchingProfiles, p) + } + } + if len(matchingProfiles) == 0 { + return "", fmt.Errorf("registry host %s does not match any profile for workspace ID %s. Verify the profile workspace host and workspace_id", registry.Host, registry.WorkspaceID) + } + if len(matchingProfiles) > 1 { + return "", fmt.Errorf("multiple Databricks profiles match workspace ID %s: %s. Remove duplicate workspace_id entries before using Docker credential helper", registry.WorkspaceID, strings.Join(matchingProfiles.Names(), " and ")) + } + return validateDockerTokenProfile(registry, matchingProfiles[0], registryHost) +} + +// validateDockerTokenProfile verifies U2M eligibility and reconstructs the registry host to prevent cross-environment matches. +func validateDockerTokenProfile(registry dockercredentials.Registry, p profile.Profile, registryHost registryHostResolver) (string, error) { + if err := validateDockerCredentialProfile(p); err != nil { + return "", err + } + // Workspace IDs can repeat across environments, so the registry must also match the profile's DNS zone. + expectedHost, err := registryHost(registry.WorkspaceID, registry.Region, p.Host) + if err != nil { + return "", fmt.Errorf("validate registry host against profile %q: %w", p.Name, err) + } + if expectedHost != registry.Host { + return "", fmt.Errorf("registry host %s does not match profile %q workspace host", registry.Host, p.Name) + } + return p.Name, nil +} + func writeTokenOutput(w io.Writer, t *oauth2.Token, textMode bool) error { if textMode { _, err := fmt.Fprintln(w, t.AccessToken) diff --git a/cmd/auth/token_test.go b/cmd/auth/token_test.go index adda6888a40..1aa775d7530 100644 --- a/cmd/auth/token_test.go +++ b/cmd/auth/token_test.go @@ -6,17 +6,25 @@ import ( "encoding/json" "errors" "net/http" + "os" + "path/filepath" + "strings" "testing" "time" "github.com/databricks/cli/libs/auth" "github.com/databricks/cli/libs/auth/storage" "github.com/databricks/cli/libs/cmdio" + "github.com/databricks/cli/libs/databrickscfg" "github.com/databricks/cli/libs/databrickscfg/profile" + "github.com/databricks/cli/libs/dockercredentials" "github.com/databricks/cli/libs/env" + "github.com/databricks/databricks-sdk-go/config" "github.com/databricks/databricks-sdk-go/credentials/u2m" "github.com/databricks/databricks-sdk-go/httpclient/fixtures" + "github.com/spf13/cobra" "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" "golang.org/x/oauth2" ) @@ -885,6 +893,396 @@ func (e errProfiler) GetPath(context.Context) (string, error) { return "", nil } +func TestTokenDockerFormatEmitsGetResponse(t *testing.T) { + ctx := cmdio.MockDiscard(t.Context()) + dir := t.TempDir() + configFile := filepath.Join(dir, ".databrickscfg") + require.NoError(t, databrickscfg.SaveToProfile(ctx, &config.Config{ + ConfigFile: configFile, + Profile: "workspace", + Host: "https://workspace.cloud.databricks.test", + WorkspaceID: "123456789", + AuthType: authTypeDatabricksCLI, + })) + + t.Setenv("DATABRICKS_CONFIG_FILE", configFile) + t.Setenv(storage.EnvVar, string(storage.StorageModePlaintext)) + t.Setenv("HOME", dir) + + var gotProfile string + loadToken := func(_ context.Context, args loadTokenArgs) (*oauth2.Token, error) { + gotProfile = args.profileName + return &oauth2.Token{AccessToken: "access-token"}, nil + } + + registryHost := "123456789.container.us-west-2.cloud.databricks.com" + var stdout bytes.Buffer + cmd := newTokenCommandWithRegistryHost(&auth.AuthArguments{}, loadToken, func(workspaceID, region, workspaceHost string) (string, error) { + require.Equal(t, "123456789", workspaceID) + require.Equal(t, "us-west-2", region) + require.Equal(t, "https://workspace.cloud.databricks.test", workspaceHost) + return registryHost, nil + }) + cmd.Flags().StringP("profile", "p", "", "~/.databrickscfg profile") + cmd.SetContext(ctx) + cmd.SetIn(strings.NewReader(registryHost + "\n")) + cmd.SetOut(&stdout) + cmd.SetArgs([]string{"--format=docker"}) + + require.NoError(t, cmd.Execute()) + require.Equal(t, "workspace", gotProfile) + + var got map[string]string + require.NoError(t, json.Unmarshal(stdout.Bytes(), &got)) + require.Equal(t, map[string]string{ + "Username": "oauthtoken", + "Secret": "access-token", + }, got) +} + +func TestWriteDockerTokenOutputUsesConfiguredProfiler(t *testing.T) { + ctx := cmdio.MockDiscard(t.Context()) + profiler := profile.InMemoryProfiler{ + Profiles: profile.Profiles{ + { + Name: "workspace", + Host: "https://workspace.cloud.databricks.test", + WorkspaceID: "123456789", + AuthType: authTypeDatabricksCLI, + }, + }, + } + + var gotProfile string + loadToken := func(_ context.Context, args loadTokenArgs) (*oauth2.Token, error) { + gotProfile = args.profileName + return &oauth2.Token{AccessToken: "access-token"}, nil + } + + cmd := &cobra.Command{Use: "token"} + var stdout bytes.Buffer + cmd.SetContext(ctx) + cmd.SetIn(strings.NewReader("123456789.container.us-west-2.cloud.databricks.com\n")) + cmd.SetOut(&stdout) + + err := writeDockerTokenOutput(ctx, cmd, loadTokenArgs{ + authArguments: &auth.AuthArguments{}, + profiler: profiler, + }, loadToken, func(string, string, string) (string, error) { + return "123456789.container.us-west-2.cloud.databricks.com", nil + }) + require.NoError(t, err) + require.Equal(t, "workspace", gotProfile) + + var got dockerGetResponse + require.NoError(t, json.Unmarshal(stdout.Bytes(), &got)) +} + +func TestDockerTokenProfileNameRejectsDifferentEnvironment(t *testing.T) { + registry := dockercredentials.Registry{ + WorkspaceID: "123456789", + Region: "us-west-2", + Host: "123456789.container.us-west-2.cloud.databricks.com", + } + profiler := profile.InMemoryProfiler{ + Profiles: profile.Profiles{{ + Name: "workspace", + Host: "https://workspace.dev.cloud.databricks.test", + WorkspaceID: registry.WorkspaceID, + AuthType: authTypeDatabricksCLI, + }}, + } + registryHost := func(workspaceID, region, workspaceHost string) (string, error) { + require.Equal(t, registry.WorkspaceID, workspaceID) + require.Equal(t, registry.Region, region) + require.Equal(t, "https://workspace.dev.cloud.databricks.test", workspaceHost) + return "123456789.container.us-west-2.dev.cloud.databricks.com", nil + } + + _, err := dockerTokenProfileName(t.Context(), registry, profiler, registryHost) + require.ErrorContains(t, err, "does not match profile") + require.ErrorContains(t, err, "workspace host") +} + +func TestDockerTokenProfileNameAllowsSameWorkspaceIDInDifferentEnvironment(t *testing.T) { + registry := dockercredentials.Registry{ + WorkspaceID: "123456789", + Region: "us-west-2", + Host: "123456789.container.us-west-2.cloud.databricks.com", + } + profiler := profile.InMemoryProfiler{ + Profiles: profile.Profiles{ + { + Name: "prod", + Host: "https://workspace.cloud.databricks.com", + WorkspaceID: registry.WorkspaceID, + AuthType: authTypeDatabricksCLI, + }, + { + Name: "dev", + Host: "https://workspace.dev.cloud.databricks.com", + WorkspaceID: registry.WorkspaceID, + AuthType: authTypeDatabricksCLI, + }, + }, + } + registryHost := func(workspaceID, region, workspaceHost string) (string, error) { + zone := ".cloud.databricks.com" + if workspaceHost == "https://workspace.dev.cloud.databricks.com" { + zone = ".dev.cloud.databricks.com" + } + return workspaceID + ".container." + region + zone, nil + } + + profileName, err := dockerTokenProfileName(t.Context(), registry, profiler, registryHost) + require.NoError(t, err) + require.Equal(t, "prod", profileName) +} + +func TestDockerTokenProfileNameIgnoresUnsupportedDuplicateProfile(t *testing.T) { + registry := dockercredentials.Registry{ + WorkspaceID: "123456789", + Region: "us-west-2", + Host: "123456789.container.us-west-2.cloud.databricks.com", + } + profiler := profile.InMemoryProfiler{ + Profiles: profile.Profiles{ + { + Name: "workspace", + Host: "https://workspace.cloud.databricks.com", + WorkspaceID: registry.WorkspaceID, + AuthType: authTypeDatabricksCLI, + }, + { + Name: "m2m", + Host: "https://workspace.cloud.databricks.com", + WorkspaceID: registry.WorkspaceID, + HasClientCredentials: true, + }, + }, + } + registryHost := func(workspaceID, region, _ string) (string, error) { + return workspaceID + ".container." + region + ".cloud.databricks.com", nil + } + + profileName, err := dockerTokenProfileName(t.Context(), registry, profiler, registryHost) + require.NoError(t, err) + require.Equal(t, "workspace", profileName) +} + +func TestTokenDockerFormatRejectsPositionalArgs(t *testing.T) { + ctx := cmdio.MockDiscard(t.Context()) + t.Setenv(storage.EnvVar, string(storage.StorageModePlaintext)) + t.Setenv("HOME", t.TempDir()) + + cmd := newTokenCommandWithLoader(&auth.AuthArguments{}, func(context.Context, loadTokenArgs) (*oauth2.Token, error) { + t.Fatal("loadToken should not be called") + return nil, nil + }) + cmd.Flags().StringP("profile", "p", "", "~/.databrickscfg profile") + cmd.SetContext(ctx) + cmd.SetIn(strings.NewReader("123456789.container.us-west-2.cloud.databricks.com\n")) + cmd.SetArgs([]string{"--format=docker", "DEFAULT"}) + + err := cmd.Execute() + require.ErrorContains(t, err, "--format=docker does not accept positional arguments") +} + +func TestTokenDockerFormatValidatesBeforeResolvingTokenStore(t *testing.T) { + ctx := cmdio.MockDiscard(t.Context()) + t.Setenv(storage.EnvVar, "invalid") + + cmd := newTokenCommandWithLoader(&auth.AuthArguments{}, func(context.Context, loadTokenArgs) (*oauth2.Token, error) { + t.Fatal("loadToken should not be called") + return nil, nil + }) + cmd.Flags().StringP("profile", "p", "", "~/.databrickscfg profile") + cmd.SetContext(ctx) + cmd.SetArgs([]string{"--format=docker", "DEFAULT"}) + + err := cmd.Execute() + require.ErrorContains(t, err, "--format=docker does not accept positional arguments") +} + +func TestTokenDockerFormatRejectsAuthSelectionFlags(t *testing.T) { + ctx := cmdio.MockDiscard(t.Context()) + dir := t.TempDir() + configFile := filepath.Join(dir, ".databrickscfg") + require.NoError(t, databrickscfg.SaveToProfile(ctx, &config.Config{ + ConfigFile: configFile, + Profile: "DEFAULT", + Host: "https://profile.cloud.databricks.test", + AuthType: authTypeDatabricksCLI, + })) + t.Setenv("DATABRICKS_CONFIG_FILE", configFile) + t.Setenv(storage.EnvVar, string(storage.StorageModePlaintext)) + t.Setenv("HOME", dir) + + cases := [][]string{ + {"--format=docker", "--profile", "DEFAULT"}, + {"--format=docker", "--host", "https://workspace.cloud.databricks.test"}, + {"--format=docker", "--profile", "DEFAULT", "--host", "https://workspace.cloud.databricks.test"}, + {"--format=docker", "--account-id", "abc"}, + {"--format=docker", "--workspace-id", "123456789"}, + } + + for _, args := range cases { + t.Run(strings.Join(args, " "), func(t *testing.T) { + var authArgs auth.AuthArguments + cmd := &cobra.Command{Use: "auth"} + cmd.PersistentFlags().StringVar(&authArgs.Host, "host", "", "Databricks Host") + cmd.PersistentFlags().StringVar(&authArgs.AccountID, "account-id", "", "Databricks Account ID") + cmd.PersistentFlags().StringVar(&authArgs.WorkspaceID, "workspace-id", "", "Databricks Workspace ID") + cmd.AddCommand(newTokenCommandWithLoader(&authArgs, func(context.Context, loadTokenArgs) (*oauth2.Token, error) { + t.Fatal("loadToken should not be called") + return nil, nil + })) + cmd.PersistentFlags().StringP("profile", "p", "", "~/.databrickscfg profile") + cmd.SetContext(ctx) + cmd.SetIn(strings.NewReader("123456789.container.us-west-2.cloud.databricks.com\n")) + cmd.SetArgs(append([]string{"token"}, args...)) + + err := cmd.Execute() + require.ErrorContains(t, err, "--format=docker does not support") + }) + } +} + +func TestTokenDockerFormatRejectsNonDARHost(t *testing.T) { + ctx := cmdio.MockDiscard(t.Context()) + t.Setenv(storage.EnvVar, string(storage.StorageModePlaintext)) + t.Setenv("HOME", t.TempDir()) + + cmd := newTokenCommandWithLoader(&auth.AuthArguments{}, func(context.Context, loadTokenArgs) (*oauth2.Token, error) { + t.Fatal("loadToken should not be called") + return nil, nil + }) + cmd.Flags().StringP("profile", "p", "", "~/.databrickscfg profile") + cmd.SetContext(ctx) + cmd.SetIn(strings.NewReader("registry.example.com\n")) + cmd.SetArgs([]string{"--format=docker"}) + + err := cmd.Execute() + require.ErrorContains(t, err, "is not a Databricks Artifact Registry host") +} + +func TestTokenDockerFormatErrorsWithoutMatchingProfile(t *testing.T) { + ctx := cmdio.MockDiscard(t.Context()) + dir := t.TempDir() + configFile := filepath.Join(dir, ".databrickscfg") + require.NoError(t, os.WriteFile(configFile, []byte(""), 0o600)) + + t.Setenv("DATABRICKS_CONFIG_FILE", configFile) + t.Setenv(storage.EnvVar, string(storage.StorageModePlaintext)) + t.Setenv("HOME", dir) + + cmd := newTokenCommandWithLoader(&auth.AuthArguments{}, func(context.Context, loadTokenArgs) (*oauth2.Token, error) { + t.Fatal("loadToken should not be called") + return nil, nil + }) + cmd.Flags().StringP("profile", "p", "", "~/.databrickscfg profile") + cmd.SetContext(ctx) + cmd.SetIn(strings.NewReader("123456789.container.us-west-2.cloud.databricks.com\n")) + cmd.SetArgs([]string{"--format=docker"}) + + err := cmd.Execute() + require.ErrorContains(t, err, "no Databricks profile found for workspace ID 123456789") + require.ErrorContains(t, err, "databricks auth login --host ") + require.ErrorContains(t, err, "workspace_id") +} + +func TestTokenDockerFormatErrorsWithMultipleMatchingProfiles(t *testing.T) { + ctx := cmdio.MockDiscard(t.Context()) + dir := t.TempDir() + configFile := filepath.Join(dir, ".databrickscfg") + for _, name := range []string{"one", "two"} { + require.NoError(t, databrickscfg.SaveToProfile(ctx, &config.Config{ + ConfigFile: configFile, + Profile: name, + Host: "https://" + name + ".cloud.databricks.test", + WorkspaceID: "123456789", + AuthType: authTypeDatabricksCLI, + })) + } + + t.Setenv("DATABRICKS_CONFIG_FILE", configFile) + t.Setenv(storage.EnvVar, string(storage.StorageModePlaintext)) + t.Setenv("HOME", dir) + + cmd := newTokenCommandWithRegistryHost(&auth.AuthArguments{}, func(context.Context, loadTokenArgs) (*oauth2.Token, error) { + t.Fatal("loadToken should not be called") + return nil, nil + }, func(workspaceID, region, _ string) (string, error) { + return workspaceID + ".container." + region + ".cloud.databricks.com", nil + }) + cmd.Flags().StringP("profile", "p", "", "~/.databrickscfg profile") + cmd.SetContext(ctx) + cmd.SetIn(strings.NewReader("123456789.container.us-west-2.cloud.databricks.com\n")) + cmd.SetArgs([]string{"--format=docker"}) + + err := cmd.Execute() + require.ErrorContains(t, err, "multiple Databricks profiles match workspace ID 123456789") + require.ErrorContains(t, err, "one and two") + require.ErrorContains(t, err, "Remove duplicate workspace_id entries") +} + +func TestTokenDockerFormatRejectsUnsupportedProfile(t *testing.T) { + ctx := cmdio.MockDiscard(t.Context()) + profiler := profile.InMemoryProfiler{ + Profiles: profile.Profiles{ + { + Name: "pat", + Host: "https://workspace.cloud.databricks.test", + WorkspaceID: "123456789", + AuthType: "pat", + }, + { + Name: "m2m", + Host: "https://m2m.cloud.databricks.test", + WorkspaceID: "987654321", + HasClientCredentials: true, + }, + { + Name: "blank-auth", + Host: "https://blank-auth.cloud.databricks.test", + WorkspaceID: "111222333", + }, + { + Name: "account", + Host: "https://accounts.cloud.databricks.test", + AccountID: "account-id", + WorkspaceID: "444555666", + AuthType: authTypeDatabricksCLI, + }, + }, + } + + for _, tc := range []struct { + registryHost string + wantError string + }{ + {"123456789.container.us-west-2.cloud.databricks.com", "requires a profile created by databricks auth login"}, + {"987654321.container.us-west-2.cloud.databricks.com", "requires a profile created by databricks auth login"}, + {"111222333.container.us-west-2.cloud.databricks.com", "requires a profile created by databricks auth login"}, + {"444555666.container.us-west-2.cloud.databricks.com", "does not target a workspace"}, + } { + t.Run(tc.registryHost, func(t *testing.T) { + cmd := &cobra.Command{Use: "token"} + cmd.SetContext(ctx) + cmd.SetIn(strings.NewReader(tc.registryHost + "\n")) + + err := writeDockerTokenOutput(ctx, cmd, loadTokenArgs{ + authArguments: &auth.AuthArguments{}, + profiler: profiler, + }, func(context.Context, loadTokenArgs) (*oauth2.Token, error) { + t.Fatal("loadToken should not be called") + return nil, nil + }, dockercredentials.RegistryHost) + require.ErrorContains(t, err, tc.wantError) + }) + } +} + func TestWriteTokenOutput(t *testing.T) { token := &oauth2.Token{ AccessToken: "my-access-token", diff --git a/libs/dockercredentials/registry.go b/libs/dockercredentials/registry.go new file mode 100644 index 00000000000..fee76b0ca64 --- /dev/null +++ b/libs/dockercredentials/registry.go @@ -0,0 +1,192 @@ +package dockercredentials + +import ( + "errors" + "fmt" + "net" + "net/url" + "strconv" + "strings" + "unicode" + + "github.com/databricks/databricks-sdk-go/common/environment" +) + +const ( + // OAuthTokenUsername is the username returned to Docker with an OAuth access token. + OAuthTokenUsername = "oauthtoken" + registryHostInfix = ".container." +) + +// Registry identifies the workspace, region, and canonical host of a Databricks Artifact Registry endpoint. +type Registry struct { + WorkspaceID string + Region string + Host string +} + +// RegistryHost derives .container.. while preserving the workspace's cloud and environment zone. +func RegistryHost(workspaceID, region, workspaceHost string) (string, error) { + workspaceID = strings.TrimSpace(workspaceID) + region = strings.TrimSpace(region) + if workspaceID == "" { + return "", errors.New("workspace ID is required") + } + if region == "" { + return "", errors.New("region is required") + } + if !isDNSLabel(workspaceID) { + return "", fmt.Errorf("invalid workspace ID %q", workspaceID) + } + if !isDNSLabel(region) { + return "", fmt.Errorf("invalid region %q", region) + } + dnsZone, err := registryDNSZoneForWorkspaceHost(workspaceHost) + if err != nil { + return "", err + } + return fmt.Sprintf("%s%s%s%s", workspaceID, registryHostInfix, region, dnsZone), nil +} + +// normalizeServerAddress canonicalizes Docker's URL-or-host input and permits only HTTPS on the default registry port. +func normalizeServerAddress(raw string) (string, error) { + value := strings.TrimSpace(raw) + if value == "" { + return "", errors.New("server address is required") + } + + if strings.Contains(value, "://") { + u, err := url.Parse(value) + if err != nil { + return "", fmt.Errorf("parse server address %q: %w", raw, err) + } + if !strings.EqualFold(u.Scheme, "https") { + return "", fmt.Errorf("unsupported registry URL scheme %q", u.Scheme) + } + value = u.Host + } else if i := strings.IndexByte(value, '/'); i >= 0 { + value = value[:i] + } + + if host, port, ok, err := splitOptionalPort(value); err != nil { + return "", err + } else if ok { + value = host + if err := validatePort(port); err != nil { + return "", err + } + if port != "443" { + return "", fmt.Errorf("unsupported registry port %q", port) + } + } + + value = strings.TrimSuffix(strings.ToLower(value), ".") + if value == "" { + return "", errors.New("server address is required") + } + return value, nil +} + +// ParseRegistryHost normalizes a Databricks Artifact Registry address and extracts its workspace and region. +func ParseRegistryHost(raw string) (Registry, error) { + host, err := normalizeServerAddress(raw) + if err != nil { + return Registry{}, err + } + + dnsZone, ok := matchingDatabricksDNSZone(host) + if !ok { + return Registry{}, fmt.Errorf("%q is not a Databricks Artifact Registry host", host) + } + + trimmed := strings.TrimSuffix(host, dnsZone) + workspaceID, region, ok := strings.Cut(trimmed, registryHostInfix) + if !ok || !isDNSLabel(workspaceID) || !isDNSLabel(region) { + return Registry{}, fmt.Errorf("%q is not a Databricks Artifact Registry host", host) + } + + return Registry{ + WorkspaceID: workspaceID, + Region: region, + Host: host, + }, nil +} + +// registryDNSZoneForWorkspaceHost derives the registry suffix from the workspace's SDK-known cloud and environment zone. +func registryDNSZoneForWorkspaceHost(raw string) (string, error) { + host, err := normalizeServerAddress(raw) + if err != nil { + return "", fmt.Errorf("parse workspace host: %w", err) + } + dnsZone, ok := matchingDatabricksDNSZone(host) + if !ok { + return "", fmt.Errorf("%q is not a supported Databricks workspace host", host) + } + return dnsZone, nil +} + +// matchingDatabricksDNSZone selects the most specific suffix from all SDK-known Databricks environments. +func matchingDatabricksDNSZone(host string) (string, bool) { + return matchingDatabricksDNSZoneInEnvironments(host, environment.AllEnvironments()) +} + +// matchingDatabricksDNSZoneInEnvironments prefers the longest suffix so environment-specific zones beat generic ones. +func matchingDatabricksDNSZoneInEnvironments(host string, envs []environment.DatabricksEnvironment) (string, bool) { + var match string + for _, e := range envs { + dnsZone := strings.ToLower(e.DnsZone) + if dnsZone == "" { + continue + } + if strings.HasSuffix(host, dnsZone) && len(dnsZone) > len(match) { + match = dnsZone + } + } + return match, match != "" +} + +// splitOptionalPort extracts bracketed or plain host ports while leaving non-port colon forms for host validation. +func splitOptionalPort(value string) (host, port string, ok bool, err error) { + host, port, err = net.SplitHostPort(value) + if err == nil { + return host, port, true, nil + } + + if strings.Count(value, ":") == 1 { + host, port, found := strings.Cut(value, ":") + if found && port != "" { + return host, port, true, nil + } + } + + return "", "", false, nil +} + +// validatePort enforces the decimal TCP port range that URL parsing alone does not validate. +func validatePort(port string) error { + n, err := strconv.Atoi(port) + if err != nil || n < 1 || n > 65535 { + return fmt.Errorf("invalid registry port %q", port) + } + return nil +} + +// isDNSLabel restricts interpolated registry components to unescaped lowercase ASCII DNS labels. +func isDNSLabel(label string) bool { + if label == "" || len(label) > 63 { + return false + } + for i, r := range label { + if r > unicode.MaxASCII { + return false + } + if (r >= 'a' && r <= 'z') || (r >= '0' && r <= '9') { + continue + } + if r == '-' && i > 0 && i < len(label)-1 { + continue + } + return false + } + return true +} diff --git a/libs/dockercredentials/registry_test.go b/libs/dockercredentials/registry_test.go new file mode 100644 index 00000000000..39f80c1124b --- /dev/null +++ b/libs/dockercredentials/registry_test.go @@ -0,0 +1,187 @@ +package dockercredentials + +import ( + "testing" + + "github.com/databricks/databricks-sdk-go/common/environment" + "github.com/stretchr/testify/require" +) + +func TestRegistryHost(t *testing.T) { + cases := []struct { + name string + workspaceHost string + region string + want string + }{ + { + name: "aws prod", + workspaceHost: "https://adb-123.456.cloud.databricks.com", + region: "us-west-2", + want: "123456789.container.us-west-2.cloud.databricks.com", + }, + { + name: "aws staging", + workspaceHost: "https://workspace.staging.cloud.databricks.com", + region: "us-west-2", + want: "123456789.container.us-west-2.staging.cloud.databricks.com", + }, + { + name: "azure prod", + workspaceHost: "https://adb-123.456.azuredatabricks.net", + region: "eastus", + want: "123456789.container.eastus.azuredatabricks.net", + }, + { + name: "azure dev", + workspaceHost: "https://workspace.dev.azuredatabricks.net", + region: "eastus", + want: "123456789.container.eastus.dev.azuredatabricks.net", + }, + { + name: "gcp prod", + workspaceHost: "https://workspace.gcp.databricks.com", + region: "us-central1", + want: "123456789.container.us-central1.gcp.databricks.com", + }, + { + name: "gcp dev", + workspaceHost: "https://workspace.dev.gcp.databricks.com", + region: "us-central1", + want: "123456789.container.us-central1.dev.gcp.databricks.com", + }, + } + + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + got, err := RegistryHost("123456789", tc.region, tc.workspaceHost) + require.NoError(t, err) + require.Equal(t, tc.want, got) + }) + } +} + +func TestRegistryHostRejectsEmptyParts(t *testing.T) { + _, err := RegistryHost("", "us-west-2", "https://workspace.cloud.databricks.test") + require.ErrorContains(t, err, "workspace ID is required") + + _, err = RegistryHost("123456789", "", "https://workspace.cloud.databricks.test") + require.ErrorContains(t, err, "region is required") +} + +func TestRegistryHostRejectsUnsupportedWorkspaceHost(t *testing.T) { + _, err := RegistryHost("123456789", "us-west-2", "https://workspace.example.test") + require.ErrorContains(t, err, `"workspace.example.test" is not a supported Databricks workspace host`) +} + +func TestParseRegistryHost(t *testing.T) { + cases := []string{ + "123456789.container.us-west-2.cloud.databricks.com", + "https://123456789.container.us-west-2.cloud.databricks.com", + "123456789.container.us-west-2.cloud.databricks.com/v2/", + } + + for _, input := range cases { + t.Run(input, func(t *testing.T) { + got, err := ParseRegistryHost(input) + require.NoError(t, err) + require.Equal(t, Registry{ + WorkspaceID: "123456789", + Region: "us-west-2", + Host: "123456789.container.us-west-2.cloud.databricks.com", + }, got) + }) + } +} + +func TestRegistryHostAndParseRegistryHostSupportAllDatabricksEnvironmentZones(t *testing.T) { + for _, env := range environment.AllEnvironments() { + dnsZone := env.DnsZone + if dnsZone == "" { + continue + } + t.Run(dnsZone, func(t *testing.T) { + wantHost := "123456789.container.test-region" + dnsZone + got, err := RegistryHost("123456789", "test-region", "https://workspace"+dnsZone) + require.NoError(t, err) + require.Equal(t, wantHost, got) + + registry, err := ParseRegistryHost("https://" + wantHost + "/v2/") + require.NoError(t, err) + require.Equal(t, Registry{ + WorkspaceID: "123456789", + Region: "test-region", + Host: wantHost, + }, registry) + }) + } +} + +func TestParseRegistryHostUsesLongestDNSZoneSuffix(t *testing.T) { + got, err := ParseRegistryHost("123456789.container.us-west-2.staging.cloud.databricks.com") + require.NoError(t, err) + require.Equal(t, Registry{ + WorkspaceID: "123456789", + Region: "us-west-2", + Host: "123456789.container.us-west-2.staging.cloud.databricks.com", + }, got) +} + +func TestMatchingDatabricksDNSZoneIgnoresEmptyDNSZones(t *testing.T) { + got, ok := matchingDatabricksDNSZoneInEnvironments("workspace.example.test", []environment.DatabricksEnvironment{ + {DnsZone: ""}, + {DnsZone: ".example.test"}, + }) + require.True(t, ok) + require.Equal(t, ".example.test", got) + + _, ok = matchingDatabricksDNSZoneInEnvironments("workspace.invalid", []environment.DatabricksEnvironment{ + {DnsZone: ""}, + }) + require.False(t, ok) +} + +func TestParseRegistryHostRejectsNonDARHost(t *testing.T) { + _, err := ParseRegistryHost("registry.example.com") + require.ErrorContains(t, err, `"registry.example.com" is not a Databricks Artifact Registry host`) +} + +func TestParseRegistryHostRejectsPluralContainersInfix(t *testing.T) { + _, err := ParseRegistryHost("123.containers.us-west-2.cloud.databricks.com") + require.ErrorContains(t, err, `"123.containers.us-west-2.cloud.databricks.com" is not a Databricks Artifact Registry host`) +} + +func TestParseRegistryHostRejectsInvalidLabels(t *testing.T) { + _, err := ParseRegistryHost("-123.container.us-west-2.cloud.databricks.com") + require.ErrorContains(t, err, `"-123.container.us-west-2.cloud.databricks.com" is not a Databricks Artifact Registry host`) + + _, err = ParseRegistryHost("123.container.-us-west-2.cloud.databricks.com") + require.ErrorContains(t, err, `"123.container.-us-west-2.cloud.databricks.com" is not a Databricks Artifact Registry host`) +} + +func TestNormalizeServerAddress(t *testing.T) { + got, err := normalizeServerAddress("HTTPS://123.container.US-WEST-2.cloud.databricks.com/v2/") + require.NoError(t, err) + require.Equal(t, "123.container.us-west-2.cloud.databricks.com", got) +} + +func TestNormalizeServerAddressRejectsNonHTTPSURL(t *testing.T) { + _, err := normalizeServerAddress("http://123.container.us-west-2.cloud.databricks.com") + require.ErrorContains(t, err, "unsupported registry URL scheme") +} + +func TestNormalizeServerAddressRejectsInvalidPort(t *testing.T) { + _, err := normalizeServerAddress("https://123.container.us-west-2.cloud.databricks.com:99999") + require.ErrorContains(t, err, "invalid registry port") +} + +func TestNormalizeServerAddressRejectsNonHTTPSPort(t *testing.T) { + _, err := normalizeServerAddress("https://123.container.us-west-2.cloud.databricks.com:5000") + require.ErrorContains(t, err, "unsupported registry port") +} + +func TestNormalizeServerAddressAllowsHTTPSPort(t *testing.T) { + got, err := normalizeServerAddress("https://123.container.us-west-2.cloud.databricks.com:443") + require.NoError(t, err) + require.Equal(t, "123.container.us-west-2.cloud.databricks.com", got) +}