master
go 203 lines 5.94 KB
Raw
1 // SPDX-License-Identifier: GPL-3.0-or-later
2
3 package cloudauth
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/stretchr/testify/assert"
14 "github.com/stretchr/testify/require"
15 )
16
17 type fakeTokenCredential struct {
18 getToken func(ctx context.Context, opts policy.TokenRequestOptions) (azcore.AccessToken, error)
19 }
20
21 func (f fakeTokenCredential) GetToken(ctx context.Context, opts policy.TokenRequestOptions) (azcore.AccessToken, error) {
22 return f.getToken(ctx, opts)
23 }
24
25 func TestNewTokenProviderValidation(t *testing.T) {
26 _, err := NewTokenProvider(nil, []string{"scope"}, time.Minute)
27 require.Error(t, err)
28
29 cred := fakeTokenCredential{
30 getToken: func(ctx context.Context, opts policy.TokenRequestOptions) (azcore.AccessToken, error) {
31 return azcore.AccessToken{}, nil
32 },
33 }
34 _, err = NewTokenProvider(cred, nil, time.Minute)
35 require.Error(t, err)
36 }
37
38 func TestNewTokenProviderDefaultsRefreshMarginWhenNonPositive(t *testing.T) {
39 cred := fakeTokenCredential{
40 getToken: func(ctx context.Context, opts policy.TokenRequestOptions) (azcore.AccessToken, error) {
41 return azcore.AccessToken{}, nil
42 },
43 }
44
45 p, err := NewTokenProvider(cred, []string{"scope"}, 0)
46 require.NoError(t, err)
47 assert.Equal(t, DefaultTokenRefreshMargin, p.refreshMargin)
48
49 p, err = NewTokenProvider(cred, []string{"scope"}, -1*time.Minute)
50 require.NoError(t, err)
51 assert.Equal(t, DefaultTokenRefreshMargin, p.refreshMargin)
52 }
53
54 func TestTokenProviderCachesToken(t *testing.T) {
55 baseNow := time.Date(2026, time.March, 6, 10, 0, 0, 0, time.UTC)
56 calls := 0
57 cred := fakeTokenCredential{
58 getToken: func(ctx context.Context, opts policy.TokenRequestOptions) (azcore.AccessToken, error) {
59 calls++
60 return azcore.AccessToken{
61 Token: "token-1",
62 ExpiresOn: baseNow.Add(30 * time.Minute),
63 }, nil
64 },
65 }
66
67 p, err := NewTokenProvider(cred, []string{"scope"}, 5*time.Minute)
68 require.NoError(t, err)
69 p.now = func() time.Time { return baseNow }
70
71 token, expiry, err := p.Token(context.Background())
72 require.NoError(t, err)
73 assert.Equal(t, "token-1", token)
74 assert.Equal(t, baseNow.Add(30*time.Minute), expiry)
75
76 token, _, err = p.Token(context.Background())
77 require.NoError(t, err)
78 assert.Equal(t, "token-1", token)
79 assert.Equal(t, 1, calls)
80 }
81
82 func TestTokenProviderRefreshesNearExpiry(t *testing.T) {
83 baseNow := time.Date(2026, time.March, 6, 10, 0, 0, 0, time.UTC)
84 currentNow := baseNow
85 calls := 0
86
87 cred := fakeTokenCredential{
88 getToken: func(ctx context.Context, opts policy.TokenRequestOptions) (azcore.AccessToken, error) {
89 calls++
90 return azcore.AccessToken{
91 Token: "token-" + string(rune('0'+calls)),
92 ExpiresOn: currentNow.Add(10 * time.Minute),
93 }, nil
94 },
95 }
96
97 p, err := NewTokenProvider(cred, []string{"scope"}, 5*time.Minute)
98 require.NoError(t, err)
99 p.now = func() time.Time { return currentNow }
100
101 first, _, err := p.Token(context.Background())
102 require.NoError(t, err)
103
104 // Past expiry-refresh boundary: now + margin >= cached expiry.
105 currentNow = baseNow.Add(6 * time.Minute)
106 second, _, err := p.Token(context.Background())
107 require.NoError(t, err)
108
109 assert.NotEqual(t, first, second)
110 assert.Equal(t, 2, calls)
111 }
112
113 func TestTokenProviderReturnsCredentialError(t *testing.T) {
114 cred := fakeTokenCredential{
115 getToken: func(ctx context.Context, opts policy.TokenRequestOptions) (azcore.AccessToken, error) {
116 return azcore.AccessToken{}, errors.New("boom")
117 },
118 }
119
120 p, err := NewTokenProvider(cred, []string{"scope"}, time.Minute)
121 require.NoError(t, err)
122
123 _, _, err = p.Token(context.Background())
124 require.Error(t, err)
125 }
126
127 func TestTokenProviderFallsBackOnRefreshFailure(t *testing.T) {
128 baseNow := time.Date(2026, time.March, 6, 10, 0, 0, 0, time.UTC)
129 currentNow := baseNow
130 calls := 0
131
132 cred := fakeTokenCredential{
133 getToken: func(ctx context.Context, opts policy.TokenRequestOptions) (azcore.AccessToken, error) {
134 calls++
135 if calls == 1 {
136 return azcore.AccessToken{
137 Token: "good-token",
138 ExpiresOn: baseNow.Add(10 * time.Minute),
139 }, nil
140 }
141 return azcore.AccessToken{}, errors.New("transient failure")
142 },
143 }
144
145 p, err := NewTokenProvider(cred, []string{"scope"}, 5*time.Minute)
146 require.NoError(t, err)
147 p.now = func() time.Time { return currentNow }
148
149 // First call succeeds
150 token, _, err := p.Token(context.Background())
151 require.NoError(t, err)
152 assert.Equal(t, "good-token", token)
153
154 // Advance into refresh margin but before expiry
155 currentNow = baseNow.Add(6 * time.Minute)
156 token, _, err = p.Token(context.Background())
157 require.NoError(t, err)
158 assert.Equal(t, "good-token", token)
159
160 // Advance past expiry — should error since no valid cache
161 currentNow = baseNow.Add(11 * time.Minute)
162 _, _, err = p.Token(context.Background())
163 require.Error(t, err)
164 }
165
166 func TestTokenProviderDoesNotFallbackToExpiredTokenAfterSlowRefreshFailure(t *testing.T) {
167 baseNow := time.Date(2026, time.March, 6, 10, 0, 0, 0, time.UTC)
168 currentNow := baseNow
169 calls := 0
170
171 cred := fakeTokenCredential{
172 getToken: func(ctx context.Context, opts policy.TokenRequestOptions) (azcore.AccessToken, error) {
173 calls++
174 if calls == 1 {
175 return azcore.AccessToken{
176 Token: "good-token",
177 ExpiresOn: baseNow.Add(10 * time.Minute),
178 }, nil
179 }
180
181 // Simulate a slow refresh attempt that crosses token expiry before failing.
182 currentNow = baseNow.Add(11 * time.Minute)
183 return azcore.AccessToken{}, errors.New("slow refresh failure")
184 },
185 }
186
187 p, err := NewTokenProvider(cred, []string{"scope"}, 5*time.Minute)
188 require.NoError(t, err)
189 p.now = func() time.Time { return currentNow }
190
191 // Seed cache.
192 token, _, err := p.Token(context.Background())
193 require.NoError(t, err)
194 assert.Equal(t, "good-token", token)
195
196 // Enter refresh window while token is still valid.
197 currentNow = baseNow.Add(6 * time.Minute)
198
199 // Refresh fails after cache has expired in wall-clock time.
200 _, _, err = p.Token(context.Background())
201 require.Error(t, err)
202 assert.Equal(t, 2, calls)
203 }