11use std:: { sync:: Arc , time:: Duration } ;
22
3- use reqwest:: {
4- Url ,
5- header:: { ACCEPT , USER_AGENT } ,
6- } ;
3+ use http:: Method ;
4+ use reqwest:: { Url , header:: ACCEPT } ;
75
86use crate :: {
9- ClientShared , Result ,
10- error:: { ConfigError , NetworkError , classify_request_error} ,
7+ ApiContext , Result ,
8+ error:: { NetworkError , classify_request_error} ,
119} ;
1210
1311#[ cfg( feature = "external-browser-sso" ) ]
@@ -23,32 +21,28 @@ const AUTHENTICATOR_REQUEST_ACCEPT: &str = "application/json";
2321
2422#[ derive( Clone ) ]
2523pub ( crate ) struct AuthApiClient {
26- shared : Arc < ClientShared > ,
24+ api : Arc < ApiContext > ,
2725 request_timeout : Duration ,
2826}
2927
3028impl AuthApiClient {
31- pub ( crate ) fn new ( shared : Arc < ClientShared > ) -> Self {
29+ pub ( crate ) fn new ( api : Arc < ApiContext > ) -> Self {
3230 Self {
33- shared ,
31+ api ,
3432 request_timeout : AUTH_REQUEST_TIMEOUT ,
3533 }
3634 }
3735
3836 #[ cfg( test) ]
39- fn with_request_timeout ( shared : Arc < ClientShared > , request_timeout : Duration ) -> Self {
37+ fn with_request_timeout ( api : Arc < ApiContext > , request_timeout : Duration ) -> Self {
4038 Self {
41- shared ,
39+ api ,
4240 request_timeout,
4341 }
4442 }
4543
4644 pub ( crate ) async fn login ( & self , request : LoginRequest < ' _ > ) -> Result < LoginSession > {
47- let url = self
48- . shared
49- . base_url
50- . join ( "session/v1/login-request" )
51- . map_err ( |e| ConfigError :: invalid_url ( e. to_string ( ) ) ) ?;
45+ let url = self . api . resolve ( "session/v1/login-request" ) ?;
5246
5347 let response = self
5448 . post ( url, LOGIN_REQUEST_ACCEPT )
@@ -72,11 +66,7 @@ impl AuthApiClient {
7266 & self ,
7367 request : AuthenticatorRequest < ' _ > ,
7468 ) -> Result < ExternalBrowserChallenge > {
75- let url = self
76- . shared
77- . base_url
78- . join ( "session/authenticator-request" )
79- . map_err ( |e| ConfigError :: invalid_url ( e. to_string ( ) ) ) ?;
69+ let url = self . api . resolve ( "session/authenticator-request" ) ?;
8070
8171 let body = request. into_body ( ClientEnvironment :: auth_defaults ( self . request_timeout ) ) ;
8272 let response = self
@@ -96,108 +86,38 @@ impl AuthApiClient {
9686 }
9787
9888 fn post ( & self , url : Url , accept : & ' static str ) -> reqwest:: RequestBuilder {
99- self . shared
100- . http
101- . post ( url)
89+ self . api
90+ . request ( Method :: POST , url)
10291 . header ( ACCEPT , accept)
103- . header ( USER_AGENT , default_user_agent ( ) )
10492 . timeout ( self . request_timeout )
10593 }
10694}
10795
108- fn default_user_agent ( ) -> String {
109- format ! ( "{}/{}" , env!( "CARGO_PKG_NAME" ) , env!( "CARGO_PKG_VERSION" ) )
110- }
111-
11296#[ cfg( test) ]
11397mod tests {
114- use std:: { io , net:: SocketAddr , time:: Duration } ;
98+ use std:: { net:: SocketAddr , time:: Duration } ;
11599
100+ use http:: StatusCode ;
116101 use serde_json:: json;
117- use tokio:: {
118- io:: { AsyncReadExt , AsyncWriteExt } ,
119- net:: { TcpListener , TcpStream } ,
120- } ;
102+ use tokio:: net:: TcpListener ;
121103
122104 use super :: * ;
123105 use crate :: {
124- ClientSharedPartial , ErrorKind ,
106+ ErrorKind ,
107+ api_context:: DEFAULT_USER_AGENT ,
125108 auth:: wire:: { LoginBody , LoginCredentialWire , LoginData , LoginQuery , LoginRequest } ,
109+ test_support:: http:: { base_url, read_http_request, write_json_response} ,
126110 } ;
127111
128- fn auth_client ( addr : SocketAddr ) -> AuthApiClient {
129- AuthApiClient :: new (
130- ClientSharedPartial :: new ( )
131- . with_base_url ( base_url_for ( addr) )
132- . build ( ) ,
133- )
134- }
135-
136112 #[ cfg( feature = "external-browser-sso" ) ]
137113 use crate :: auth:: wire:: AuthenticatorRequest ;
138114
139- async fn read_http_message ( stream : & mut TcpStream ) -> io:: Result < String > {
140- let mut buf = Vec :: new ( ) ;
141- let mut header_end = None ;
142-
143- while header_end. is_none ( ) {
144- let mut chunk = [ 0_u8 ; 1024 ] ;
145- let n = stream. read ( & mut chunk) . await ?;
146- if n == 0 {
147- return Err ( io:: Error :: new (
148- io:: ErrorKind :: UnexpectedEof ,
149- "stream closed before headers completed" ,
150- ) ) ;
151- }
152- buf. extend_from_slice ( & chunk[ ..n] ) ;
153- header_end = buf. windows ( 4 ) . position ( |w| w == b"\r \n \r \n " ) ;
154- }
155-
156- let header_end = header_end. expect ( "header end must exist" ) + 4 ;
157- let headers = String :: from_utf8_lossy ( & buf[ ..header_end] ) ;
158- let content_length = headers
159- . lines ( )
160- . find_map ( |line| {
161- let ( name, value) = line. split_once ( ':' ) ?;
162- name. eq_ignore_ascii_case ( "content-length" )
163- . then ( || value. trim ( ) . parse :: < usize > ( ) . ok ( ) )
164- . flatten ( )
165- } )
166- . unwrap_or ( 0 ) ;
167-
168- while buf. len ( ) < header_end + content_length {
169- let mut chunk = [ 0_u8 ; 1024 ] ;
170- let n = stream. read ( & mut chunk) . await ?;
171- if n == 0 {
172- return Err ( io:: Error :: new (
173- io:: ErrorKind :: UnexpectedEof ,
174- "stream closed before body completed" ,
175- ) ) ;
176- }
177- buf. extend_from_slice ( & chunk[ ..n] ) ;
178- }
179-
180- Ok ( String :: from_utf8_lossy ( & buf) . into_owned ( ) )
181- }
182-
183- async fn write_http_response (
184- stream : & mut TcpStream ,
185- status : & str ,
186- body : & str ,
187- ) -> io:: Result < ( ) > {
188- stream
189- . write_all (
190- format ! (
191- "HTTP/1.1 {status}\r \n content-type: application/json\r \n content-length: {}\r \n connection: close\r \n \r \n {body}" ,
192- body. len( )
193- )
194- . as_bytes ( ) ,
195- )
196- . await
115+ fn test_api_context ( http : reqwest:: Client , base_url : Url ) -> Arc < ApiContext > {
116+ Arc :: new ( ApiContext :: new ( http, base_url) )
197117 }
198118
199- fn base_url_for ( addr : SocketAddr ) -> Url {
200- Url :: parse ( & format ! ( "http://{ addr}/" ) ) . expect ( "test base URL must parse" )
119+ fn auth_client ( addr : SocketAddr ) -> AuthApiClient {
120+ AuthApiClient :: new ( test_api_context ( reqwest :: Client :: new ( ) , base_url ( addr) ) )
201121 }
202122
203123 fn sample_login_request < ' a > ( ) -> LoginRequest < ' a > {
@@ -229,7 +149,7 @@ mod tests {
229149 let addr = listener. local_addr ( ) . unwrap ( ) ;
230150 let server = tokio:: spawn ( async move {
231151 let ( mut socket, _) = listener. accept ( ) . await . unwrap ( ) ;
232- let request = read_http_message ( & mut socket) . await . unwrap ( ) ;
152+ let request = read_http_request ( & mut socket) . await . unwrap ( ) ;
233153 assert ! (
234154 request. starts_with(
235155 "POST /session/v1/login-request?warehouse=warehouse&databaseName=database&schemaName=schema&roleName=role HTTP/1.1"
@@ -241,7 +161,7 @@ mod tests {
241161 assert ! ( lowered. contains( "\r \n accept: application/snowflake\r \n " ) ) ;
242162 assert ! ( lowered. contains( & format!(
243163 "\r \n user-agent: {}\r \n " ,
244- default_user_agent ( ) . to_ascii_lowercase( )
164+ DEFAULT_USER_AGENT . to_ascii_lowercase( )
245165 ) ) ) ;
246166
247167 let body = request
@@ -260,9 +180,9 @@ mod tests {
260180 } )
261181 ) ;
262182
263- write_http_response (
183+ write_json_response (
264184 & mut socket,
265- "200 OK" ,
185+ StatusCode :: OK ,
266186 r#"{"success":true,"data":{"token":"session-token"}}"# ,
267187 )
268188 . await
@@ -283,10 +203,14 @@ mod tests {
283203 let addr = listener. local_addr ( ) . unwrap ( ) ;
284204 let server = tokio:: spawn ( async move {
285205 let ( mut socket, _) = listener. accept ( ) . await . unwrap ( ) ;
286- let _ = read_http_message ( & mut socket) . await . unwrap ( ) ;
287- write_http_response ( & mut socket, "401 Unauthorized" , r#"{"message":"nope"}"# )
288- . await
289- . unwrap ( ) ;
206+ let _ = read_http_request ( & mut socket) . await . unwrap ( ) ;
207+ write_json_response (
208+ & mut socket,
209+ StatusCode :: UNAUTHORIZED ,
210+ r#"{"message":"nope"}"# ,
211+ )
212+ . await
213+ . unwrap ( ) ;
290214 } ) ;
291215
292216 let client = auth_client ( addr) ;
@@ -303,8 +227,8 @@ mod tests {
303227 let addr = listener. local_addr ( ) . unwrap ( ) ;
304228 let server = tokio:: spawn ( async move {
305229 let ( mut socket, _) = listener. accept ( ) . await . unwrap ( ) ;
306- let _ = read_http_message ( & mut socket) . await . unwrap ( ) ;
307- write_http_response ( & mut socket, "200 OK" , "not-json" )
230+ let _ = read_http_request ( & mut socket) . await . unwrap ( ) ;
231+ write_json_response ( & mut socket, StatusCode :: OK , "not-json" )
308232 . await
309233 . unwrap ( ) ;
310234 } ) ;
@@ -327,9 +251,7 @@ mod tests {
327251 } ) ;
328252
329253 let client = AuthApiClient :: with_request_timeout (
330- ClientSharedPartial :: new ( )
331- . with_base_url ( base_url_for ( addr) )
332- . build ( ) ,
254+ test_api_context ( reqwest:: Client :: new ( ) , base_url ( addr) ) ,
333255 Duration :: from_millis ( 50 ) ,
334256 ) ;
335257 let err = client. login ( sample_login_request ( ) ) . await . unwrap_err ( ) ;
@@ -348,13 +270,17 @@ mod tests {
348270 let addr = listener. local_addr ( ) . unwrap ( ) ;
349271 let server = tokio:: spawn ( async move {
350272 let ( mut socket, _) = listener. accept ( ) . await . unwrap ( ) ;
351- let request = read_http_message ( & mut socket) . await . unwrap ( ) ;
273+ let request = read_http_request ( & mut socket) . await . unwrap ( ) ;
352274 assert ! (
353275 request. starts_with( "POST /session/authenticator-request HTTP/1.1" ) ,
354276 "{request}"
355277 ) ;
356278 let lowered = request. to_ascii_lowercase ( ) ;
357279 assert ! ( lowered. contains( "\r \n accept: application/json\r \n " ) ) ;
280+ assert ! ( lowered. contains( & format!(
281+ "\r \n user-agent: {}\r \n " ,
282+ DEFAULT_USER_AGENT . to_ascii_lowercase( )
283+ ) ) ) ;
358284
359285 let body = request
360286 . split ( "\r \n \r \n " )
@@ -380,9 +306,9 @@ mod tests {
380306 } )
381307 ) ;
382308
383- write_http_response (
309+ write_json_response (
384310 & mut socket,
385- "200 OK" ,
311+ StatusCode :: OK ,
386312 r#"{"success":true,"data":{"ssoUrl":"https://example.com/sso","proofKey":"proof-key"}}"# ,
387313 )
388314 . await
0 commit comments