| 1 | // SPDX-License-Identifier: GPL-3.0-or-later |
| 2 | |
| 3 | package sqladapter |
| 4 | |
| 5 | import ( |
| 6 | "net/url" |
| 7 | "testing" |
| 8 | |
| 9 | "github.com/stretchr/testify/assert" |
| 10 | "github.com/stretchr/testify/require" |
| 11 | |
| 12 | "github.com/netdata/netdata/go/plugins/plugin/go.d/pkg/cloudauth" |
| 13 | ) |
| 14 | |
| 15 | func TestMSSQLDriver(t *testing.T) { |
| 16 | assert.Equal(t, MSSQLDriverName, MSSQLDriver(cloudauth.Config{})) |
| 17 | assert.Equal(t, MSSQLAzureDriverName, MSSQLDriver(cloudauth.Config{Provider: cloudauth.ProviderAzureAD})) |
| 18 | } |
| 19 | |
| 20 | func TestBuildMSSQLAzureADDSN(t *testing.T) { |
| 21 | base := "sqlserver://localhost:1433?database=master" |
| 22 | |
| 23 | tests := map[string]struct { |
| 24 | baseDSN string |
| 25 | cfg cloudauth.Config |
| 26 | wantErr bool |
| 27 | validate func(t *testing.T, dsn string, u *url.URL) |
| 28 | parseDSN bool |
| 29 | }{ |
| 30 | "provider disabled returns original dsn": { |
| 31 | baseDSN: base, |
| 32 | cfg: cloudauth.Config{Provider: cloudauth.ProviderNone}, |
| 33 | validate: func(t *testing.T, dsn string, _ *url.URL) { |
| 34 | assert.Equal(t, base, dsn) |
| 35 | }, |
| 36 | }, |
| 37 | "service principal": { |
| 38 | baseDSN: base, |
| 39 | cfg: cloudauth.Config{ |
| 40 | Provider: cloudauth.ProviderAzureAD, |
| 41 | AzureAD: &cloudauth.AzureADAuthConfig{ |
| 42 | Mode: cloudauth.AzureADAuthModeServicePrincipal, |
| 43 | ModeServicePrincipal: &cloudauth.AzureADModeServicePrincipalConfig{ |
| 44 | TenantID: "tenant", |
| 45 | ClientID: "client", |
| 46 | ClientSecret: "secret", |
| 47 | }, |
| 48 | }, |
| 49 | }, |
| 50 | parseDSN: true, |
| 51 | validate: func(t *testing.T, _ string, u *url.URL) { |
| 52 | assert.Equal(t, "ActiveDirectoryServicePrincipal", u.Query().Get("fedauth")) |
| 53 | assert.Equal(t, "client@tenant", u.User.Username()) |
| 54 | pass, ok := u.User.Password() |
| 55 | require.True(t, ok) |
| 56 | assert.Equal(t, "secret", pass) |
| 57 | }, |
| 58 | }, |
| 59 | "managed identity with client id": { |
| 60 | baseDSN: base, |
| 61 | cfg: cloudauth.Config{ |
| 62 | Provider: cloudauth.ProviderAzureAD, |
| 63 | AzureAD: &cloudauth.AzureADAuthConfig{ |
| 64 | Mode: cloudauth.AzureADAuthModeManagedIdentity, |
| 65 | ModeManagedIdentity: &cloudauth.AzureADModeManagedIdentityConfig{ |
| 66 | ClientID: "mi-client-id", |
| 67 | }, |
| 68 | }, |
| 69 | }, |
| 70 | parseDSN: true, |
| 71 | validate: func(t *testing.T, _ string, u *url.URL) { |
| 72 | assert.Equal(t, "ActiveDirectoryManagedIdentity", u.Query().Get("fedauth")) |
| 73 | assert.Equal(t, "mi-client-id", u.Query().Get("user id")) |
| 74 | assert.Nil(t, u.User) |
| 75 | }, |
| 76 | }, |
| 77 | "default credential": { |
| 78 | baseDSN: base, |
| 79 | cfg: cloudauth.Config{ |
| 80 | Provider: cloudauth.ProviderAzureAD, |
| 81 | AzureAD: &cloudauth.AzureADAuthConfig{Mode: cloudauth.AzureADAuthModeDefault}, |
| 82 | }, |
| 83 | parseDSN: true, |
| 84 | validate: func(t *testing.T, _ string, u *url.URL) { |
| 85 | assert.Equal(t, "ActiveDirectoryDefault", u.Query().Get("fedauth")) |
| 86 | assert.Nil(t, u.User) |
| 87 | }, |
| 88 | }, |
| 89 | "cleans mixed-case stale params": { |
| 90 | baseDSN: "sqlserver://olduser:oldpass@localhost:1433?database=master&FedAuth=old&User+ID=old&Password=old", |
| 91 | cfg: cloudauth.Config{ |
| 92 | Provider: cloudauth.ProviderAzureAD, |
| 93 | AzureAD: &cloudauth.AzureADAuthConfig{Mode: cloudauth.AzureADAuthModeDefault}, |
| 94 | }, |
| 95 | parseDSN: true, |
| 96 | validate: func(t *testing.T, _ string, u *url.URL) { |
| 97 | assert.Empty(t, u.Query().Get("FedAuth")) |
| 98 | assert.Empty(t, u.Query().Get("User ID")) |
| 99 | assert.Empty(t, u.Query().Get("Password")) |
| 100 | assert.Equal(t, "ActiveDirectoryDefault", u.Query().Get("fedauth")) |
| 101 | assert.Nil(t, u.User) |
| 102 | }, |
| 103 | }, |
| 104 | "invalid scheme": { |
| 105 | baseDSN: "server=localhost;database=master", |
| 106 | cfg: cloudauth.Config{ |
| 107 | Provider: cloudauth.ProviderAzureAD, |
| 108 | AzureAD: &cloudauth.AzureADAuthConfig{Mode: cloudauth.AzureADAuthModeDefault}, |
| 109 | }, |
| 110 | wantErr: true, |
| 111 | }, |
| 112 | } |
| 113 | |
| 114 | for name, tc := range tests { |
| 115 | t.Run(name, func(t *testing.T) { |
| 116 | dsn, err := BuildMSSQLAzureADDSN(tc.baseDSN, tc.cfg) |
| 117 | if tc.wantErr { |
| 118 | require.Error(t, err) |
| 119 | return |
| 120 | } |
| 121 | |
| 122 | require.NoError(t, err) |
| 123 | if tc.validate == nil { |
| 124 | return |
| 125 | } |
| 126 | |
| 127 | var u *url.URL |
| 128 | if tc.parseDSN { |
| 129 | parsed, parseErr := url.Parse(dsn) |
| 130 | require.NoError(t, parseErr) |
| 131 | u = parsed |
| 132 | } |
| 133 | tc.validate(t, dsn, u) |
| 134 | }) |
| 135 | } |
| 136 | } |