1use bitvec::prelude::*;
4use derive_deftly::Deftly;
5use oneshot_fused_workaround as oneshot;
6
7use tor_cell::relaycell::{RelayCellFormat, RelayCmd, StreamId, UnparsedRelayMsg, msg};
8use tor_cell::restricted_msg;
9use tor_error::internal;
10use tor_memquota::derive_deftly_template_HasMemoryCost;
11use tor_memquota::mq_queue::{self, MpscSpec};
12use tor_rtcompat::DynTimeProvider;
13
14use crate::circuit::CircHopSyncView;
15use crate::circuit::circhop::ReactorStreamComponents;
16use crate::stream::cmdcheck::{AnyCmdChecker, CmdChecker, StreamStatus};
17use crate::stream::{CloseStreamBehavior, StreamComponents};
18use crate::{Error, Result};
19
20use crate::client::stream::DataStream;
22
23use crate::memquota::StreamAccount;
24use crate::{HopLocation, HopNum};
25
26#[derive(Debug, Default)]
28pub(crate) struct InboundDataCmdChecker;
29
30restricted_msg! {
31 enum IncomingDataStreamMsg:RelayMsg {
33 Data, End,
35 }
36}
37
38impl CmdChecker for InboundDataCmdChecker {
39 fn check_msg(&mut self, msg: &tor_cell::relaycell::UnparsedRelayMsg) -> Result<StreamStatus> {
40 use StreamStatus::*;
41 match msg.cmd() {
42 RelayCmd::DATA => Ok(Open),
43 RelayCmd::END => Ok(Closed),
44 _ => Err(Error::StreamProto(format!(
45 "Unexpected {} on an incoming data stream!",
46 msg.cmd()
47 ))),
48 }
49 }
50
51 fn consume_checked_msg(&mut self, msg: tor_cell::relaycell::UnparsedRelayMsg) -> Result<()> {
52 let _ = msg
53 .decode::<IncomingDataStreamMsg>()
54 .map_err(|err| Error::from_bytes_err(err, "cell on half-closed stream"))?;
55 Ok(())
56 }
57}
58
59impl InboundDataCmdChecker {
60 pub(crate) fn new_connected() -> AnyCmdChecker {
66 Box::new(Self)
67 }
68}
69
70#[derive(Debug)]
80pub struct IncomingStream {
81 time_provider: DynTimeProvider,
83 request: IncomingStreamRequest,
85 components: StreamComponents,
87}
88
89impl IncomingStream {
90 pub(crate) fn new(
92 time_provider: DynTimeProvider,
93 request: IncomingStreamRequest,
94 components: StreamComponents,
95 ) -> Self {
96 Self {
97 time_provider,
98 request,
99 components,
100 }
101 }
102
103 pub fn request(&self) -> &IncomingStreamRequest {
105 &self.request
106 }
107
108 pub async fn accept_data(self, message: msg::Connected) -> Result<DataStream> {
113 let Self {
114 time_provider,
115 request,
116 components:
117 StreamComponents {
118 mut target,
119 stream_receiver,
120 xon_xoff_reader_ctrl,
121 memquota,
122 },
123 } = self;
124
125 match request {
126 IncomingStreamRequest::Begin(_) | IncomingStreamRequest::BeginDir(_) => {
127 target.send(message.into()).await?;
128 Ok(DataStream::new_connected(
129 time_provider,
130 stream_receiver,
131 xon_xoff_reader_ctrl,
132 target,
133 memquota,
134 ))
135 }
136 IncomingStreamRequest::Resolve(_) => {
137 Err(internal!("Cannot accept data on a RESOLVE stream").into())
138 }
139 }
140 }
141
142 #[cfg(feature = "relay")]
148 pub async fn resolve(mut self, message: msg::Resolved) -> Result<()> {
149 match self.request {
150 IncomingStreamRequest::Begin(_) | IncomingStreamRequest::BeginDir(_) => {
151 Err(internal!("Cannot send RESOLVED on a data or directory stream").into())
152 }
153 IncomingStreamRequest::Resolve(_) => {
154 let rx = self.close(CloseStreamBehavior::SendResolved(message))?;
155
156 rx.await.map_err(|_| Error::CircuitClosed)?
157 }
158 }
159 }
160
161 pub async fn reject(mut self, message: msg::End) -> Result<()> {
163 let rx = self.close(CloseStreamBehavior::SendEnd(message))?;
164
165 rx.await.map_err(|_| Error::CircuitClosed)?
166 }
167
168 fn close(&mut self, message: CloseStreamBehavior) -> Result<oneshot::Receiver<Result<()>>> {
172 self.components.target.close_pending(message)
173 }
174
175 pub async fn discard(mut self) -> Result<()> {
181 let rx = self.close(CloseStreamBehavior::SendNothing)?;
182
183 rx.await.map_err(|_| Error::CircuitClosed)?.map(|_| ())
184 }
185}
186
187restricted_msg! {
193 #[derive(Clone, Debug, Deftly)]
195 #[derive_deftly(HasMemoryCost)]
196 #[non_exhaustive]
197 pub enum IncomingStreamRequest: RelayMsg {
198 Begin,
200 BeginDir,
202 Resolve,
204 }
205}
206
207type RelayCmdSet = bitvec::BitArr!(for 256);
212
213#[derive(Debug)]
216pub(crate) struct IncomingCmdChecker {
217 allow_commands: RelayCmdSet,
225}
226
227impl IncomingCmdChecker {
228 pub(crate) fn new_any(allow_commands: &[RelayCmd]) -> AnyCmdChecker {
230 let mut array = BitArray::ZERO;
231 for c in allow_commands {
232 array.set(u8::from(*c) as usize, true);
233 }
234 Box::new(Self {
235 allow_commands: array,
236 })
237 }
238}
239
240impl CmdChecker for IncomingCmdChecker {
241 fn check_msg(&mut self, msg: &UnparsedRelayMsg) -> Result<StreamStatus> {
242 if self.allow_commands[u8::from(msg.cmd()) as usize] {
243 Ok(StreamStatus::Open)
244 } else {
245 Err(Error::StreamProto(format!(
246 "Unexpected {} on incoming stream",
247 msg.cmd()
248 )))
249 }
250 }
251
252 fn consume_checked_msg(&mut self, msg: UnparsedRelayMsg) -> Result<()> {
253 let _ = msg
254 .decode::<IncomingStreamRequest>()
255 .map_err(|err| Error::from_bytes_err(err, "invalid message on incoming stream"))?;
256
257 Ok(())
258 }
259}
260
261pub trait IncomingStreamRequestFilter: Send + 'static {
269 fn disposition(
271 &mut self,
272 ctx: &IncomingStreamRequestContext<'_>,
273 circ: &CircHopSyncView<'_>,
274 ) -> Result<IncomingStreamRequestDisposition>;
275}
276
277#[derive(Clone, Debug)]
279#[non_exhaustive]
280pub enum IncomingStreamRequestDisposition {
281 Accept,
284 CloseCircuit,
286 RejectRequest(msg::End),
288}
289
290pub struct IncomingStreamRequestContext<'a> {
292 pub(crate) request: &'a IncomingStreamRequest,
294}
295impl<'a> IncomingStreamRequestContext<'a> {
296 pub fn request(&self) -> &'a IncomingStreamRequest {
298 self.request
299 }
300}
301
302#[cfg(test)]
304#[derive(Copy, Clone, Debug, Default)]
305pub(crate) struct NoOpRequestFilter;
306
307#[cfg(test)]
308impl IncomingStreamRequestFilter for NoOpRequestFilter {
309 fn disposition(
310 &mut self,
311 _ctx: &IncomingStreamRequestContext<'_>,
312 _circ: &CircHopSyncView<'_>,
313 ) -> crate::Result<IncomingStreamRequestDisposition> {
314 Ok(IncomingStreamRequestDisposition::Accept)
315 }
316}
317
318#[derive(Debug, Deftly)]
320#[derive_deftly(HasMemoryCost)]
321pub(crate) struct StreamReqInfo {
322 pub(crate) req: IncomingStreamRequest,
324 pub(crate) stream_id: StreamId,
326 pub(crate) hop: Option<HopLocation>,
333 #[deftly(has_memory_cost(indirect_size = "0"))]
335 pub(crate) relay_cell_format: RelayCellFormat,
336 pub(crate) stream_components: ReactorStreamComponents,
338 #[deftly(has_memory_cost(indirect_size = "0"))] pub(crate) memquota: StreamAccount,
341}
342
343#[cfg(any(feature = "hs-service", feature = "relay"))]
345pub(crate) type StreamReqSender = mq_queue::Sender<StreamReqInfo, MpscSpec>;
346
347#[derive(educe::Educe)]
349#[educe(Debug)]
350#[cfg(any(feature = "hs-service", feature = "relay"))]
351pub(crate) struct IncomingStreamRequestHandler {
352 pub(crate) incoming_sender: StreamReqSender,
354 pub(crate) hop_num: Option<HopNum>,
358 pub(crate) cmd_checker: AnyCmdChecker,
360 #[educe(Debug(ignore))]
363 pub(crate) filter: Box<dyn IncomingStreamRequestFilter>,
364}
365
366#[cfg(test)]
367mod test {
368 #![allow(clippy::bool_assert_comparison)]
370 #![allow(clippy::clone_on_copy)]
371 #![allow(clippy::dbg_macro)]
372 #![allow(clippy::mixed_attributes_style)]
373 #![allow(clippy::print_stderr)]
374 #![allow(clippy::print_stdout)]
375 #![allow(clippy::single_char_pattern)]
376 #![allow(clippy::unwrap_used)]
377 #![allow(clippy::unchecked_time_subtraction)]
378 #![allow(clippy::useless_vec)]
379 #![allow(clippy::needless_pass_by_value)]
380 #![allow(clippy::string_slice)] use tor_cell::relaycell::{
384 AnyRelayMsgOuter, RelayCellFormat,
385 msg::{Begin, BeginDir, Data, Resolve},
386 };
387
388 use super::*;
389
390 #[test]
391 fn incoming_cmd_checker() {
392 let u = |msg| {
394 let body = AnyRelayMsgOuter::new(None, msg)
395 .encode(RelayCellFormat::V0, &mut rand::rng())
396 .unwrap();
397 UnparsedRelayMsg::from_singleton_body(RelayCellFormat::V0, body).unwrap()
398 };
399 let begin = u(Begin::new("allium.example.com", 443, 0).unwrap().into());
400 let begin_dir = u(BeginDir::default().into());
401 let resolve = u(Resolve::new("allium.example.com").into());
402 let data = u(Data::new(&[1, 2, 3]).unwrap().into());
403
404 {
405 let mut cc_none = IncomingCmdChecker::new_any(&[]);
406 for m in [&begin, &begin_dir, &resolve, &data] {
407 assert!(cc_none.check_msg(m).is_err());
408 }
409 }
410
411 {
412 let mut cc_begin = IncomingCmdChecker::new_any(&[RelayCmd::BEGIN]);
413 assert_eq!(cc_begin.check_msg(&begin).unwrap(), StreamStatus::Open);
414 for m in [&begin_dir, &resolve, &data] {
415 assert!(cc_begin.check_msg(m).is_err());
416 }
417 }
418
419 {
420 let mut cc_any = IncomingCmdChecker::new_any(&[
421 RelayCmd::BEGIN,
422 RelayCmd::BEGIN_DIR,
423 RelayCmd::RESOLVE,
424 ]);
425 for m in [&begin, &begin_dir, &resolve] {
426 assert_eq!(cc_any.check_msg(m).unwrap(), StreamStatus::Open);
427 }
428 assert!(cc_any.check_msg(&data).is_err());
429 }
430 }
431}