Skip to content

Commit cb393d2

Browse files
authored
Merge pull request #211 from estie-inc/fix/unify-snowflake-origin-user-agent
fix(api): send the connector User-Agent on every Snowflake origin request
2 parents 8b75849 + 2bee882 commit cb393d2

16 files changed

Lines changed: 589 additions & 405 deletions

File tree

src/api_context.rs

Lines changed: 95 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,95 @@
1+
use http::Method;
2+
use reqwest::{Url, header::USER_AGENT};
3+
4+
use crate::{Result, error::ConfigError};
5+
6+
pub(crate) const DEFAULT_USER_AGENT: &str =
7+
concat!(env!("CARGO_PKG_NAME"), "/", env!("CARGO_PKG_VERSION"));
8+
9+
/// The connector's authority to talk to one trusted Snowflake origin.
10+
///
11+
/// It represents what every Snowflake API request has in common regardless of endpoint: the origin it is addressed
12+
/// to and the connector it originates from.
13+
pub(crate) struct ApiContext {
14+
http: reqwest::Client,
15+
base_url: Url,
16+
}
17+
18+
impl ApiContext {
19+
pub(crate) fn new(http: reqwest::Client, base_url: Url) -> Self {
20+
Self { http, base_url }
21+
}
22+
23+
pub(crate) fn resolve(&self, relative: &str) -> Result<Url> {
24+
self.base_url
25+
.join(relative)
26+
.map_err(|e| ConfigError::invalid_url(e.to_string()).into())
27+
}
28+
29+
pub(crate) fn request(&self, method: Method, url: Url) -> reqwest::RequestBuilder {
30+
self.http
31+
.request(method, url)
32+
.header(USER_AGENT, DEFAULT_USER_AGENT)
33+
}
34+
35+
pub(crate) fn base_url(&self) -> &Url {
36+
&self.base_url
37+
}
38+
}
39+
40+
#[cfg(test)]
41+
mod tests {
42+
use super::*;
43+
use crate::ErrorKind;
44+
45+
fn context() -> ApiContext {
46+
ApiContext::new(
47+
reqwest::Client::new(),
48+
Url::parse("https://example.com/").expect("test base URL must parse"),
49+
)
50+
}
51+
52+
#[test]
53+
fn resolve_joins_relative_endpoint_against_base_url() {
54+
let url = context().resolve("session/v1/login-request").unwrap();
55+
assert_eq!(url.as_str(), "https://example.com/session/v1/login-request");
56+
}
57+
58+
#[test]
59+
fn resolve_invalid_url_is_config_error() {
60+
// An empty host makes the join fail; it must surface as the same config error the endpoints used before.
61+
let err = context().resolve("https://").unwrap_err();
62+
assert_eq!(err.kind(), ErrorKind::Config);
63+
}
64+
65+
#[test]
66+
fn request_carries_connector_user_agent() {
67+
let ctx = context();
68+
let url = ctx.resolve("queries/v1/query-request").unwrap();
69+
let request = ctx.request(Method::POST, url).build().unwrap();
70+
assert_eq!(
71+
request.headers().get(USER_AGENT).unwrap(),
72+
DEFAULT_USER_AGENT
73+
);
74+
}
75+
76+
#[test]
77+
fn request_does_not_attach_endpoint_headers() {
78+
let ctx = context();
79+
let url = ctx.resolve("queries/v1/query-request").unwrap();
80+
let request = ctx.request(Method::POST, url).build().unwrap();
81+
assert!(request.headers().get(reqwest::header::ACCEPT).is_none());
82+
assert!(
83+
request
84+
.headers()
85+
.get(reqwest::header::AUTHORIZATION)
86+
.is_none()
87+
);
88+
assert!(
89+
request
90+
.headers()
91+
.get(reqwest::header::CONTENT_TYPE)
92+
.is_none()
93+
);
94+
}
95+
}

src/auth/api.rs

Lines changed: 45 additions & 119 deletions
Original file line numberDiff line numberDiff line change
@@ -1,13 +1,11 @@
11
use 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

86
use 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)]
2523
pub(crate) struct AuthApiClient {
26-
shared: Arc<ClientShared>,
24+
api: Arc<ApiContext>,
2725
request_timeout: Duration,
2826
}
2927

3028
impl 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)]
11397
mod 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\ncontent-type: application/json\r\ncontent-length: {}\r\nconnection: 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\naccept: application/snowflake\r\n"));
242162
assert!(lowered.contains(&format!(
243163
"\r\nuser-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\naccept: application/json\r\n"));
280+
assert!(lowered.contains(&format!(
281+
"\r\nuser-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

Comments
 (0)