master
go 108 lines 3.52 KB
Raw
1 // SPDX-License-Identifier: GPL-3.0-or-later
2
3 package sql
4
5 import (
6 "context"
7 "errors"
8 "testing"
9 "time"
10
11 "github.com/Azure/azure-sdk-for-go/sdk/azcore"
12 "github.com/Azure/azure-sdk-for-go/sdk/azcore/policy"
13 "github.com/jackc/pgx/v5"
14 "github.com/netdata/netdata/go/plugins/plugin/go.d/pkg/cloudauth"
15 "github.com/netdata/netdata/go/plugins/plugin/go.d/pkg/cloudauth/sqladapter"
16 "github.com/stretchr/testify/assert"
17 "github.com/stretchr/testify/require"
18 )
19
20 type fakeAADTokenCredential struct {
21 getToken func(ctx context.Context, opts policy.TokenRequestOptions) (azcore.AccessToken, error)
22 }
23
24 func (f fakeAADTokenCredential) GetToken(ctx context.Context, opts policy.TokenRequestOptions) (azcore.AccessToken, error) {
25 return f.getToken(ctx, opts)
26 }
27
28 func TestCollector_openConnection_AzureADSQLServerRequiresURLDSN(t *testing.T) {
29 c := New()
30 c.Driver = "sqlserver"
31 c.DSN = "server=localhost;database=master"
32 c.CloudAuth.Provider = cloudauth.ProviderAzureAD
33 c.CloudAuth.AzureAD = &cloudauth.AzureADAuthConfig{Mode: cloudauth.AzureADAuthModeDefault}
34
35 err := c.openConnection(context.Background())
36 assert.ErrorContains(t, err, "prepare cloud_auth SQL Server DSN")
37 assert.Nil(t, c.db)
38 }
39
40 func TestCollector_resolveConnectionParams_AzureADSQLServerRewritesDSNAndDriver(t *testing.T) {
41 c := New()
42 c.Driver = "sqlserver"
43 c.DSN = "sqlserver://localhost:1433?database=master"
44 c.CloudAuth.Provider = cloudauth.ProviderAzureAD
45 c.CloudAuth.AzureAD = &cloudauth.AzureADAuthConfig{Mode: cloudauth.AzureADAuthModeDefault}
46
47 driverName, dsn, err := c.resolveConnectionParams()
48 require.NoError(t, err)
49
50 assert.Equal(t, sqladapter.MSSQLAzureDriverName, driverName)
51 assert.Contains(t, dsn, "fedauth=ActiveDirectoryDefault")
52 }
53
54 func TestCollector_openConnection_AzureADPGXWithoutTokenProvider(t *testing.T) {
55 c := New()
56 c.Driver = "pgx"
57 c.DSN = "postgres://netdata@127.0.0.1:5432/postgres"
58 c.CloudAuth.Provider = cloudauth.ProviderAzureAD
59 c.CloudAuth.AzureAD = &cloudauth.AzureADAuthConfig{Mode: cloudauth.AzureADAuthModeDefault}
60
61 err := c.openConnection(context.Background())
62 assert.ErrorContains(t, err, "cloud auth token provider is not initialized for pgx")
63 assert.Nil(t, c.db)
64 }
65
66 func TestCollector_azureADBeforeConnect_SetsPassword(t *testing.T) {
67 cred := fakeAADTokenCredential{
68 getToken: func(context.Context, policy.TokenRequestOptions) (azcore.AccessToken, error) {
69 return azcore.AccessToken{
70 Token: "aad-token",
71 ExpiresOn: time.Now().Add(30 * time.Minute),
72 }, nil
73 },
74 }
75 provider, err := cloudauth.NewTokenProvider(cred, []string{"scope"}, time.Minute)
76 require.NoError(t, err)
77
78 c := New()
79 c.azureTokenProvider = provider
80
81 cfg, err := pgx.ParseConfig("postgres://netdata@127.0.0.1:5432/postgres")
82 require.NoError(t, err)
83
84 err = c.azureADBeforeConnect(context.Background(), cfg)
85 require.NoError(t, err)
86 assert.Equal(t, "aad-token", cfg.Password)
87 }
88
89 func TestCollector_azureADBeforeConnect_ProviderError(t *testing.T) {
90 cred := fakeAADTokenCredential{
91 getToken: func(context.Context, policy.TokenRequestOptions) (azcore.AccessToken, error) {
92 return azcore.AccessToken{}, errors.New("token failure")
93 },
94 }
95 provider, err := cloudauth.NewTokenProvider(cred, []string{"scope"}, time.Minute)
96 require.NoError(t, err)
97
98 c := New()
99 c.azureTokenProvider = provider
100
101 cfg, err := pgx.ParseConfig("postgres://netdata@127.0.0.1:5432/postgres")
102 require.NoError(t, err)
103 cfg.Password = "old-password"
104
105 err = c.azureADBeforeConnect(context.Background(), cfg)
106 assert.ErrorContains(t, err, "token failure")
107 assert.Equal(t, "old-password", cfg.Password)
108 }