use crate::FlowCtrlParameters;
use crate::ccparams::{
AlgorithmDiscriminants, CongestionWindowParams, FixedWindowParams, RoundTripEstimatorParams,
VegasParams,
};
use crate::channel::Channel;
use crate::circuit::celltypes::{CreateRequest, CreateResponse};
use crate::circuit::circhop::{HandshakeParamsError, HopSettings};
use crate::circuit::{CircuitRxSender, HandshakeSubprotocols, UniqId};
use crate::client::circuit::padding::PaddingController;
use crate::crypto::binding::CircuitBinding;
use crate::crypto::cell::CryptInit as _;
use crate::crypto::cell::{InboundRelayLayer, OutboundRelayLayer, RelayLayer, tor1};
use crate::crypto::handshake::RelayHandshakeError;
use crate::crypto::handshake::ServerHandshake as _;
use crate::crypto::handshake::fast::CreateFastServer;
use crate::crypto::handshake::ntor::{NtorSecretKey, NtorServer};
use crate::memquota::SpecificAccount as _;
use crate::memquota::{ChannelAccount, CircuitAccount};
use crate::relay::channel_provider::ChannelProvider;
use crate::relay::reactor::Reactor;
use crate::relay::{IncomingStreamRequestFilter, RelayCirc};
use crate::stream::IncomingStream;
use futures::channel::mpsc;
use futures::{SinkExt, Stream};
use smallvec::SmallVec;
use std::sync::{Arc, RwLock, Weak};
use tor_cell::chancell::ChanMsg as _;
use tor_cell::chancell::CircId;
use tor_cell::chancell::msg::{
CreateFast, Created2, CreatedFast, Destroy, DestroyReason, HandshakeType,
};
use tor_cell::relaycell::RelayCmd;
use tor_error::{ErrorKind, HasKind, debug_report, internal, into_internal, warn_report};
use tor_linkspec::OwnedChanTarget;
use tor_llcrypto::cipher::aes::Aes128Ctr;
use tor_llcrypto::d::Sha1;
use tor_llcrypto::pk::ed25519::Ed25519Identity;
use tor_llcrypto::pk::rsa::RsaIdentity;
use tor_memquota::mq_queue::ChannelSpec as _;
use tor_memquota::mq_queue::MpscSpec;
use tor_relay_crypto::pk::{RelayNtorKeypair, RelayNtorKeys};
use tor_rtcompat::SpawnExt as _;
use tor_rtcompat::{DynTimeProvider, Runtime};
use tracing::trace;
#[derive(derive_more::Debug)]
pub struct CreateRequestHandler {
chan_provider: Weak<dyn ChannelProvider<BuildSpec = OwnedChanTarget> + Send + Sync>,
circ_net_params: RwLock<CircNetParameters>,
#[debug(skip)]
ntor_keys: RwLock<RelayNtorKeys>,
#[debug(skip)]
incoming_filter_factory: Box<dyn IncomingStreamRequestFilterFactory + Send + Sync>,
allowed_stream_cmds: SmallVec<[RelayCmd; 3]>,
#[debug(skip)]
circuit_stream_tx: mpsc::Sender<Box<dyn Stream<Item = IncomingStream> + Send + Sync + Unpin>>,
}
impl CreateRequestHandler {
pub fn new(
chan_provider: Weak<dyn ChannelProvider<BuildSpec = OwnedChanTarget> + Send + Sync>,
circ_net_params: CircNetParameters,
ntor_keys: RelayNtorKeys,
incoming_filter_factory: Box<dyn IncomingStreamRequestFilterFactory + Send + Sync>,
allowed_stream_cmds: &[RelayCmd],
) -> (Self, CircuitIncomingStreamReceiver) {
const CIRC_STREAM_BUF_SIZE: usize = 1024;
#[allow(clippy::disallowed_methods)]
let (stream_tx, stream_rx) = mpsc::channel(CIRC_STREAM_BUF_SIZE);
let handler = Self {
chan_provider,
circ_net_params: RwLock::new(circ_net_params),
ntor_keys: RwLock::new(ntor_keys),
incoming_filter_factory,
allowed_stream_cmds: allowed_stream_cmds.into(),
circuit_stream_tx: stream_tx,
};
let circuit_stream_rx = CircuitIncomingStreamReceiver {
circuit_stream_rx: stream_rx,
};
(handler, circuit_stream_rx)
}
pub fn update_params(&self, circ_net_params: CircNetParameters) {
*self.circ_net_params.write().expect("rwlock poisoned") = circ_net_params;
}
pub fn update_ntor_keys(&self, ntor_keys: RelayNtorKeys) {
*self.ntor_keys.write().expect("rwlock poisoned") = ntor_keys;
}
#[allow(clippy::too_many_arguments)]
pub(crate) fn handle_create<R: Runtime>(
&self,
runtime: &R,
channel: &Arc<Channel>,
our_ed25519_id: &Ed25519Identity,
our_rsa_id: &RsaIdentity,
circ_id: CircId,
msg: &CreateRequest,
memquota: &ChannelAccount,
circ_unique_id: UniqId,
) -> Result<(CreateResponse, RelayCircComponents), Destroy> {
let result = self.handle_create_inner(
runtime,
channel,
our_ed25519_id,
our_rsa_id,
circ_id,
msg,
memquota,
circ_unique_id,
);
match result {
Ok(x) => Ok(x),
Err(e) => {
let cmd = msg.cmd();
debug_report!(&e, %cmd, "Failed to handle circuit create request");
Err(Destroy::new(DestroyReason::NONE))
}
}
}
#[allow(clippy::too_many_arguments)]
fn handle_create_inner<R: Runtime>(
&self,
runtime: &R,
channel: &Arc<Channel>,
our_ed25519_id: &Ed25519Identity,
our_rsa_id: &RsaIdentity,
circ_id: CircId,
msg: &CreateRequest,
memquota: &ChannelAccount,
circ_unique_id: UniqId,
) -> Result<(CreateResponse, RelayCircComponents), HandleCreateError> {
let handshake_components = match msg {
CreateRequest::CreateFast(msg) => self.handle_create_fast(msg)?,
CreateRequest::Create2(msg) => match msg.handshake_type() {
HandshakeType::NTOR_V3 => self.handle_create2_ntorv3(msg.body(), our_ed25519_id)?,
HandshakeType::NTOR => self.handle_create2_ntor(msg.body(), our_rsa_id)?,
x @ HandshakeType::TAP | x => {
return Err(HandleCreateError::Create2HandshakeType(x));
}
},
};
let memquota = CircuitAccount::new(memquota)?;
let time_provider = DynTimeProvider::new(runtime.clone());
let account = memquota.as_raw_account();
let (sender, receiver) =
MpscSpec::new(10_000_000).new_mq(time_provider.clone(), account)?;
let (sender, receiver) = crate::circuit::circ_sender::channel(sender, receiver);
let (padding_ctrl, padding_stream) =
crate::client::circuit::padding::new_padding(DynTimeProvider::new(runtime.clone()));
let Some(chan_provider) = self.chan_provider.upgrade() else {
return Err(internal!("Unable to upgrade weak `ChannelProvider`").into());
};
let incoming_filter = self.incoming_filter_factory.current_filter();
let (reactor, circ, incoming_streams) = Reactor::new(
runtime.clone(),
channel,
circ_id,
circ_unique_id,
receiver,
handshake_components.crypto_in,
handshake_components.crypto_out,
&handshake_components.hop_settings,
chan_provider,
padding_ctrl.clone(),
padding_stream,
incoming_filter,
&self.allowed_stream_cmds,
&memquota,
)
.map_err(into_internal!("Failed to start circuit reactor"))?;
let mut circuit_stream_tx = self.circuit_stream_tx.clone();
let () = runtime.spawn(async move {
if let Err(e) = circuit_stream_tx.send(Box::new(incoming_streams)).await {
warn_report!(e, "IncomingStream handler disappeared?!");
drop(reactor);
} else {
match reactor.run().await {
Ok(()) => {}
Err(e) => {
debug_report!(e, "Relay circuit reactor exited with an error");
}
}
}
})?;
Ok((
handshake_components.response,
RelayCircComponents {
circ,
sender,
padding_ctrl,
},
))
}
fn handle_create_fast(
&self,
msg: &CreateFast,
) -> Result<CompletedHandshakeComponents, HandleCreateError> {
let (keygen, handshake_msg) = CreateFastServer::server(
&mut rand::rng(),
&mut |_: &()| Some(()),
&[()],
msg.handshake(),
)?;
let circ_net_params = self
.circ_net_params
.read()
.expect("rwlock poisoned")
.clone();
let subprotos = HandshakeSubprotocols::default();
let hop_settings = HopSettings::from_handshake_params(
circ_net_params,
AlgorithmDiscriminants::FixedWindow,
subprotos,
)?;
let crypt = tor1::CryptStatePair::<Aes128Ctr, Sha1>::construct(keygen)
.map_err(into_internal!("Circuit crypt state construction failed"))?;
let (crypto_out, crypto_in, _binding) = split_relay_layer(crypt);
let response = CreatedFast::new(handshake_msg);
let response = CreateResponse::CreatedFast(response);
trace!("Completed CREATE_FAST handshake");
Ok(CompletedHandshakeComponents {
response,
hop_settings,
crypto_out,
crypto_in,
})
}
fn handle_create2_ntor(
&self,
msg_body: &[u8],
our_rsa_id: &RsaIdentity,
) -> Result<CompletedHandshakeComponents, HandleCreateError> {
let ntor_keys = self.ntor_keys(|k| {
NtorSecretKey::new(k.secret().clone(), *k.public().inner(), *our_rsa_id)
});
let (keygen, handshake_msg) = NtorServer::server(
&mut rand::rng(),
&mut |_: &()| Some(()),
ntor_keys.as_ref(),
msg_body,
)?;
let circ_net_params = self
.circ_net_params
.read()
.expect("rwlock poisoned")
.clone();
let subprotos = HandshakeSubprotocols::default();
let hop_settings = HopSettings::from_handshake_params(
circ_net_params,
AlgorithmDiscriminants::FixedWindow,
subprotos,
)?;
let crypt = tor1::CryptStatePair::<Aes128Ctr, Sha1>::construct(keygen)
.map_err(into_internal!("Circuit crypt state construction failed"))?;
let (crypto_out, crypto_in, _binding) = split_relay_layer(crypt);
let response = Created2::new(handshake_msg);
let response = CreateResponse::Created2(response);
trace!("Completed ntor handshake");
Ok(CompletedHandshakeComponents {
response,
hop_settings,
crypto_out,
crypto_in,
})
}
fn handle_create2_ntorv3(
&self,
_msg_body: &[u8],
_our_ed25519_id: &Ed25519Identity,
) -> Result<CompletedHandshakeComponents, HandleCreateError> {
Err(HandleCreateError::Create2HandshakeType(
HandshakeType::NTOR_V3,
))
}
fn ntor_keys<T>(&self, map: impl FnMut(&RelayNtorKeypair) -> T) -> impl AsRef<[T]> {
let ntor_keys = self.ntor_keys.read().expect("rwlock poisoned");
let ntor_keys = [Some(ntor_keys.latest()), ntor_keys.previous()];
ntor_keys
.into_iter()
.flatten()
.map(map)
.collect::<SmallVec<[T; 2]>>()
}
}
pub struct CircuitIncomingStreamReceiver {
circuit_stream_rx: mpsc::Receiver<<Self as Stream>::Item>,
}
impl Stream for CircuitIncomingStreamReceiver {
type Item = Box<dyn Stream<Item = IncomingStream> + Send + Sync + Unpin>;
fn poll_next(
mut self: std::pin::Pin<&mut Self>,
cx: &mut std::task::Context<'_>,
) -> std::task::Poll<Option<Self::Item>> {
use futures::StreamExt as _;
self.circuit_stream_rx.poll_next_unpin(cx)
}
}
fn split_relay_layer<F, B>(
crypt: impl RelayLayer<F, B>,
) -> (
Box<dyn OutboundRelayLayer + Send>,
Box<dyn InboundRelayLayer + Send>,
CircuitBinding,
)
where
F: OutboundRelayLayer + Send + 'static,
B: InboundRelayLayer + Send + 'static,
{
let (crypto_out, crypto_in, binding) = crypt.split_relay_layer();
let (crypto_out, crypto_in) = (Box::new(crypto_out), Box::new(crypto_in));
(crypto_out, crypto_in, binding)
}
#[derive(Debug, thiserror::Error)]
enum HandleCreateError {
#[error("Circuit relay handshake failed")]
Handshake(#[from] RelayHandshakeError),
#[error("Failed to process the circuit relay handshake parameters")]
HandshakeParameters(#[from] HandshakeParamsError),
#[error("Unsupported handshake type {0}")]
Create2HandshakeType(HandshakeType),
#[error("Memquota error")]
Memquota(#[from] tor_memquota::Error),
#[error("Runtime task spawn error")]
Spawn(#[from] futures::task::SpawnError),
#[error("Internal error")]
Internal(#[from] tor_error::Bug),
}
impl HasKind for HandleCreateError {
fn kind(&self) -> ErrorKind {
match self {
Self::Handshake(e) => e.kind(),
Self::HandshakeParameters(e) => e.kind(),
Self::Create2HandshakeType(_) => ErrorKind::NotImplemented,
Self::Memquota(e) => e.kind(),
Self::Spawn(e) => e.kind(),
Self::Internal(_) => ErrorKind::Internal,
}
}
}
struct CompletedHandshakeComponents {
response: CreateResponse,
hop_settings: HopSettings,
crypto_out: Box<dyn OutboundRelayLayer + Send>,
crypto_in: Box<dyn InboundRelayLayer + Send>,
}
pub(crate) struct RelayCircComponents {
pub(crate) circ: Arc<RelayCirc>,
pub(crate) sender: CircuitRxSender,
pub(crate) padding_ctrl: PaddingController,
}
#[derive(Debug, Clone)]
#[allow(clippy::exhaustive_structs)]
pub struct CongestionControlNetParams {
pub fixed_window: FixedWindowParams,
pub vegas_exit: VegasParams,
pub cwnd: CongestionWindowParams,
pub rtt: RoundTripEstimatorParams,
pub flow_ctrl: FlowCtrlParameters,
}
impl CongestionControlNetParams {
#[cfg(test)]
pub(crate) fn defaults_for_tests() -> Self {
Self {
fixed_window: FixedWindowParams::defaults_for_tests(),
vegas_exit: VegasParams::defaults_for_tests(),
cwnd: CongestionWindowParams::defaults_for_tests(),
rtt: RoundTripEstimatorParams::defaults_for_tests(),
flow_ctrl: FlowCtrlParameters::defaults_for_tests(),
}
}
}
#[derive(Debug, Clone)]
#[allow(clippy::exhaustive_structs)]
pub struct CircNetParameters {
pub cc: CongestionControlNetParams,
}
pub trait IncomingStreamRequestFilterFactory {
fn current_filter(&self) -> Box<dyn IncomingStreamRequestFilter>;
}
impl<F> IncomingStreamRequestFilterFactory for F
where
F: Fn() -> Box<dyn IncomingStreamRequestFilter>,
{
fn current_filter(&self) -> Box<dyn IncomingStreamRequestFilter> {
(self)()
}
}
#[cfg(test)]
mod test {
#![allow(clippy::bool_assert_comparison)]
#![allow(clippy::clone_on_copy)]
#![allow(clippy::dbg_macro)]
#![allow(clippy::mixed_attributes_style)]
#![allow(clippy::print_stderr)]
#![allow(clippy::print_stdout)]
#![allow(clippy::single_char_pattern)]
#![allow(clippy::unwrap_used)]
#![allow(clippy::unchecked_time_subtraction)]
#![allow(clippy::useless_vec)]
#![allow(clippy::needless_pass_by_value)]
#![allow(clippy::string_slice)]
use tor_cell::chancell::{ChanCmd, ChanMsg as _};
use tor_rtcompat::test_with_one_runtime;
use crate::channel::test_utils;
use crate::circuit::CircParameters;
#[test]
fn create_fast() {
test_with_one_runtime!(|rt| async move {
let mut conn_inspector = test_utils::ConnInspector::new();
let (client_chan, _relay_chan, _circuit_stream_rx, _target_builder) =
test_utils::new_channel_pair_with_keys(&rt, &conn_inspector);
let pending_tunnel = test_utils::new_pending_tunnel(&rt, &client_chan).await;
let circ_params = CircParameters::default();
let tunnel = pending_tunnel
.create_firsthop_fast(circ_params)
.await
.unwrap();
assert_eq!(
conn_inspector.try_client_cell().unwrap().msg().cmd(),
ChanCmd::CREATE_FAST,
);
assert_eq!(
conn_inspector.try_relay_cell().unwrap().msg().cmd(),
ChanCmd::CREATED_FAST,
);
drop(tunnel);
assert_eq!(
conn_inspector.client_cell().await.unwrap().msg().cmd(),
ChanCmd::DESTROY,
);
assert_eq!(
conn_inspector.relay_cell().await.unwrap().msg().cmd(),
ChanCmd::DESTROY,
);
});
}
#[test]
fn tap() {
test_with_one_runtime!(|rt| async move {
let mut conn_inspector = test_utils::ConnInspector::new();
let (client_chan, _relay_chan, _circuit_stream_rx, mut target_builder) =
test_utils::new_channel_pair_with_keys(&rt, &conn_inspector);
let pending_tunnel = test_utils::new_pending_tunnel(&rt, &client_chan).await;
let circ_params = CircParameters::default();
let protocols = "Relay=1".parse().unwrap();
let target = target_builder.protocols(protocols).build().unwrap();
let _tunnel = pending_tunnel
.create_firsthop(&target, circ_params)
.await
.unwrap();
assert_eq!(
conn_inspector.try_client_cell().unwrap().msg().cmd(),
ChanCmd::CREATE2,
);
assert_eq!(
conn_inspector.try_relay_cell().unwrap().msg().cmd(),
ChanCmd::CREATED2,
);
});
}
#[test]
fn ntor() {
test_with_one_runtime!(|rt| async move {
let mut conn_inspector = test_utils::ConnInspector::new();
let (client_chan, _relay_chan, _circuit_stream_rx, mut target_builder) =
test_utils::new_channel_pair_with_keys(&rt, &conn_inspector);
for relay_version in [2, 3] {
let pending_tunnel = test_utils::new_pending_tunnel(&rt, &client_chan).await;
let circ_params = CircParameters::default();
let protocols = format!("Relay=2-{relay_version}").parse().unwrap();
let target = target_builder.protocols(protocols).build().unwrap();
let tunnel = pending_tunnel
.create_firsthop(&target, circ_params)
.await
.unwrap();
assert_eq!(
conn_inspector.try_client_cell().unwrap().msg().cmd(),
ChanCmd::CREATE2,
);
assert_eq!(
conn_inspector.try_relay_cell().unwrap().msg().cmd(),
ChanCmd::CREATED2,
);
drop(tunnel);
assert_eq!(
conn_inspector.client_cell().await.unwrap().msg().cmd(),
ChanCmd::DESTROY,
);
assert_eq!(
conn_inspector.relay_cell().await.unwrap().msg().cmd(),
ChanCmd::DESTROY,
);
}
});
}
}