@@ -17,24 +17,35 @@ package auth
1717import (
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
5061func 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