openai/openai-go

Public

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

CodeCommitsIssuesPull requestsActionsInsightsSecurity
codex/split-release-publish-gate

Branches

Tags

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

Clone

HTTPS

Download ZIP

auth/subjecttokenprovider_test.go

292lines · modecode

1package auth_test
2
3import (
4 "context"
5 "io"
6 "net/http"
7 "os"
8 "strings"
9 "testing"
10
11 "github.com/openai/openai-go/v3/auth"
12)
13
14func TestK8sProviderFileReading(t *testing.T) {
15 tmpFile, err := os.CreateTemp("", "k8s-token-*")
16 if err != nil {
17 t.Fatalf("Failed to create temp file: %v", err)
18 }
19 defer os.Remove(tmpFile.Name())
20
21 tokenContent := " test-jwt-token-123 \n"
22 if _, err := tmpFile.WriteString(tokenContent); err != nil {
23 t.Fatalf("Failed to write to temp file: %v", err)
24 }
25 tmpFile.Close()
26
27 provider := auth.K8sServiceAccountTokenProvider(tmpFile.Name())
28
29 token, err := provider.GetToken(context.Background(), nil)
30 if err != nil {
31 t.Fatalf("GetToken() error = %v", err)
32 }
33
34 expectedToken := "test-jwt-token-123"
35 if token != expectedToken {
36 t.Errorf("GetToken() = %q, want %q", token, expectedToken)
37 }
38
39 if provider.TokenType() != auth.SubjectTokenTypeJWT {
40 t.Errorf("TokenType() = %v, want %v", provider.TokenType(), auth.SubjectTokenTypeJWT)
41 }
42}
43
44func TestK8sProviderDefaultPath(t *testing.T) {
45 provider := auth.K8sServiceAccountTokenProvider("")
46
47 defaultPath := "/var/run/secrets/kubernetes.io/serviceaccount/token"
48
49 _, err := provider.GetToken(context.Background(), nil)
50 if err == nil {
51 t.Log("Default path file exists, skipping validation")
52 return
53 }
54
55 providerErr, ok := err.(*auth.SubjectTokenProviderError)
56 if !ok {
57 t.Fatalf("Expected *SubjectTokenProviderError, got %T", err)
58 }
59
60 if providerErr.Provider != "kubernetes" {
61 t.Errorf("Provider = %q, want %q", providerErr.Provider, "kubernetes")
62 }
63
64 if !strings.Contains(providerErr.Error(), defaultPath) {
65 t.Errorf("Error should reference default path %q: %v", defaultPath, providerErr)
66 }
67}
68
69func TestK8sProviderErrorHandling(t *testing.T) {
70 provider := auth.K8sServiceAccountTokenProvider("/nonexistent/path/to/token")
71
72 _, err := provider.GetToken(context.Background(), nil)
73 if err == nil {
74 t.Fatal("Expected error, got nil")
75 }
76
77 providerErr, ok := err.(*auth.SubjectTokenProviderError)
78 if !ok {
79 t.Fatalf("Expected *SubjectTokenProviderError, got %T", err)
80 }
81
82 if providerErr.Provider != "kubernetes" {
83 t.Errorf("Provider = %q, want %q", providerErr.Provider, "kubernetes")
84 }
85
86 if providerErr.Cause == nil {
87 t.Error("Expected Cause to be set")
88 }
89}
90
91func TestK8sProviderEmptyToken(t *testing.T) {
92 tmpFile, err := os.CreateTemp("", "k8s-token-empty-*")
93 if err != nil {
94 t.Fatalf("Failed to create temp file: %v", err)
95 }
96 defer os.Remove(tmpFile.Name())
97
98 if _, err := tmpFile.WriteString(" \n "); err != nil {
99 t.Fatalf("Failed to write to temp file: %v", err)
100 }
101 tmpFile.Close()
102
103 provider := auth.K8sServiceAccountTokenProvider(tmpFile.Name())
104
105 _, err = provider.GetToken(context.Background(), nil)
106 if err == nil {
107 t.Fatal("Expected error for empty token, got nil")
108 }
109
110 providerErr, ok := err.(*auth.SubjectTokenProviderError)
111 if !ok {
112 t.Fatalf("Expected *SubjectTokenProviderError, got %T", err)
113 }
114
115 if providerErr.Provider != "kubernetes" {
116 t.Errorf("Provider = %q, want %q", providerErr.Provider, "kubernetes")
117 }
118
119 if !strings.Contains(err.Error(), "empty") {
120 t.Errorf("Error should mention 'empty': %v", err)
121 }
122}
123
124func TestAzureProviderTokenType(t *testing.T) {
125 provider := auth.AzureManagedIdentityTokenProvider(nil)
126
127 if provider.TokenType() != auth.SubjectTokenTypeJWT {
128 t.Errorf("TokenType() = %v, want %v", provider.TokenType(), auth.SubjectTokenTypeJWT)
129 }
130}
131
132func TestAzureProviderGetToken(t *testing.T) {
133 mockClient := &http.Client{
134 Transport: roundTripFunc(func(req *http.Request) (*http.Response, error) {
135 if req.Header.Get("Metadata") != "true" {
136 t.Fatalf("Metadata header = %q, want true", req.Header.Get("Metadata"))
137 }
138 return &http.Response{
139 StatusCode: http.StatusOK,
140 Body: io.NopCloser(strings.NewReader(`{"access_token":"azure-token-123"}`)),
141 Header: make(http.Header),
142 }, nil
143 }),
144 }
145
146 provider := auth.AzureManagedIdentityTokenProvider(nil)
147 token, err := provider.GetToken(context.Background(), mockClient)
148 if err != nil {
149 t.Fatalf("GetToken() error = %v", err)
150 }
151 if token != "azure-token-123" {
152 t.Errorf("GetToken() = %q, want %q", token, "azure-token-123")
153 }
154}
155
156func TestGCPProviderTokenType(t *testing.T) {
157 provider := auth.GCPIDTokenProvider(nil)
158
159 if provider.TokenType() != auth.SubjectTokenTypeID {
160 t.Errorf("TokenType() = %v, want %v", provider.TokenType(), auth.SubjectTokenTypeID)
161 }
162}
163
164func TestAzureProviderCustomResource(t *testing.T) {
165 mockClient := &http.Client{
166 Transport: roundTripFunc(func(req *http.Request) (*http.Response, error) {
167 if !strings.Contains(req.URL.RawQuery, "resource=https%3A%2F%2Fcustom.openai.com") {
168 t.Fatalf("Expected custom resource in URL, got %q", req.URL.RawQuery)
169 }
170 return &http.Response{
171 StatusCode: http.StatusOK,
172 Body: io.NopCloser(strings.NewReader(`{"access_token":"azure-custom-token"}`)),
173 Header: make(http.Header),
174 }, nil
175 }),
176 }
177
178 provider := auth.AzureManagedIdentityTokenProvider(&auth.AzureManagedIdentityTokenProviderConfig{
179 Resource: "https://custom.openai.com",
180 })
181 token, err := provider.GetToken(context.Background(), mockClient)
182 if err != nil {
183 t.Fatalf("GetToken() error = %v", err)
184 }
185 if token != "azure-custom-token" {
186 t.Errorf("GetToken() = %q, want %q", token, "azure-custom-token")
187 }
188}
189
190func TestGCPProviderCustomAudience(t *testing.T) {
191 mockClient := &http.Client{
192 Transport: roundTripFunc(func(req *http.Request) (*http.Response, error) {
193 if !strings.Contains(req.URL.RawQuery, "audience=https%3A%2F%2Fcustom.openai.com") {
194 t.Fatalf("Expected custom audience in URL, got %q", req.URL.RawQuery)
195 }
196 return &http.Response{
197 StatusCode: http.StatusOK,
198 Body: io.NopCloser(strings.NewReader("gcp-custom-token-jwt")),
199 Header: make(http.Header),
200 }, nil
201 }),
202 }
203
204 provider := auth.GCPIDTokenProvider(&auth.GCPIDTokenProviderConfig{
205 Audience: "https://custom.openai.com",
206 })
207 token, err := provider.GetToken(context.Background(), mockClient)
208 if err != nil {
209 t.Fatalf("GetToken() error = %v", err)
210 }
211 if token != "gcp-custom-token-jwt" {
212 t.Errorf("GetToken() = %q, want %q", token, "gcp-custom-token-jwt")
213 }
214}
215
216func TestAzureProviderNilConfigBackwardCompatibility(t *testing.T) {
217 mockClient := &http.Client{
218 Transport: roundTripFunc(func(req *http.Request) (*http.Response, error) {
219 if !strings.Contains(req.URL.RawQuery, "resource=https%3A%2F%2Fmanagement.azure.com%2F") {
220 t.Fatalf("Expected default resource in URL, got %q", req.URL.RawQuery)
221 }
222 if !strings.Contains(req.URL.RawQuery, "api-version=2018-02-01") {
223 t.Fatalf("Expected default API version in URL, got %q", req.URL.RawQuery)
224 }
225 return &http.Response{
226 StatusCode: http.StatusOK,
227 Body: io.NopCloser(strings.NewReader(`{"access_token":"default-token"}`)),
228 Header: make(http.Header),
229 }, nil
230 }),
231 }
232
233 provider := auth.AzureManagedIdentityTokenProvider(nil)
234 _, err := provider.GetToken(context.Background(), mockClient)
235 if err != nil {
236 t.Fatalf("GetToken() error = %v", err)
237 }
238}
239
240func TestGCPProviderNilConfigBackwardCompatibility(t *testing.T) {
241 mockClient := &http.Client{
242 Transport: roundTripFunc(func(req *http.Request) (*http.Response, error) {
243 if !strings.Contains(req.URL.RawQuery, "audience=https%3A%2F%2Fapi.openai.com") {
244 t.Fatalf("Expected default audience in URL, got %q", req.URL.RawQuery)
245 }
246 expectedPath := "/computeMetadata/v1/instance/service-accounts/default/identity"
247 if req.URL.Path != expectedPath {
248 t.Fatalf("Expected path %q, got %q", expectedPath, req.URL.Path)
249 }
250 return &http.Response{
251 StatusCode: http.StatusOK,
252 Body: io.NopCloser(strings.NewReader("default-gcp-token")),
253 Header: make(http.Header),
254 }, nil
255 }),
256 }
257
258 provider := auth.GCPIDTokenProvider(nil)
259 _, err := provider.GetToken(context.Background(), mockClient)
260 if err != nil {
261 t.Fatalf("GetToken() error = %v", err)
262 }
263}
264
265func TestAzureProviderResourceURLEncoding(t *testing.T) {
266 mockClient := &http.Client{
267 Transport: roundTripFunc(func(req *http.Request) (*http.Response, error) {
268 if !strings.Contains(req.URL.RawQuery, "resource=https%3A%2F%2Fapi.openai.com%2Fv1%2Fspecial") {
269 t.Fatalf("Expected URL-encoded resource with path, got %q", req.URL.RawQuery)
270 }
271 return &http.Response{
272 StatusCode: http.StatusOK,
273 Body: io.NopCloser(strings.NewReader(`{"access_token":"encoded-token"}`)),
274 Header: make(http.Header),
275 }, nil
276 }),
277 }
278
279 provider := auth.AzureManagedIdentityTokenProvider(&auth.AzureManagedIdentityTokenProviderConfig{
280 Resource: "https://api.openai.com/v1/special",
281 })
282 _, err := provider.GetToken(context.Background(), mockClient)
283 if err != nil {
284 t.Fatalf("GetToken() error = %v", err)
285 }
286}
287
288type roundTripFunc func(*http.Request) (*http.Response, error)
289
290func (f roundTripFunc) RoundTrip(req *http.Request) (*http.Response, error) {
291 return f(req)
292}
293