diff --git a/crates/jmux-proxy/src/lib.rs b/crates/jmux-proxy/src/lib.rs index 7373fce85..acf3dbf8f 100644 --- a/crates/jmux-proxy/src/lib.rs +++ b/crates/jmux-proxy/src/lib.rs @@ -12,7 +12,10 @@ mod id_allocator; use std::collections::{HashMap, HashSet}; use std::convert::TryFrom; +use std::future::Future; use std::io; +use std::net::IpAddr; +use std::pin::Pin; use std::sync::Arc; use std::sync::atomic::{AtomicBool, AtomicU64, AtomicUsize, Ordering}; use std::time::SystemTime; @@ -22,7 +25,6 @@ use bytes::Bytes; use jmux_proto::{ChannelData, DistantChannelId, Header, LocalChannelId, Message, ReasonCode}; use tokio::io::{AsyncRead, AsyncWrite, AsyncWriteExt}; use tokio::net::TcpStream; -use tokio::net::tcp::{OwnedReadHalf, OwnedWriteHalf}; use tokio::sync::{Notify, mpsc, oneshot}; use tokio::task::JoinHandle; use tokio_util::codec::FramedRead; @@ -53,6 +55,26 @@ pub type ApiResponseReceiver = oneshot::Receiver; pub type ApiRequestSender = mpsc::Sender; pub type ApiRequestReceiver = mpsc::Receiver; +trait TargetStream: AsyncRead + AsyncWrite + Unpin + Send {} + +impl TargetStream for T where T: AsyncRead + AsyncWrite + Unpin + Send {} + +type ErasedTargetStream = Box; +type TargetConnectorFuture = Pin>> + Send>>; +type TargetConnector = Arc TargetConnectorFuture + Send + Sync>; + +pub struct ConnectedTarget { + stream: ErasedTargetStream, +} + +impl ConnectedTarget { + pub fn new(stream: impl AsyncRead + AsyncWrite + Unpin + Send + 'static) -> Self { + Self { + stream: Box::new(stream), + } + } +} + #[derive(Debug)] pub enum JmuxApiRequest { OpenChannel { @@ -84,6 +106,7 @@ pub struct JmuxProxy { jmux_reader: Box, jmux_writer: Box, traffic_callback: Option, + target_connector: Option, } impl JmuxProxy { @@ -98,6 +121,7 @@ impl JmuxProxy { jmux_reader, jmux_writer, traffic_callback: None, + target_connector: None, } } @@ -113,6 +137,19 @@ impl JmuxProxy { self } + /// Tries a custom target connection before falling back to direct TCP. + /// + /// Return `Ok(None)` when the target should use the default direct connection. + #[must_use] + pub fn with_target_connector(mut self, connector: C) -> Self + where + C: Fn(DestinationUrl) -> F + Send + Sync + 'static, + F: Future>> + Send + 'static, + { + self.target_connector = Some(Arc::new(move |destination_url| Box::pin(connector(destination_url)))); + self + } + /// Configures an outgoing-traffic callback for lifecycle event monitoring. /// /// The provided callback will be invoked exactly once per outgoing stream at the end of its @@ -186,6 +223,7 @@ async fn run_proxy_impl(proxy: JmuxProxy, span: Span) -> anyhow::Result<()> { jmux_reader, jmux_writer, traffic_callback, + target_connector, } = proxy; let (msg_to_send_tx, msg_to_send_rx) = mpsc::channel::(JMUX_MESSAGE_MPSC_CHANNEL_SIZE); @@ -206,6 +244,7 @@ async fn run_proxy_impl(proxy: JmuxProxy, span: Span) -> anyhow::Result<()> { msg_to_send_tx, api_request_rx, traffic_callback, + target_connector, parent_span: span, } .spawn(); @@ -255,7 +294,7 @@ struct JmuxChannelCtx { // Traffic audit metadata target_host: String, /// Target server resolved address IP - target_ip: Option, + target_ip: Option, /// Target server port target_port: u16, /// Time the connection with target peer was established at @@ -344,7 +383,6 @@ type DataReceiver = mpsc::Receiver; type DataSender = mpsc::Sender; type InternalMessageSender = mpsc::Sender; -#[derive(Debug)] enum InternalMessage { Eof { id: LocalChannelId, @@ -353,7 +391,11 @@ enum InternalMessage { // Boxing reduces enum size from 224 bytes to ~16 bytes // (clippy::large_enum_variant) channel: Box, - stream: TcpStream, + stream: ErasedTargetStream, + }, + TargetConnectionFailed { + id: LocalChannelId, + distant_id: DistantChannelId, }, AbnormalTermination { id: LocalChannelId, @@ -427,6 +469,7 @@ struct JmuxSchedulerTask { msg_to_send_tx: MessageSender, api_request_rx: ApiRequestReceiver, traffic_callback: Option, + target_connector: Option, parent_span: Span, } @@ -448,6 +491,7 @@ async fn scheduler_task_impl(task: JmuxSc msg_to_send_tx, mut api_request_rx, traffic_callback, + target_connector, parent_span, } = task; @@ -501,7 +545,8 @@ async fn scheduler_task_impl(task: JmuxSc error!(%error, "Couldn't send leftover bytes"); } - let (reader, writer) = stream.into_split(); + let stream = Box::new(stream) as ErasedTargetStream; + let (reader, writer) = tokio::io::split(stream); DataWriterTask { writer, @@ -626,7 +671,7 @@ async fn scheduler_task_impl(task: JmuxSc debug!("Channel accepted"); }); - let (reader, writer) = stream.into_split(); + let (reader, writer) = tokio::io::split(stream); DataWriterTask { writer, @@ -652,6 +697,17 @@ async fn scheduler_task_impl(task: JmuxSc .spawn(channel_span) .detach(); } + InternalMessage::TargetConnectionFailed { id, distant_id } => { + jmux_ctx.id_allocator.free(id); + msg_to_send_tx + .send(Message::open_failure( + distant_id, + ReasonCode::GENERAL_FAILURE, + "target connection failed", + )) + .await + .context("couldn't send OPEN FAILURE message through mpsc channel")?; + } } } msg = jmux_stream.next() => { @@ -760,6 +816,7 @@ async fn scheduler_task_impl(task: JmuxSc internal_msg_tx: internal_msg_tx.clone(), msg_to_send_tx: msg_to_send_tx.clone(), traffic_callback: traffic_callback.clone(), + target_connector: target_connector.clone(), } .spawn() .detach(); @@ -958,7 +1015,7 @@ async fn scheduler_task_impl(task: JmuxSc // ---------------------- // struct DataReaderTask { - reader: OwnedReadHalf, + reader: tokio::io::ReadHalf, local_id: LocalChannelId, distant_id: DistantChannelId, window_size_updated: Arc, @@ -1087,7 +1144,7 @@ impl DataReaderTask { // ---------------------- // struct DataWriterTask { - writer: OwnedWriteHalf, + writer: tokio::io::WriteHalf, data_rx: DataReceiver, /// Tracks bytes written into the stream. bytes_tx: Arc, @@ -1123,6 +1180,8 @@ impl DataWriterTask { bytes_tx.fetch_add(data.len() as u64, Ordering::SeqCst); } + + let _ = writer.shutdown().await; } .instrument(span), ); @@ -1139,6 +1198,7 @@ struct StreamResolverTask { internal_msg_tx: InternalMessageSender, msg_to_send_tx: MessageSender, traffic_callback: Option, + target_connector: Option, } impl StreamResolverTask { @@ -1164,6 +1224,7 @@ impl StreamResolverTask { internal_msg_tx, msg_to_send_tx, traffic_callback, + target_connector, } = self; let scheme = destination_url.scheme(); @@ -1172,6 +1233,42 @@ impl StreamResolverTask { match scheme { "tcp" => { + if let Some(connector) = target_connector { + match connector(destination_url.clone()).await { + Ok(Some(ConnectedTarget { stream })) => { + channel.connect_at = SystemTime::now(); + + internal_msg_tx + .send(InternalMessage::StreamResolved { + channel: Box::new(channel), + stream, + }) + .await + .map_err(|_| { + anyhow::anyhow!("couldn't send back resolved stream through internal mpsc channel") + })?; + + return Ok(()); + } + Ok(None) => {} + Err(error) => { + internal_msg_tx + .send(InternalMessage::TargetConnectionFailed { + id: channel.local_id, + distant_id: channel.distant_id, + }) + .await + .map_err(|_| { + anyhow::anyhow!( + "couldn't report target connection failure through internal mpsc channel" + ) + })?; + + return Err(error.context(format!("couldn't connect to {host}:{port}"))); + } + } + } + // Perform DNS resolution first to get concrete IP addresses. let socket_addrs = match tokio::net::lookup_host((host, port)).await { Ok(addrs) => addrs, @@ -1203,10 +1300,12 @@ impl StreamResolverTask { internal_msg_tx .send(InternalMessage::StreamResolved { channel: Box::new(channel), - stream, + stream: Box::new(stream), }) .await - .context("couldn't send back resolved stream through internal mpsc channel")?; + .map_err(|_| { + anyhow::anyhow!("couldn't send back resolved stream through internal mpsc channel") + })?; return Ok(()); } diff --git a/crates/jmux-proxy/tests/target_connector.rs b/crates/jmux-proxy/tests/target_connector.rs new file mode 100644 index 000000000..885aebd68 --- /dev/null +++ b/crates/jmux-proxy/tests/target_connector.rs @@ -0,0 +1,131 @@ +#![allow(unused_crate_dependencies)] +#![allow(clippy::unwrap_used)] + +use std::time::Duration; + +use jmux_proto::{BytesMut, DistantChannelId, Header, LocalChannelId, Message, ReasonCode}; +use jmux_proxy::{ConnectedTarget, DestinationUrl, JmuxConfig, JmuxProxy}; +use tokio::io::{AsyncRead, AsyncReadExt, AsyncWrite, AsyncWriteExt}; +use tokio::time::timeout; + +const TEST_TIMEOUT: Duration = Duration::from_secs(5); + +async fn send_message(writer: &mut (impl AsyncWrite + Unpin), message: Message) { + let mut bytes = BytesMut::new(); + message.encode(&mut bytes).expect("encode JMUX message"); + writer.write_all(&bytes).await.expect("send JMUX message"); +} + +async fn receive_message(reader: &mut (impl AsyncRead + Unpin)) -> Message { + timeout(TEST_TIMEOUT, async { + let mut header = [0; Header::SIZE]; + reader.read_exact(&mut header).await.expect("read JMUX header"); + let message_size = usize::from(u16::from_be_bytes([header[1], header[2]])); + let mut body = vec![0; message_size - Header::SIZE]; + reader.read_exact(&mut body).await.expect("read JMUX body"); + + let mut bytes = BytesMut::with_capacity(message_size); + bytes.extend_from_slice(&header); + bytes.extend_from_slice(&body); + Message::decode(bytes.freeze()).expect("decode JMUX message") + }) + .await + .expect("JMUX response timed out") +} + +#[tokio::test] +async fn connector_success_opens_channel() { + let (proxy_stream, peer_stream) = tokio::io::duplex(8192); + let (proxy_reader, proxy_writer) = tokio::io::split(proxy_stream); + let (mut peer_reader, mut peer_writer) = tokio::io::split(peer_stream); + let proxy = JmuxProxy::new(Box::new(proxy_reader), Box::new(proxy_writer)) + .with_config(JmuxConfig::permissive()) + .with_target_connector(move |destination| async move { + assert_eq!(destination.host(), "agent.example"); + let (target_stream, mut target_peer) = tokio::io::duplex(64); + tokio::spawn(async move { + target_peer.shutdown().await.expect("close target stream"); + }); + Ok(Some(ConnectedTarget::new(target_stream))) + }); + let proxy_task = tokio::spawn(proxy.run()); + + send_message( + &mut peer_writer, + Message::open( + LocalChannelId::from(7), + 4096, + DestinationUrl::new("tcp", "agent.example", 443), + ), + ) + .await; + + let Message::OpenSuccess(open_success) = receive_message(&mut peer_reader).await else { + panic!("expected OPEN SUCCESS"); + }; + let local_id = DistantChannelId::from(open_success.sender_channel_id); + + assert!(matches!(receive_message(&mut peer_reader).await, Message::Eof(_))); + send_message(&mut peer_writer, Message::eof(local_id)).await; + assert!(matches!(receive_message(&mut peer_reader).await, Message::Close(_))); + send_message(&mut peer_writer, Message::close(local_id)).await; + + proxy_task.abort(); +} + +#[tokio::test] +async fn connector_failure_is_bounded_and_does_not_stop_direct_fallback() { + let (proxy_stream, peer_stream) = tokio::io::duplex(8192); + let (proxy_reader, proxy_writer) = tokio::io::split(proxy_stream); + let (mut peer_reader, mut peer_writer) = tokio::io::split(peer_stream); + let proxy = JmuxProxy::new(Box::new(proxy_reader), Box::new(proxy_writer)) + .with_config(JmuxConfig::permissive()) + .with_target_connector(|destination| async move { + if destination.host() == "fail.example" { + anyhow::bail!("{}", "agent error ".repeat(8192)); + } + Ok(None) + }); + let proxy_task = tokio::spawn(proxy.run()); + + send_message( + &mut peer_writer, + Message::open( + LocalChannelId::from(11), + 4096, + DestinationUrl::new("tcp", "fail.example", 443), + ), + ) + .await; + + let Message::OpenFailure(open_failure) = receive_message(&mut peer_reader).await else { + panic!("expected OPEN FAILURE"); + }; + assert_eq!(open_failure.reason_code, ReasonCode::GENERAL_FAILURE); + assert_eq!(open_failure.description, "target connection failed"); + let listener = tokio::net::TcpListener::bind("127.0.0.1:0") + .await + .expect("bind direct target"); + let target_port = listener.local_addr().expect("read direct target address").port(); + let server_task = tokio::spawn(async move { + let (_stream, _) = listener.accept().await.expect("accept direct connection"); + std::future::pending::<()>().await; + }); + + send_message( + &mut peer_writer, + Message::open( + LocalChannelId::from(12), + 4096, + DestinationUrl::new("tcp", "127.0.0.1", target_port), + ), + ) + .await; + let Message::OpenSuccess(open_success) = receive_message(&mut peer_reader).await else { + panic!("expected OPEN SUCCESS"); + }; + assert_eq!(open_success.sender_channel_id, 0); + + server_task.abort(); + proxy_task.abort(); +} diff --git a/devolutions-agent/src/tunnel.rs b/devolutions-agent/src/tunnel.rs index e19c69d65..8ca7e7980 100644 --- a/devolutions-agent/src/tunnel.rs +++ b/devolutions-agent/src/tunnel.rs @@ -756,7 +756,7 @@ async fn run_session_proxy( // Whatever went wrong has to travel back as a ConnectResponse::Error — returning // early instead drops the stream and the Gateway just sees an unexplained EOF. - let (tcp_stream, selected_target) = match connect_result { + let (mut tcp_stream, selected_target) = match connect_result { Ok(connected) => connected, Err(error) => { let reason = format!("{error:#}"); @@ -783,20 +783,10 @@ async fn run_session_proxy( .context("send ConnectResponse")?; info!("Sent ConnectResponse::Success"); - let (mut send, mut recv) = session.into_inner(); - let (mut tcp_read, mut tcp_write) = tcp_stream.into_split(); - - // Use join! (not select!) to wait for BOTH directions to finish. - // select! would cancel in-flight data when one direction closes first. - let (r1, r2) = tokio::join!( - tokio::io::copy(&mut recv, &mut tcp_write), - tokio::io::copy(&mut tcp_read, &mut send), - ); - r1.inspect_err(|e| debug!(%e, "QUIC->TCP copy ended"))?; - r2.inspect_err(|e| debug!(%e, "TCP->QUIC copy ended"))?; - - // Gracefully finish the QUIC send stream (signals EOF to peer). - let _ = send.finish(); + let (send, recv) = session.into_inner(); + tokio::io::copy_bidirectional(&mut tokio::io::join(recv, send), &mut tcp_stream) + .await + .context("proxy session traffic")?; Ok(()) } diff --git a/devolutions-gateway/src/api/jmux.rs b/devolutions-gateway/src/api/jmux.rs index b9199c01f..e40bcf3e6 100644 --- a/devolutions-gateway/src/api/jmux.rs +++ b/devolutions-gateway/src/api/jmux.rs @@ -22,6 +22,7 @@ pub async fn handler( shutdown_signal, conf_handle, traffic_audit_handle, + agent_tunnel_handle, .. }): State, JmuxToken(claims): JmuxToken, @@ -35,6 +36,7 @@ pub async fn handler( sessions, subscriber_tx, traffic_audit_handle, + agent_tunnel_handle, claims, source_addr, Duration::from_secs(conf_handle.get_conf().debug.ws_keep_alive_interval), @@ -54,6 +56,7 @@ async fn handle_socket( sessions: SessionMessageSender, subscriber_tx: SubscriberSender, traffic_audit_handle: TrafficAuditHandle, + agent_tunnel_handle: Option>, claims: JmuxTokenClaims, source_addr: SocketAddr, keep_alive_interval: Duration, @@ -65,9 +68,16 @@ async fn handle_socket( ); let session_id = claims.jet_aid; - let result = crate::jmux::handle(stream, claims, sessions, subscriber_tx, traffic_audit_handle) - .instrument(info_span!("jmux", client = %source_addr, %session_id)) - .await; + let result = crate::jmux::handle( + stream, + claims, + sessions, + subscriber_tx, + traffic_audit_handle, + agent_tunnel_handle, + ) + .instrument(info_span!("jmux", client = %source_addr, %session_id)) + .await; if let Err(error) = result { close_handle.server_error("JMUX failure".to_owned()).await; diff --git a/devolutions-gateway/src/jmux.rs b/devolutions-gateway/src/jmux.rs index 500ec0963..cbf6f3dec 100644 --- a/devolutions-gateway/src/jmux.rs +++ b/devolutions-gateway/src/jmux.rs @@ -2,7 +2,7 @@ use std::sync::Arc; use anyhow::Context as _; use devolutions_gateway_task::ChildTask; -use jmux_proxy::{FilteringRule, JmuxConfig, JmuxProxy}; +use jmux_proxy::{ConnectedTarget, DestinationUrl, FilteringRule, JmuxConfig, JmuxProxy}; use tap::prelude::*; use tokio::io::{AsyncRead, AsyncWrite}; use tokio::sync::Notify; @@ -10,8 +10,10 @@ use transport::{ErasedRead, ErasedWrite}; use crate::session::{ConnectionModeDetails, SessionInfo, SessionMessageSender}; use crate::subscriber::SubscriberSender; +use crate::target_addr::TargetAddr; use crate::token::{JmuxTokenClaims, RecordingPolicy}; use crate::traffic_audit::TrafficAuditHandle; +use crate::upstream::route_target_from_target_addr; pub async fn handle( stream: impl AsyncRead + AsyncWrite + Send + 'static, @@ -19,6 +21,7 @@ pub async fn handle( sessions: SessionMessageSender, subscriber_tx: SubscriberSender, traffic_audit_handle: TrafficAuditHandle, + agent_tunnel_handle: Option>, ) -> anyhow::Result<()> { match claims.jet_rec { RecordingPolicy::None | RecordingPolicy::Stream => (), @@ -105,10 +108,43 @@ pub async fn handle( }); }; - let proxy_fut = JmuxProxy::new(reader, writer) + let mut proxy = JmuxProxy::new(reader, writer) .with_config(config) - .with_outgoing_traffic_event_callback(traffic_event_callback) - .run(); + .with_outgoing_traffic_event_callback(traffic_event_callback); + + if let Some(agent_tunnel_handle) = agent_tunnel_handle { + proxy = proxy.with_target_connector(move |destination_url: DestinationUrl| { + let agent_tunnel_handle = Arc::clone(&agent_tunnel_handle); + + async move { + let target = TargetAddr::from_components( + destination_url.scheme(), + destination_url.host(), + destination_url.port(), + ) + .context("invalid JMUX target")?; + let route_target = route_target_from_target_addr(&target); + + let routed = agent_tunnel::routing::try_route( + Some(agent_tunnel_handle.as_ref()), + // TODO: Pass `jet_agent_id` after JMUX consumers start issuing it. + None, + &route_target, + session_id, + target.as_addr(), + ) + .await?; + + let Some((stream, _agent)) = routed else { + return Ok(None); + }; + + Ok(Some(ConnectedTarget::new(stream))) + } + }); + } + + let proxy_fut = proxy.run(); let proxy_handle = ChildTask::spawn(proxy_fut); let join_fut = proxy_handle.join(); tokio::pin!(join_fut);