Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
42 changes: 33 additions & 9 deletions rs/src/connections/relay_tunnel_client.rs
Original file line number Diff line number Diff line change
Expand Up @@ -81,14 +81,36 @@ impl RelayTunnelClient {
.as_deref()
.ok_or(TunnelError::MissingClientEndpoint)?;

let req = build_websocket_request(
client_relay_uri,
&[
("Sec-WebSocket-Protocol", "tunnel-relay-client"),
("Authorization", &format!("tunnel {}", access_token)),
("User-Agent", self.mgmt.user_agent.to_str().unwrap()),
],
)?;
let mut headers = vec![
(
"Sec-WebSocket-Protocol".to_string(),
"tunnel-relay-client".to_string(),
),
(
"Authorization".to_string(),
format!("tunnel {}", access_token),
),
(
"User-Agent".to_string(),
self.mgmt.user_agent.to_str().unwrap().to_string(),
),
];
headers.extend(
self.mgmt
.request_headers()
.into_iter()
.filter_map(|(name, value)| {
value
.to_str()
.ok()
.map(|value| (name.as_str().to_string(), value.to_string()))
}),
);
let header_refs: Vec<(&str, &str)> = headers
.iter()
.map(|(name, value)| (name.as_str(), value.as_str()))
.collect();
let req = build_websocket_request(client_relay_uri, &header_refs)?;

let cnx = if let Some(proxy) = &self.proxy {
log::debug!("connecting via http_proxy on {}", proxy);
Expand Down Expand Up @@ -182,7 +204,9 @@ impl ClientRelayHandle {
/// The listener will be stopped when `close()` is called on this handle.
pub async fn forward_port_locally(&self, port: u16) -> Result<SocketAddr, TunnelError> {
let addr = SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), port);
let listener = TcpListener::bind(addr).await.map_err(TunnelError::ErrorListeningOnAddress)?;
let listener = TcpListener::bind(addr)
.await
.map_err(TunnelError::ErrorListeningOnAddress)?;
let local_addr = listener
.local_addr()
.map_err(TunnelError::ErrorListeningOnAddress)?;
Expand Down
43 changes: 32 additions & 11 deletions rs/src/connections/relay_tunnel_host.rs
Original file line number Diff line number Diff line change
Expand Up @@ -427,14 +427,36 @@ impl RelayTunnelHost {
.as_deref()
.ok_or(TunnelError::MissingHostEndpoint)?;

let req = build_websocket_request(
url,
&[
("Sec-WebSocket-Protocol", "tunnel-relay-host"),
("Authorization", &format!("tunnel {}", host_token)),
("User-Agent", self.mgmt.user_agent.to_str().unwrap()),
],
)?;
let mut headers = vec![
(
"Sec-WebSocket-Protocol".to_string(),
"tunnel-relay-host".to_string(),
),
(
"Authorization".to_string(),
format!("tunnel {}", host_token),
),
(
"User-Agent".to_string(),
self.mgmt.user_agent.to_str().unwrap().to_string(),
),
];
headers.extend(
self.mgmt
.request_headers()
.into_iter()
.filter_map(|(name, value)| {
value
.to_str()
.ok()
.map(|value| (name.as_str().to_string(), value.to_string()))
}),
);
let header_refs: Vec<(&str, &str)> = headers
.iter()
.map(|(name, value)| (name.as_str(), value.as_str()))
.collect();
let req = build_websocket_request(url, &header_refs)?;

let cnx = if let Some(proxy) = &self.proxy {
log::debug!("connecting via http_proxy on {}", proxy);
Expand Down Expand Up @@ -1166,8 +1188,8 @@ mod tests {
fn encodes_rsa_public_key_in_canonical_ssh_wire_form() {
// Generate an RSA keypair with a non-default hash variant — this is
// what `RelayTunnelHost::new` does (SHA2_512).
let keypair = KeyPair::generate_rsa(2048, SignatureHash::SHA2_512)
.expect("generate rsa keypair");
let keypair =
KeyPair::generate_rsa(2048, SignatureHash::SHA2_512).expect("generate rsa keypair");

let advertised = encode_host_public_keys(&keypair);
assert_eq!(advertised.len(), 1, "expected exactly one host public key");
Expand All @@ -1182,4 +1204,3 @@ mod tests {
assert_eq!(advertised[0], public_key_b64);
}
}

123 changes: 122 additions & 1 deletion rs/src/management/http_client.rs
Original file line number Diff line number Diff line change
Expand Up @@ -33,6 +33,9 @@ pub struct TunnelManagementClient {
environment: TunnelServiceProperties,
api_version: String,
is_custom_domain: bool,
additional_headers: Vec<(HeaderName, HeaderValue)>,
additional_headers_provider:
Option<Arc<dyn Fn() -> Vec<(HeaderName, HeaderValue)> + Send + Sync>>,
}

const TUNNELS_API_PATH: &str = "/tunnels";
Expand All @@ -57,9 +60,25 @@ impl TunnelManagementClient {
environment: self.environment.clone(),
api_version: self.api_version.clone(),
is_custom_domain: self.is_custom_domain,
additional_headers: self.additional_headers.clone(),
additional_headers_provider: self.additional_headers_provider.clone(),
}
}

/// Gets headers applied to every tunnel service request.
pub fn additional_headers(&self) -> &[(HeaderName, HeaderValue)] {
&self.additional_headers
}

/// Gets headers applied to this tunnel service request.
pub fn request_headers(&self) -> Vec<(HeaderName, HeaderValue)> {
let mut headers = self.additional_headers.clone();
if let Some(provider) = &self.additional_headers_provider {
headers.extend(provider());
}
headers
}

/// Lists tunnels owned by the user.
pub async fn list_all_tunnels(
&self,
Expand Down Expand Up @@ -595,6 +614,9 @@ impl TunnelManagementClient {
for (name, value) in &tunnel_opts.headers {
headers.append(name, value.to_owned());
}
for (name, value) in self.request_headers() {
headers.append(name, value);
}

request
}
Expand Down Expand Up @@ -715,6 +737,9 @@ pub struct TunnelClientBuilder {
environment: TunnelServiceProperties,
api_version: String,
is_custom_domain: bool,
additional_headers: Vec<(HeaderName, HeaderValue)>,
additional_headers_provider:
Option<Arc<dyn Fn() -> Vec<(HeaderName, HeaderValue)> + Send + Sync>>,
}

/// Creates a new tunnel client builder. You can set options, then use `into()`
Expand All @@ -731,6 +756,8 @@ pub fn new_tunnel_management(user_agent: &str) -> TunnelClientBuilder {
environment: env_production(),
api_version: API_VERSIONS[0].to_owned(),
is_custom_domain: false,
additional_headers: Vec::new(),
additional_headers_provider: None,
}
}

Expand Down Expand Up @@ -806,6 +833,21 @@ impl TunnelClientBuilder {
self.environment = environment;
self
}

/// Adds headers to every tunnel service request and relay WebSocket handshake.
pub fn additional_headers(&mut self, headers: Vec<(HeaderName, HeaderValue)>) -> &mut Self {
self.additional_headers = headers;
self
}

/// Provides headers for each tunnel service request and relay WebSocket handshake.
pub fn additional_headers_provider(
&mut self,
provider: impl Fn() -> Vec<(HeaderName, HeaderValue)> + Send + Sync + 'static,
) -> &mut Self {
self.additional_headers_provider = Some(Arc::new(provider));
self
}
}

impl From<TunnelClientBuilder> for TunnelManagementClient {
Expand All @@ -817,6 +859,8 @@ impl From<TunnelClientBuilder> for TunnelManagementClient {
environment: builder.environment,
api_version: builder.api_version,
is_custom_domain: builder.is_custom_domain,
additional_headers: builder.additional_headers,
additional_headers_provider: builder.additional_headers_provider,
}
}
}
Expand Down Expand Up @@ -989,7 +1033,10 @@ mod tests {
};

use regex::Regex;
use reqwest::Url;
use reqwest::{
header::{HeaderName, HeaderValue},
Url,
};
use tokio::{
io::{AsyncReadExt, AsyncWriteExt},
net::{TcpListener, TcpStream},
Expand Down Expand Up @@ -1294,6 +1341,80 @@ mod tests {
);
}

#[tokio::test]
async fn applies_additional_headers_to_tunnel_requests() {
let server = TestServer::start(vec![TestResponse::ok("{\"value\":[]}".to_string())]).await;
let client = test_client(&server, Authorization::Anonymous);
let mut builder = client.build();
builder.additional_headers(vec![
(
HeaderName::from_static("x-tunnels-vscode-session-id"),
HeaderValue::from_static("session-id"),
),
(
HeaderName::from_static("x-tunnels-vscode-client-operation-id"),
HeaderValue::from_static("operation-id"),
),
(
HeaderName::from_static("x-tunnels-vscode-client-request-id"),
HeaderValue::from_static("request-id"),
),
]);
let client: super::TunnelManagementClient = builder.into();

client.list_all_tunnels(NO_REQUEST_OPTIONS).await.unwrap();

let requests = server.finish().await;
assert_eq!(
requests[0].header_values("x-tunnels-vscode-session-id"),
["session-id"]
);
assert_eq!(
requests[0].header_values("x-tunnels-vscode-client-operation-id"),
["operation-id"]
);
assert_eq!(
requests[0].header_values("x-tunnels-vscode-client-request-id"),
["request-id"]
);
}

#[tokio::test]
async fn generates_additional_headers_for_each_tunnel_request() {
let server = TestServer::start(vec![
TestResponse::ok("{\"value\":[]}".to_string()),
TestResponse::ok("{\"value\":[]}".to_string()),
])
.await;
let client = test_client(&server, Authorization::Anonymous);
let mut builder = client.build();
let request_number = Arc::new(std::sync::atomic::AtomicUsize::new(0));
builder.additional_headers_provider({
let request_number = request_number.clone();
move || {
let request_id = request_number.fetch_add(1, std::sync::atomic::Ordering::SeqCst);
vec![(
HeaderName::from_static("x-tunnels-vscode-client-request-id"),
HeaderValue::from_str(&format!("request-{request_id}")).unwrap(),
)]
}
});
let client: super::TunnelManagementClient = builder.into();

client.list_all_tunnels(NO_REQUEST_OPTIONS).await.unwrap();
client.list_all_tunnels(NO_REQUEST_OPTIONS).await.unwrap();

let requests = server.finish().await;
assert_eq!(
requests[0].header_values("x-tunnels-vscode-client-request-id"),
["request-0"]
);
assert_eq!(
requests[1].header_values("x-tunnels-vscode-client-request-id"),
["request-1"]
);
}

fn recommendation_response(recommended_cluster_id: Option<&str>) -> String {
let recommended_cluster_id = recommended_cluster_id
.map(|value| format!("\"{value}\""))
Expand Down
4 changes: 3 additions & 1 deletion ts/src/connections/defaultTunnelRelayStreamFactory.ts
Original file line number Diff line number Diff line change
Expand Up @@ -15,19 +15,21 @@ export class DefaultTunnelRelayStreamFactory implements TunnelRelayStreamFactory
protocols: string[],
accessToken?: string,
clientConfig?: IClientConfig,
additionalHeaders?: { [header: string]: string },
): Promise<{ stream: Stream, protocol: string }> {
if (isNode()) {
const stream = await SshHelpers.openConnection(
relayUri,
protocols,
{
...additionalHeaders,
...(accessToken && { Authorization: `tunnel ${accessToken}` }),
},
clientConfig,
);
return { stream, protocol: stream.protocol! };
} else {
// Web sockets don't support auth. Authenticate TunnelRelay by sending accessToken as a subprotocol.
// Browser WebSocket APIs cannot set auth or custom handshake headers.
if (accessToken) {
protocols = [...protocols, accessToken];
}
Expand Down
3 changes: 2 additions & 1 deletion ts/src/connections/tunnelConnectionSession.ts
Original file line number Diff line number Diff line change
Expand Up @@ -330,7 +330,8 @@ export class TunnelConnectionSession extends TunnelConnectionBase implements Tun
this.relayUri,
this.connectionProtocols,
this.accessToken,
clientConfig
clientConfig,
this.managementClient?.additionalRequestHeaders,
);

this.trace(
Expand Down
3 changes: 3 additions & 0 deletions ts/src/connections/tunnelRelayStreamFactory.ts
Original file line number Diff line number Diff line change
Expand Up @@ -14,11 +14,14 @@ export interface TunnelRelayStreamFactory {
* @param protocols Array of supported connection protocols (websocket sub-protocols).
* @param accessToken Tunnel host access token, or null if anonymous.
* @param clientConfig Client config for websocket.
* @param additionalHeaders Additional headers for the Node.js relay WebSocket request.
* Browser WebSocket APIs cannot set custom handshake headers.
*/
createRelayStream(
relayUri: string,
protocols: string[],
accessToken?: string,
clientConfig?: IClientConfig,
additionalHeaders?: { [header: string]: string },
): Promise<{ stream: Stream, protocol: string }>;
}
7 changes: 7 additions & 0 deletions ts/src/management/tunnelManagementClient.ts
Original file line number Diff line number Diff line change
Expand Up @@ -24,6 +24,13 @@ export interface TunnelManagementClient {
*/
httpsAgent?: https.Agent;

/**
* Additional headers included in every tunnel service request, including
* Node.js relay WebSocket connection requests made with this client.
* Browser WebSocket APIs cannot set custom relay handshake headers.
*/
additionalRequestHeaders?: { [header: string]: string };

/**
* Lists tunnels that are owned by the caller.
*
Expand Down
1 change: 1 addition & 0 deletions ts/test/tunnels-test/mocks/mockTunnelManagementClient.ts
Original file line number Diff line number Diff line change
Expand Up @@ -17,6 +17,7 @@ import {

export class MockTunnelManagementClient implements TunnelManagementClient {
private idCounter: number = 0;
public additionalRequestHeaders?: { [header: string]: string };
public tunnels: Tunnel[] = [];
public hostRelayUri?: string;
public clientRelayUri?: string;
Expand Down
3 changes: 3 additions & 0 deletions ts/test/tunnels-test/mocks/mockTunnelRelayStreamFactory.ts
Original file line number Diff line number Diff line change
Expand Up @@ -14,6 +14,7 @@ import { IClientConfig } from 'websocket';
export class MockTunnelRelayStreamFactory implements TunnelRelayStreamFactory {
private readonly connectionType: string;
private readonly stream: Stream;
public lastAdditionalHeaders?: { [header: string]: string };

constructor(
connectionType: string,
Expand All @@ -32,10 +33,12 @@ export class MockTunnelRelayStreamFactory implements TunnelRelayStreamFactory {
protocols: string[],
accessToken?: string,
clientConfig?: IClientConfig,
additionalHeaders?: { [header: string]: string },
) => {
if (!relayUri || !accessToken || !protocols.includes(this.connectionType)) {
throw new Error('Invalid params');
}
this.lastAdditionalHeaders = additionalHeaders;
return Promise.resolve({ stream: this.stream, protocol: this.connectionType });
};

Expand Down
Loading
Loading