openai/openai-go

Public

mirrored from https://github.com/openai/openai-goAvailable

CodeCommitsIssuesPull requestsActionsInsightsSecurity
next

Branches

Tags

  • No tags available.
0Branches0Tags
Go to file
Add file
Code

Clone

HTTPS

Download ZIP

auth/workloadidentity_test.go

515lines · modecode

1package auth_test
2
3import (
4 "context"
5 "encoding/json"
6 "io"
7 "net/http"
8 "strings"
9 "sync"
10 "testing"
11 "time"
12
13 "github.com/openai/openai-go/v3/auth"
14)
15
16type mockProvider struct {
17 token string
18 tokenType auth.SubjectTokenType
19 callCount int
20 delay time.Duration
21 err error
22 mu sync.Mutex
23}
24
25func (m *mockProvider) TokenType() auth.SubjectTokenType {
26 return m.tokenType
27}
28
29func (m *mockProvider) GetToken(ctx context.Context, _ auth.HTTPDoer) (string, error) {
30 m.mu.Lock()
31 m.callCount++
32 m.mu.Unlock()
33
34 if m.delay > 0 {
35 time.Sleep(m.delay)
36 }
37
38 if m.err != nil {
39 return "", m.err
40 }
41
42 return m.token, nil
43}
44
45func (m *mockProvider) GetCallCount() int {
46 m.mu.Lock()
47 defer m.mu.Unlock()
48 return m.callCount
49}
50
51type closureTransport struct {
52 fn func(req *http.Request) (*http.Response, error)
53}
54
55func (t *closureTransport) RoundTrip(req *http.Request) (*http.Response, error) {
56 return t.fn(req)
57}
58
59func mockOAuthServer(responseBody string, statusCode int) *http.Client {
60 return &http.Client{
61 Transport: &closureTransport{
62 fn: func(req *http.Request) (*http.Response, error) {
63 return &http.Response{
64 StatusCode: statusCode,
65 Body: io.NopCloser(strings.NewReader(responseBody)),
66 Header: make(http.Header),
67 }, nil
68 },
69 },
70 }
71}
72
73func TestWorkloadIdentityClientIDOptional(t *testing.T) {
74 provider := &mockProvider{
75 token: "test-subject-token",
76 tokenType: auth.SubjectTokenTypeJWT,
77 }
78
79 var requestBody map[string]any
80 httpClient := &http.Client{
81 Transport: &closureTransport{
82 fn: func(req *http.Request) (*http.Response, error) {
83 body, err := io.ReadAll(req.Body)
84 if err != nil {
85 t.Fatalf("failed reading request body: %v", err)
86 return nil, err
87 }
88 if err := json.Unmarshal(body, &requestBody); err != nil {
89 t.Fatalf("failed decoding request body: %v", err)
90 return nil, err
91 }
92
93 return &http.Response{
94 StatusCode: 200,
95 Body: io.NopCloser(strings.NewReader(`{"access_token": "exchanged-token-123", "expires_in": 3600}`)),
96 Header: make(http.Header),
97 }, nil
98 },
99 },
100 }
101
102 config := auth.WorkloadIdentity{IdentityProviderID: "idp-id", ServiceAccountID: "sa-id", Provider: provider}
103 wa, err := auth.NewWorkloadIdentityAuth(config)
104 if err != nil {
105 t.Fatalf("NewWorkloadIdentityAuth() error = %v", err)
106 }
107
108 _, err = wa.GetToken(context.Background(), httpClient)
109 if err != nil {
110 t.Fatalf("GetToken() error = %v", err)
111 }
112
113 if _, ok := requestBody["client_id"]; ok {
114 t.Errorf("request body contains client_id = %v, want omitted", requestBody["client_id"])
115 }
116
117 if requestBody["identity_provider_id"] != "idp-id" {
118 t.Errorf("identity_provider_id = %v, want idp-id", requestBody["identity_provider_id"])
119 }
120 if requestBody["service_account_id"] != "sa-id" {
121 t.Errorf("service_account_id = %v, want sa-id", requestBody["service_account_id"])
122 }
123}
124
125func TestWorkloadIdentityClientIDIncludedWhenConfigured(t *testing.T) {
126 provider := &mockProvider{
127 token: "test-subject-token",
128 tokenType: auth.SubjectTokenTypeJWT,
129 }
130
131 var requestBody map[string]any
132 httpClient := &http.Client{
133 Transport: &closureTransport{
134 fn: func(req *http.Request) (*http.Response, error) {
135 body, err := io.ReadAll(req.Body)
136 if err != nil {
137 t.Fatalf("failed reading request body: %v", err)
138 return nil, err
139 }
140 if err := json.Unmarshal(body, &requestBody); err != nil {
141 t.Fatalf("failed decoding request body: %v", err)
142 return nil, err
143 }
144
145 return &http.Response{
146 StatusCode: 200,
147 Body: io.NopCloser(strings.NewReader(`{"access_token": "exchanged-token-123", "expires_in": 3600}`)),
148 Header: make(http.Header),
149 }, nil
150 },
151 },
152 }
153
154 config := auth.WorkloadIdentity{ClientID: "client-id", IdentityProviderID: "idp-id", ServiceAccountID: "sa-id", Provider: provider}
155 wa, err := auth.NewWorkloadIdentityAuth(config)
156 if err != nil {
157 t.Fatalf("NewWorkloadIdentityAuth() error = %v", err)
158 }
159
160 _, err = wa.GetToken(context.Background(), httpClient)
161 if err != nil {
162 t.Fatalf("GetToken() error = %v", err)
163 }
164
165 if requestBody["client_id"] != "client-id" {
166 t.Errorf("client_id = %v, want client-id", requestBody["client_id"])
167 }
168}
169
170func TestTokenCaching(t *testing.T) {
171 provider := &mockProvider{
172 token: "test-subject-token",
173 tokenType: auth.SubjectTokenTypeJWT,
174 }
175
176 responseBody := `{"access_token": "exchanged-token-123", "expires_in": 3600}`
177 httpClient := mockOAuthServer(responseBody, 200)
178
179 config := auth.WorkloadIdentity{ClientID: "client-id", IdentityProviderID: "idp-id", ServiceAccountID: "sa-id", Provider: provider}
180 wa, err := auth.NewWorkloadIdentityAuth(config)
181 if err != nil {
182 t.Fatalf("NewWorkloadIdentityAuth() error = %v", err)
183 }
184
185 token1, err := wa.GetToken(context.Background(), httpClient)
186 if err != nil {
187 t.Fatalf("First GetToken() error = %v", err)
188 }
189
190 token2, err := wa.GetToken(context.Background(), httpClient)
191 if err != nil {
192 t.Fatalf("Second GetToken() error = %v", err)
193 }
194
195 if token1 != token2 {
196 t.Errorf("Tokens don't match: %q != %q", token1, token2)
197 }
198
199 if provider.GetCallCount() != 1 {
200 t.Errorf("Provider call count = %d, want 1", provider.GetCallCount())
201 }
202}
203
204func TestTokenExpiration(t *testing.T) {
205 callCount := 0
206 mu := sync.Mutex{}
207
208 transport := &closureTransport{
209 fn: func(req *http.Request) (*http.Response, error) {
210 mu.Lock()
211 callCount++
212 mu.Unlock()
213
214 responseBody := `{"access_token": "token-` + string(rune('0'+callCount)) + `", "expires_in": 0}`
215 return &http.Response{
216 StatusCode: 200,
217 Body: io.NopCloser(strings.NewReader(responseBody)),
218 Header: make(http.Header),
219 }, nil
220 },
221 }
222 httpClient := &http.Client{Transport: transport}
223
224 provider := &mockProvider{
225 token: "test-subject-token",
226 tokenType: auth.SubjectTokenTypeJWT,
227 }
228
229 config := auth.WorkloadIdentity{ClientID: "client-id", IdentityProviderID: "idp-id", ServiceAccountID: "sa-id", Provider: provider}
230 wa, err := auth.NewWorkloadIdentityAuth(config)
231 if err != nil {
232 t.Fatalf("NewWorkloadIdentityAuth() error = %v", err)
233 }
234
235 token1, err := wa.GetToken(context.Background(), httpClient)
236 if err != nil {
237 t.Fatalf("First GetToken() error = %v", err)
238 }
239
240 time.Sleep(100 * time.Millisecond)
241
242 token2, err := wa.GetToken(context.Background(), httpClient)
243 if err != nil {
244 t.Fatalf("Second GetToken() error = %v", err)
245 }
246
247 if token1 == token2 {
248 t.Errorf("Expected different tokens after expiration, got same token: %q", token1)
249 }
250
251 if provider.GetCallCount() < 2 {
252 t.Errorf("Provider call count = %d, want at least 2", provider.GetCallCount())
253 }
254}
255
256func TestConcurrentDeduplication(t *testing.T) {
257 provider := &mockProvider{
258 token: "test-subject-token",
259 tokenType: auth.SubjectTokenTypeJWT,
260 delay: 100 * time.Millisecond,
261 }
262
263 oauthCallCount := 0
264 var oauthMu sync.Mutex
265
266 transport := &closureTransport{
267 fn: func(req *http.Request) (*http.Response, error) {
268 oauthMu.Lock()
269 oauthCallCount++
270 oauthMu.Unlock()
271
272 responseBody := `{"access_token": "exchanged-token-123", "expires_in": 3600}`
273 return &http.Response{
274 StatusCode: 200,
275 Body: io.NopCloser(strings.NewReader(responseBody)),
276 Header: make(http.Header),
277 }, nil
278 },
279 }
280 httpClient := &http.Client{Transport: transport}
281
282 config := auth.WorkloadIdentity{ClientID: "client-id", IdentityProviderID: "idp-id", ServiceAccountID: "sa-id", Provider: provider}
283 wa, err := auth.NewWorkloadIdentityAuth(config)
284 if err != nil {
285 t.Fatalf("NewWorkloadIdentityAuth() error = %v", err)
286 }
287
288 const numGoroutines = 5
289 type result struct {
290 token string
291 err error
292 }
293 resultsChan := make(chan result, numGoroutines)
294
295 for i := 0; i < numGoroutines; i++ {
296 go func() {
297 token, err := wa.GetToken(context.Background(), httpClient)
298 resultsChan <- result{token: token, err: err}
299 }()
300 }
301
302 results := make([]result, 0, numGoroutines)
303 for i := 0; i < numGoroutines; i++ {
304 results = append(results, <-resultsChan)
305 }
306
307 for i, r := range results {
308 if r.err != nil {
309 t.Errorf("Goroutine %d error = %v", i, r.err)
310 }
311 if r.token == "" {
312 t.Errorf("Goroutine %d got empty token (error: %v)", i, r.err)
313 }
314 }
315
316 expectedToken := "exchanged-token-123"
317 for i, r := range results {
318 if r.token != expectedToken {
319 t.Errorf("Goroutine %d got token %q, want %q", i, r.token, expectedToken)
320 }
321 }
322
323 if provider.GetCallCount() != 1 {
324 t.Errorf("Provider call count = %d, want 1", provider.GetCallCount())
325 }
326
327 oauthMu.Lock()
328 finalOAuthCallCount := oauthCallCount
329 oauthMu.Unlock()
330
331 if finalOAuthCallCount != 1 {
332 t.Errorf("OAuth call count = %d, want 1 (deduplication should prevent multiple calls)", finalOAuthCallCount)
333 }
334}
335
336func TestProactiveRefresh(t *testing.T) {
337 provider := &mockProvider{
338 token: "test-subject-token",
339 tokenType: auth.SubjectTokenTypeJWT,
340 }
341
342 callCount := 0
343 mu := sync.Mutex{}
344
345 refreshSignal := make(chan struct{}, 10)
346 transport := &closureTransport{
347 fn: func(req *http.Request) (*http.Response, error) {
348 mu.Lock()
349 callCount++
350 currentCount := callCount
351 mu.Unlock()
352
353 select {
354 case refreshSignal <- struct{}{}:
355 default:
356 }
357
358 responseBody := `{"access_token": "token-` + string(rune('0'+currentCount)) + `", "expires_in": 1}`
359 return &http.Response{
360 StatusCode: 200,
361 Body: io.NopCloser(strings.NewReader(responseBody)),
362 Header: make(http.Header),
363 }, nil
364 },
365 }
366 httpClient := &http.Client{Transport: transport}
367
368 bufferSeconds := 0
369 config := auth.WorkloadIdentity{
370 ClientID: "client-id",
371 IdentityProviderID: "idp-id",
372 ServiceAccountID: "sa-id",
373 Provider: provider,
374 RefreshBufferSeconds: bufferSeconds,
375 }
376 wa, err := auth.NewWorkloadIdentityAuth(config)
377 if err != nil {
378 t.Fatalf("NewWorkloadIdentityAuth() error = %v", err)
379 }
380
381 token1, err := wa.GetToken(context.Background(), httpClient)
382 if err != nil {
383 t.Fatalf("First GetToken() error = %v", err)
384 }
385 <-refreshSignal
386
387 token2, err := wa.GetToken(context.Background(), httpClient)
388 if err != nil {
389 t.Fatalf("Second GetToken() error = %v", err)
390 }
391
392 if token1 != token2 {
393 t.Log("Tokens differ, background refresh may have completed")
394 }
395
396 select {
397 case <-refreshSignal:
398 case <-time.After(5 * time.Second):
399 t.Error("timed out waiting for background refresh")
400 }
401
402 mu.Lock()
403 finalCallCount := callCount
404 mu.Unlock()
405
406 if finalCallCount < 2 {
407 t.Errorf("OAuth call count = %d, want at least 2 (initial + background refresh)", finalCallCount)
408 }
409
410 if provider.GetCallCount() < 2 {
411 t.Errorf("Provider call count = %d, want at least 2", provider.GetCallCount())
412 }
413}
414
415func TestOAuthErrorHandling(t *testing.T) {
416 provider := &mockProvider{
417 token: "test-subject-token",
418 tokenType: auth.SubjectTokenTypeJWT,
419 }
420
421 testCases := []struct {
422 statusCode int
423 shouldBeOAuthError bool
424 }{
425 {400, true},
426 {401, true},
427 {403, true},
428 {500, false},
429 }
430
431 for _, tc := range testCases {
432 errorBody := `{"error": "invalid_grant", "error_description": "Token exchange failed"}`
433 httpClient := mockOAuthServer(errorBody, tc.statusCode)
434
435 config := auth.WorkloadIdentity{ClientID: "client-id", IdentityProviderID: "idp-id", ServiceAccountID: "sa-id", Provider: provider}
436 wa, err := auth.NewWorkloadIdentityAuth(config)
437 if err != nil {
438 t.Fatalf("NewWorkloadIdentityAuth() error = %v", err)
439 }
440
441 _, err = wa.GetToken(context.Background(), httpClient)
442 if err == nil {
443 t.Errorf("Status %d: expected error, got nil", tc.statusCode)
444 continue
445 }
446
447 _, isOAuthError := err.(*auth.OAuthError)
448 if isOAuthError != tc.shouldBeOAuthError {
449 t.Errorf("Status %d: isOAuthError = %v, want %v", tc.statusCode, isOAuthError, tc.shouldBeOAuthError)
450 }
451
452 if tc.shouldBeOAuthError {
453 oauthErr := err.(*auth.OAuthError)
454 if oauthErr.StatusCode != tc.statusCode {
455 t.Errorf("StatusCode = %d, want %d", oauthErr.StatusCode, tc.statusCode)
456 }
457 if oauthErr.ErrorCode != "invalid_grant" {
458 t.Errorf("ErrorCode = %q, want %q", oauthErr.ErrorCode, "invalid_grant")
459 }
460 if oauthErr.ErrorDescription != "Token exchange failed" {
461 t.Errorf("ErrorDescription = %q, want %q", oauthErr.ErrorDescription, "Token exchange failed")
462 }
463 }
464 }
465}
466
467func TestDefaultValues(t *testing.T) {
468 provider := &mockProvider{
469 token: "test-subject-token",
470 tokenType: auth.SubjectTokenTypeJWT,
471 }
472
473 responseBody := `{"access_token": "test-token"}`
474 httpClient := mockOAuthServer(responseBody, 200)
475
476 config := auth.WorkloadIdentity{ClientID: "client-id", IdentityProviderID: "idp-id", ServiceAccountID: "sa-id", Provider: provider}
477 wa, err := auth.NewWorkloadIdentityAuth(config)
478 if err != nil {
479 t.Fatalf("NewWorkloadIdentityAuth() error = %v", err)
480 }
481
482 _, err = wa.GetToken(context.Background(), httpClient)
483 if err != nil {
484 t.Fatalf("GetToken() error = %v", err)
485 }
486
487 time.Sleep(100 * time.Millisecond)
488
489 _, err = wa.GetToken(context.Background(), httpClient)
490 if err != nil {
491 t.Fatalf("Second GetToken() error = %v", err)
492 }
493
494 if provider.GetCallCount() != 1 {
495 t.Errorf("Provider call count = %d, want 1 (token should be cached with 3600s default)", provider.GetCallCount())
496 }
497
498 customBuffer := 300
499 config2 := auth.WorkloadIdentity{
500 ClientID: "client-id",
501 IdentityProviderID: "idp-id",
502 ServiceAccountID: "sa-id",
503 Provider: provider,
504 RefreshBufferSeconds: customBuffer,
505 }
506 wa2, err2 := auth.NewWorkloadIdentityAuth(config2)
507 if err2 != nil {
508 t.Fatalf("NewWorkloadIdentityAuth() error = %v", err2)
509 }
510
511 _, err = wa2.GetToken(context.Background(), httpClient)
512 if err != nil {
513 t.Fatalf("GetToken() with custom buffer error = %v", err)
514 }
515}