Skip to content

Commit e4a4d27

Browse files
feat: support HTTP path for unix scheme in token-url using fragments (#4950)
* Supporting path for UDS endpoint in token-url * PR comments
1 parent 96651c5 commit e4a4d27

2 files changed

Lines changed: 145 additions & 8 deletions

File tree

internal/auth/token_source.go

Lines changed: 21 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -47,7 +47,27 @@ func newProxyTokenSource(
4747
return dialer.DialContext(ctx, u.Scheme, u.Path)
4848
},
4949
}
50-
endpoint = "http://unix?" + u.RawQuery
50+
if escapedFragment := u.EscapedFragment(); escapedFragment != "" {
51+
relURL, err := url.Parse(escapedFragment)
52+
if err != nil {
53+
return nil, fmt.Errorf("parsing fragment %q: %w", escapedFragment, err)
54+
}
55+
targetURL := url.URL{
56+
Scheme: "http",
57+
Host: "unix",
58+
Path: relURL.Path,
59+
RawPath: relURL.RawPath,
60+
RawQuery: relURL.RawQuery,
61+
}
62+
endpoint = targetURL.String()
63+
} else {
64+
targetURL := url.URL{
65+
Scheme: "http",
66+
Host: "unix",
67+
RawQuery: u.RawQuery,
68+
}
69+
endpoint = targetURL.String()
70+
}
5171
}
5272

5373
ts = proxyTokenSource{

internal/auth/token_source_test.go

Lines changed: 124 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -17,24 +17,35 @@ package auth
1717
import (
1818
"context"
1919
"encoding/json"
20+
"fmt"
21+
"net"
2022
"net/http"
2123
"net/http/httptest"
24+
"os"
2225
"testing"
2326

2427
"github.com/stretchr/testify/assert"
2528
"github.com/stretchr/testify/require"
2629
"golang.org/x/oauth2"
2730
)
2831

29-
func Test_NewTokenSourceFromURL_Success(t *testing.T) {
30-
// Create fake token server.
31-
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
32+
func tokenHandler(accessToken string) http.HandlerFunc {
33+
return func(w http.ResponseWriter, r *http.Request) {
3234
token := oauth2.Token{
33-
AccessToken: "test-access-token",
35+
AccessToken: accessToken,
3436
TokenType: "Bearer",
3537
}
36-
require.NoError(t, json.NewEncoder(w).Encode(token))
37-
}))
38+
_ = json.NewEncoder(w).Encode(token)
39+
}
40+
}
41+
42+
func startFakeTCPTokenServer(accessToken string) *httptest.Server {
43+
return httptest.NewServer(tokenHandler(accessToken))
44+
}
45+
46+
func Test_NewTokenSourceFromURL_Success(t *testing.T) {
47+
accessToken := "test-access-token"
48+
server := startFakeTCPTokenServer(accessToken)
3849
defer server.Close()
3950

4051
ts, err := NewTokenSourceFromURL(context.Background(), server.URL, false)
@@ -44,7 +55,7 @@ func Test_NewTokenSourceFromURL_Success(t *testing.T) {
4455
// Fetch token
4556
token, err := ts.Token()
4657
assert.NoError(t, err)
47-
assert.Equal(t, "test-access-token", token.AccessToken)
58+
assert.Equal(t, accessToken, token.AccessToken)
4859
}
4960

5061
func Test_NewTokenSourceFromURL_InvalidURL(t *testing.T) {
@@ -86,3 +97,109 @@ func TestProxyTokenSource_TokenFetch_InvalidJSON(t *testing.T) {
8697
assert.Contains(t, err.Error(), "decode body")
8798
assert.Nil(t, token)
8899
}
100+
101+
func startFakeUDSTokenServer(t *testing.T, listener net.Listener, expectedEscapedPath, expectedQuery, accessToken string) *http.Server {
102+
server := &http.Server{
103+
Handler: http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
104+
assert.Equal(t, expectedEscapedPath, r.URL.EscapedPath())
105+
assert.Equal(t, expectedQuery, r.URL.RawQuery)
106+
assert.Equal(t, "unix", r.Host)
107+
tokenHandler(accessToken)(w, r)
108+
}),
109+
}
110+
go func() {
111+
_ = server.Serve(listener)
112+
}()
113+
return server
114+
}
115+
116+
func Test_NewTokenSourceFromURL_UnixSocket_WithFragment_Success(t *testing.T) {
117+
// Create a temp file for the socket.
118+
tmpFile, err := os.CreateTemp("", "gcsfuse-uds-test-*.sock")
119+
require.NoError(t, err)
120+
socketPath := tmpFile.Name()
121+
require.NoError(t, tmpFile.Close())
122+
require.NoError(t, os.Remove(socketPath)) // remove it so net.Listen can create it
123+
defer func() { _ = os.Remove(socketPath) }()
124+
125+
listener, err := net.Listen("unix", socketPath)
126+
require.NoError(t, err)
127+
defer func() { _ = listener.Close() }()
128+
129+
expectedEscapedPath := "/computeMetadata/v1/instance/service-accounts/default/token"
130+
expectedQuery := "foo=bar&baz=qux"
131+
accessToken := "uds-access-token"
132+
133+
server := startFakeUDSTokenServer(t, listener, expectedEscapedPath, expectedQuery, accessToken)
134+
defer func() { _ = server.Close() }()
135+
136+
// unix:///path/to/socket#/http_path?query
137+
tokenURL := fmt.Sprintf("unix://%s#%s?%s", socketPath, expectedEscapedPath, expectedQuery)
138+
ts, err := NewTokenSourceFromURL(context.Background(), tokenURL, false)
139+
require.NoError(t, err)
140+
require.NotNil(t, ts)
141+
142+
token, err := ts.Token()
143+
assert.NoError(t, err)
144+
assert.Equal(t, accessToken, token.AccessToken)
145+
}
146+
147+
func Test_NewTokenSourceFromURL_UnixSocket_WithFragment_EscapedChars_Success(t *testing.T) {
148+
tmpFile, err := os.CreateTemp("", "gcsfuse-uds-test-*.sock")
149+
require.NoError(t, err)
150+
socketPath := tmpFile.Name()
151+
require.NoError(t, tmpFile.Close())
152+
require.NoError(t, os.Remove(socketPath))
153+
defer func() { _ = os.Remove(socketPath) }()
154+
155+
listener, err := net.Listen("unix", socketPath)
156+
require.NoError(t, err)
157+
defer func() { _ = listener.Close() }()
158+
159+
// Path contains escaped slash (%2F) and query contains escaped space (%20)
160+
expectedEscapedPath := "/computeMetadata%2Fv1%2Finstance%2Fservice-accounts%2Fdefault%2Ftoken"
161+
expectedQuery := "foo=bar%20baz"
162+
accessToken := "uds-access-token-escaped"
163+
164+
server := startFakeUDSTokenServer(t, listener, expectedEscapedPath, expectedQuery, accessToken)
165+
defer func() { _ = server.Close() }()
166+
167+
tokenURL := fmt.Sprintf("unix://%s#%s?%s", socketPath, expectedEscapedPath, expectedQuery)
168+
ts, err := NewTokenSourceFromURL(context.Background(), tokenURL, false)
169+
require.NoError(t, err)
170+
require.NotNil(t, ts)
171+
172+
token, err := ts.Token()
173+
assert.NoError(t, err)
174+
assert.Equal(t, accessToken, token.AccessToken)
175+
}
176+
177+
func Test_NewTokenSourceFromURL_UnixSocket_BackwardCompatibility_Success(t *testing.T) {
178+
tmpFile, err := os.CreateTemp("", "gcsfuse-uds-test-*.sock")
179+
require.NoError(t, err)
180+
socketPath := tmpFile.Name()
181+
require.NoError(t, tmpFile.Close())
182+
require.NoError(t, os.Remove(socketPath))
183+
defer func() { _ = os.Remove(socketPath) }()
184+
185+
listener, err := net.Listen("unix", socketPath)
186+
require.NoError(t, err)
187+
defer func() { _ = listener.Close() }()
188+
189+
expectedEscapedPath := "/"
190+
expectedQuery := "foo=bar"
191+
accessToken := "uds-access-token-compat"
192+
193+
server := startFakeUDSTokenServer(t, listener, expectedEscapedPath, expectedQuery, accessToken)
194+
defer func() { _ = server.Close() }()
195+
196+
// unix:///path/to/socket?query (old way, but with query)
197+
tokenURL := fmt.Sprintf("unix://%s?%s", socketPath, expectedQuery)
198+
ts, err := NewTokenSourceFromURL(context.Background(), tokenURL, false)
199+
require.NoError(t, err)
200+
require.NotNil(t, ts)
201+
202+
token, err := ts.Token()
203+
assert.NoError(t, err)
204+
assert.Equal(t, accessToken, token.AccessToken)
205+
}

0 commit comments

Comments
 (0)