diff --git a/rs/src/connections/relay_tunnel_client.rs b/rs/src/connections/relay_tunnel_client.rs index 63ffa8b0..6f7f760c 100644 --- a/rs/src/connections/relay_tunnel_client.rs +++ b/rs/src/connections/relay_tunnel_client.rs @@ -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); @@ -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 { 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)?; diff --git a/rs/src/connections/relay_tunnel_host.rs b/rs/src/connections/relay_tunnel_host.rs index c3fddaa9..2acc532b 100644 --- a/rs/src/connections/relay_tunnel_host.rs +++ b/rs/src/connections/relay_tunnel_host.rs @@ -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); @@ -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"); @@ -1182,4 +1204,3 @@ mod tests { assert_eq!(advertised[0], public_key_b64); } } - diff --git a/rs/src/management/http_client.rs b/rs/src/management/http_client.rs index 7e47f880..253b4df0 100644 --- a/rs/src/management/http_client.rs +++ b/rs/src/management/http_client.rs @@ -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 Vec<(HeaderName, HeaderValue)> + Send + Sync>>, } const TUNNELS_API_PATH: &str = "/tunnels"; @@ -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, @@ -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 } @@ -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 Vec<(HeaderName, HeaderValue)> + Send + Sync>>, } /// Creates a new tunnel client builder. You can set options, then use `into()` @@ -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, } } @@ -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 for TunnelManagementClient { @@ -817,6 +859,8 @@ impl From 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, } } } @@ -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}, @@ -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}\"")) diff --git a/ts/src/connections/defaultTunnelRelayStreamFactory.ts b/ts/src/connections/defaultTunnelRelayStreamFactory.ts index 1ba6b7ab..b5d0c1af 100644 --- a/ts/src/connections/defaultTunnelRelayStreamFactory.ts +++ b/ts/src/connections/defaultTunnelRelayStreamFactory.ts @@ -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]; } diff --git a/ts/src/connections/tunnelConnectionSession.ts b/ts/src/connections/tunnelConnectionSession.ts index 2a99e22a..ed74cb46 100644 --- a/ts/src/connections/tunnelConnectionSession.ts +++ b/ts/src/connections/tunnelConnectionSession.ts @@ -330,7 +330,8 @@ export class TunnelConnectionSession extends TunnelConnectionBase implements Tun this.relayUri, this.connectionProtocols, this.accessToken, - clientConfig + clientConfig, + this.managementClient?.additionalRequestHeaders, ); this.trace( diff --git a/ts/src/connections/tunnelRelayStreamFactory.ts b/ts/src/connections/tunnelRelayStreamFactory.ts index 39656d94..a6cf5e55 100644 --- a/ts/src/connections/tunnelRelayStreamFactory.ts +++ b/ts/src/connections/tunnelRelayStreamFactory.ts @@ -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 }>; } diff --git a/ts/src/management/tunnelManagementClient.ts b/ts/src/management/tunnelManagementClient.ts index 3ea609e8..e67f8fe8 100644 --- a/ts/src/management/tunnelManagementClient.ts +++ b/ts/src/management/tunnelManagementClient.ts @@ -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. * diff --git a/ts/test/tunnels-test/mocks/mockTunnelManagementClient.ts b/ts/test/tunnels-test/mocks/mockTunnelManagementClient.ts index 0ba23d3d..facffe8a 100644 --- a/ts/test/tunnels-test/mocks/mockTunnelManagementClient.ts +++ b/ts/test/tunnels-test/mocks/mockTunnelManagementClient.ts @@ -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; diff --git a/ts/test/tunnels-test/mocks/mockTunnelRelayStreamFactory.ts b/ts/test/tunnels-test/mocks/mockTunnelRelayStreamFactory.ts index d0ed4de2..10aeab71 100644 --- a/ts/test/tunnels-test/mocks/mockTunnelRelayStreamFactory.ts +++ b/ts/test/tunnels-test/mocks/mockTunnelRelayStreamFactory.ts @@ -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, @@ -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 }); }; diff --git a/ts/test/tunnels-test/tunnelHostAndClientTests.ts b/ts/test/tunnels-test/tunnelHostAndClientTests.ts index e02d9c3d..5c50d1eb 100644 --- a/ts/test/tunnels-test/tunnelHostAndClientTests.ts +++ b/ts/test/tunnels-test/tunnelHostAndClientTests.ts @@ -502,6 +502,27 @@ export class TunnelHostAndClientTests { serverSession.dispose(); } + @test + public async forwardsTunnelHeadersToRelayHandshake() { + const managementClient = new MockTunnelManagementClient(); + managementClient.additionalRequestHeaders = { + 'X-Tunnels-VSCode-Session-Id': 'session-id', + 'X-Tunnels-VSCode-Client-Operation-Id': 'operation-id', + 'X-Tunnels-VSCode-Client-Request-Id': 'request-id', + }; + const relayClient = new TestTunnelRelayTunnelClient(managementClient); + const serverSession = await this.connectRelayClient({ + relayClient, + tunnel: this.createRelayTunnel(), + }); + + const factory = relayClient.streamFactory as MockTunnelRelayStreamFactory; + assert.deepStrictEqual(factory.lastAdditionalHeaders, managementClient.additionalRequestHeaders); + + relayClient.dispose(); + serverSession.dispose(); + } + @test async connectRelayClientWithStaleHostKey() { // A good tunnel with the correct host public key. diff --git a/ts/test/tunnels-test/tunnelManagementTests.ts b/ts/test/tunnels-test/tunnelManagementTests.ts index 5f63e639..bc58f226 100644 --- a/ts/test/tunnels-test/tunnelManagementTests.ts +++ b/ts/test/tunnels-test/tunnelManagementTests.ts @@ -129,6 +129,50 @@ export class TunnelManagementTests { assert(this.lastRequest.uri.includes('api-version=2023-09-27-preview')); } + @test + public async additionalRequestHeaders() { + this.managementClient.additionalRequestHeaders = { + 'X-Tunnels-VSCode-Session-Id': 'session-id', + 'X-Tunnels-VSCode-Client-Operation-Id': 'operation-id', + 'X-Tunnels-VSCode-Client-Request-Id': 'request-id', + }; + this.nextResponse = []; + + await this.managementClient.listUserLimits(); + + assert.deepStrictEqual(this.lastRequest?.config.headers, { + 'X-Tunnels-VSCode-Session-Id': 'session-id', + 'X-Tunnels-VSCode-Client-Operation-Id': 'operation-id', + 'X-Tunnels-VSCode-Client-Request-Id': 'request-id', + 'User-Agent': this.lastRequest?.config.headers?.['User-Agent'], + }); + } + + @test + public async generatesRequestHeadersForEachServiceRequest() { + let requestNumber = 0; + const headers: { [header: string]: string } = { + 'X-Tunnels-VSCode-Session-Id': 'session-id', + 'X-Tunnels-VSCode-Client-Operation-Id': 'operation-id', + }; + Object.defineProperty(headers, 'X-Tunnels-VSCode-Client-Request-Id', { + enumerable: true, + get: () => `request-${++requestNumber}`, + }); + this.managementClient.additionalRequestHeaders = headers; + this.nextResponse = []; + + await this.managementClient.listUserLimits(); + const firstRequestId = this.lastRequest?.config.headers?.['X-Tunnels-VSCode-Client-Request-Id']; + await this.managementClient.listUserLimits(); + const secondRequestId = this.lastRequest?.config.headers?.['X-Tunnels-VSCode-Client-Request-Id']; + + assert.strictEqual(this.lastRequest?.config.headers?.['X-Tunnels-VSCode-Session-Id'], 'session-id'); + assert.strictEqual(this.lastRequest?.config.headers?.['X-Tunnels-VSCode-Client-Operation-Id'], 'operation-id'); + assert.strictEqual(firstRequestId, 'request-1'); + assert.strictEqual(secondRequestId, 'request-2'); + } + @test public async timeoutServerResponse() { this.nextResponse = [];