From 6931b8672bc8517350d5bd3e3a35709020ab4290 Mon Sep 17 00:00:00 2001 From: Connor Peet Date: Tue, 29 Sep 2026 14:20:37 -0700 Subject: [PATCH] Forward additional headers on tunnel service requests Allow Rust management clients to configure static and per-request headers for management calls and relay handshakes. Forward TypeScript management headers to Node.js relay connections so VS Code can correlate each service request. Browser WebSocket handshakes cannot set custom headers. Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> --- rs/src/connections/relay_tunnel_client.rs | 42 ++++-- rs/src/connections/relay_tunnel_host.rs | 43 ++++-- rs/src/management/http_client.rs | 123 +++++++++++++++++- .../defaultTunnelRelayStreamFactory.ts | 4 +- ts/src/connections/tunnelConnectionSession.ts | 3 +- .../connections/tunnelRelayStreamFactory.ts | 3 + ts/src/management/tunnelManagementClient.ts | 7 + .../mocks/mockTunnelManagementClient.ts | 1 + .../mocks/mockTunnelRelayStreamFactory.ts | 3 + .../tunnels-test/tunnelHostAndClientTests.ts | 21 +++ ts/test/tunnels-test/tunnelManagementTests.ts | 44 +++++++ 11 files changed, 271 insertions(+), 23 deletions(-) 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 = [];