use crate::client_pool::ClientPool;
use crate::mixnet::{IncludedSurbs, MixnetClientBuilder, MixnetMessageSender, NymNetworkDetails};
use crate::tcp_proxy::utils::{MessageBuffer, Payload, ProxiedMessage};
use anyhow::Result;
use dashmap::DashSet;
use nym_network_defaults::setup_env;
use nym_sphinx::addressing::Recipient;
use std::sync::Arc;
use tokio::{
net::{TcpListener, TcpStream},
sync::oneshot,
};
use tokio_stream::StreamExt;
use tokio_util::codec::{BytesCodec, FramedRead};
use tokio_util::sync::CancellationToken;
use tracing::{debug, info, instrument};
const DEFAULT_CLOSE_TIMEOUT: u64 = 60;
const DEFAULT_LISTEN_HOST: &str = "127.0.0.1";
const DEFAULT_LISTEN_PORT: &str = "8080";
const DEFAULT_CLIENT_POOL_SIZE: usize = 2;
#[derive(Clone)]
pub struct NymProxyClient {
server_address: Recipient,
listen_address: String,
listen_port: String,
close_timeout: u64,
conn_pool: ClientPool,
cancel_token: CancellationToken,
}
impl NymProxyClient {
pub async fn new(
server_address: Recipient,
listen_address: &str,
listen_port: &str,
close_timeout: u64,
env: Option<String>,
default_client_amount: usize,
) -> Result<Self> {
debug!("Loading env file: {:?}", env);
setup_env(env); Ok(NymProxyClient {
server_address,
listen_address: listen_address.to_string(),
listen_port: listen_port.to_string(),
close_timeout,
conn_pool: ClientPool::new(default_client_amount),
cancel_token: CancellationToken::new(),
})
}
pub async fn new_with_defaults(server_address: Recipient, env: Option<String>) -> Result<Self> {
NymProxyClient::new(
server_address,
DEFAULT_LISTEN_HOST,
DEFAULT_LISTEN_PORT,
DEFAULT_CLOSE_TIMEOUT,
env,
DEFAULT_CLIENT_POOL_SIZE,
)
.await
}
pub async fn run(&self) -> Result<()> {
info!("Connecting to mixnet server at {}", self.server_address);
let listener =
TcpListener::bind(format!("{}:{}", self.listen_address, self.listen_port)).await?;
let client_maker = self.conn_pool.clone();
tokio::spawn(async move {
client_maker.start().await?;
Ok::<(), anyhow::Error>(())
});
loop {
tokio::select! {
stream = listener.accept() => {
let (stream, _) = stream?;
tokio::spawn(NymProxyClient::handle_incoming(
stream,
self.server_address,
self.close_timeout,
self.conn_pool.clone(),
self.cancel_token.clone(),
));
}
_ = self.cancel_token.cancelled() => {
break Ok(());
}
}
}
}
pub async fn disconnect(&self) {
self.cancel_token.cancel();
self.conn_pool.disconnect_pool().await;
}
#[instrument(skip(stream, server_address, close_timeout, conn_pool, cancel_token))]
async fn handle_incoming(
stream: TcpStream,
server_address: Recipient,
close_timeout: u64,
conn_pool: ClientPool,
cancel_token: CancellationToken,
) -> Result<()> {
let session_id = uuid::Uuid::new_v4();
let (tx, mut rx) = oneshot::channel();
info!("Starting session: {}", session_id);
let mut client = match conn_pool.get_mixnet_client().await {
Some(client) => {
info!("Grabbed client {} from pool", client.nym_address());
client
}
None => {
info!("Not enough clients in pool, creating ephemeral client");
let net = NymNetworkDetails::new_from_env();
let client = MixnetClientBuilder::new_ephemeral()
.network_details(net)
.build()?
.connect_to_mixnet()
.await?;
info!(
"Using {} for the moment, created outside of the connection pool",
client.nym_address()
);
client
}
};
let (read, mut write) = stream.into_split();
let codec = BytesCodec::new();
let mut framed_read = FramedRead::new(read, codec);
let sender = client.split_sender();
let server_addr = server_address;
let messages_account = Arc::new(DashSet::new());
let sent_messages_account = Arc::clone(&messages_account);
tokio::spawn(async move {
let mut message_id = 0;
while let Some(Ok(bytes)) = framed_read.next().await {
message_id += 1;
sent_messages_account.insert(message_id);
let message =
ProxiedMessage::new(Payload::Data(bytes.to_vec()), session_id, message_id);
let coded_message = bincode::serialize(&message)?;
sender
.send_message(server_addr, &coded_message, IncludedSurbs::Amount(100))
.await?;
info!(
"Sent message with id {} for session {} of {} bytes",
message_id,
session_id,
bytes.len()
);
}
message_id += 1;
let message = ProxiedMessage::new(Payload::Close, session_id, message_id);
let coded_message = bincode::serialize(&message)?;
sender
.send_message(server_addr, &coded_message, IncludedSurbs::Amount(100))
.await?;
info!("Closing read end of session: {}", session_id);
tx.send(true)
.map_err(|_| anyhow::anyhow!("Could not send close signal"))?;
Ok::<(), anyhow::Error>(())
});
tokio::spawn(async move {
let mut msg_buffer = MessageBuffer::new();
loop {
tokio::select! {
_ = &mut rx => {
info!("Closing write end of session: {session_id} in {close_timeout} seconds");
break
}
Some(message) = client.next() => {
let message = bincode::deserialize::<ProxiedMessage>(&message.message)?;
msg_buffer.push(message);
msg_buffer.tick(&mut write).await?;
},
_ = cancel_token.cancelled() => {
info!("CTRL_C triggered in thread, triggering loop shutdown");
break
},
_ = tokio::time::sleep(tokio::time::Duration::from_millis(100)) => {
msg_buffer.tick(&mut write).await?;
}
}
}
loop {
tokio::select! {
Some(message) = client.next() => {
let message = bincode::deserialize::<ProxiedMessage>(&message.message)?;
msg_buffer.push(message);
msg_buffer.tick(&mut write).await?;
},
_ = cancel_token.cancelled() => {
info!("CTRL_C triggered in thread, triggering client shutdown");
client.disconnect().await;
return Ok::<(), anyhow::Error>(())
},
_ = tokio::time::sleep(tokio::time::Duration::from_secs(close_timeout)) => {
info!("Closing write end of session: {}", session_id);
info!("Triggering client shutdown");
client.disconnect().await;
return Ok::<(), anyhow::Error>(())
},
}
}
});
Ok(())
}
}