1use crate::circuit::CircuitRxSender;
7use crate::client::circuit::padding::{PaddingController, QueuedCellPaddingInfo};
8use crate::{Error, Result};
9use tor_basic_utils::RngExt;
10use tor_cell::chancell::CircId;
11use tor_cell::chancell::msg::DestroyReason;
12
13use crate::circuit::celltypes::CreateResponse;
14use crate::client::circuit::halfcirc::HalfCirc;
15
16use oneshot_fused_workaround as oneshot;
17
18use rand::Rng;
19use rand::distr::Distribution;
20use std::collections::{HashMap, hash_map::Entry};
21use std::ops::{Deref, DerefMut};
22use std::result::Result as StdResult;
23use std::sync::Arc;
24
25#[cfg(feature = "relay")]
26use crate::relay::RelayCirc;
27
28#[derive(Copy, Clone)]
33pub(crate) enum CircIdRange {
34 #[allow(dead_code)] Low,
37 High,
39 }
43
44impl CircIdRange {
45 const fn integer_range(&self) -> std::ops::RangeInclusive<u32> {
48 const MIDPOINT: u32 = 0x8000_0000;
49
50 match self {
51 Self::Low => 1..=(MIDPOINT - 1),
53 Self::High => MIDPOINT..=u32::MAX,
54 }
55 }
56
57 pub(crate) fn is_allowed_for_peer(&self, id: CircId) -> bool {
59 !self.integer_range().contains(&id.into())
63 }
64}
65
66impl rand::distr::Distribution<CircId> for CircIdRange {
67 fn sample<R: Rng + ?Sized>(&self, mut rng: &mut R) -> CircId {
69 let v = rng.gen_range_checked(self.integer_range());
70 let v = v.expect("Unexpected empty range passed to gen_range_checked");
71 CircId::new(v).expect("Unexpected zero value")
72 }
73}
74
75#[derive(Debug)]
79pub(super) enum CircEnt {
80 Opening {
89 create_response_sender: oneshot::Sender<CreateResponse>,
91 cell_sender: CircuitRxSender,
94 padding_ctrl: PaddingController,
96 },
97
98 OpenOrigin {
101 cell_sender: CircuitRxSender,
104 padding_ctrl: PaddingController,
106 },
107
108 #[cfg(feature = "relay")]
111 OpenRelay {
112 _circ: Arc<RelayCirc>,
118 cell_sender: CircuitRxSender,
121 padding_ctrl: PaddingController,
123 },
124
125 DestroySent(HalfCirc),
128}
129
130pub(super) struct MutCircEnt<'a> {
136 value: &'a mut CircEnt,
138 open_count: &'a mut usize,
141 was_open: bool,
143}
144
145impl<'a> Drop for MutCircEnt<'a> {
146 fn drop(&mut self) {
147 let is_open = !matches!(self.value, CircEnt::DestroySent(_));
148 match (self.was_open, is_open) {
149 (false, true) => *self.open_count = self.open_count.saturating_add(1),
150 (true, false) => *self.open_count = self.open_count.saturating_sub(1),
151 (_, _) => (),
152 };
153 }
154}
155
156impl<'a> Deref for MutCircEnt<'a> {
157 type Target = CircEnt;
158 fn deref(&self) -> &Self::Target {
159 self.value
160 }
161}
162
163impl<'a> DerefMut for MutCircEnt<'a> {
164 fn deref_mut(&mut self) -> &mut Self::Target {
165 self.value
166 }
167}
168
169pub(super) struct CircMap {
171 m: HashMap<CircId, CircEnt>,
173 range: CircIdRange,
175 open_count: usize,
177}
178
179impl CircMap {
180 pub(super) fn new(idrange: CircIdRange) -> Self {
182 CircMap {
183 m: HashMap::new(),
184 range: idrange,
185 open_count: 0,
186 }
187 }
188
189 pub(super) fn add_origin_ent<R: Rng>(
195 &mut self,
196 rng: &mut R,
197 createdsink: oneshot::Sender<CreateResponse>,
198 sink: CircuitRxSender,
199 padding_ctrl: PaddingController,
200 ) -> Result<CircId> {
201 const N_ATTEMPTS: usize = 16;
206 let iter = self.range.sample_iter(rng).take(N_ATTEMPTS);
207 let circ_ent = CircEnt::Opening {
208 create_response_sender: createdsink,
209 cell_sender: sink,
210 padding_ctrl,
211 };
212 for id in iter {
213 let ent = self.m.entry(id);
214 if let Entry::Vacant(_) = &ent {
215 ent.or_insert(circ_ent);
216 self.open_count += 1;
217 return Ok(id);
218 }
219 }
220 Err(Error::IdRangeFull)
221 }
222
223 #[cfg(feature = "relay")]
228 pub(super) fn add_relay_ent(
229 &mut self,
230 circ_id: CircId,
231 circ: Arc<RelayCirc>,
232 sink: CircuitRxSender,
233 padding_ctrl: PaddingController,
234 ) -> StdResult<(), DestroyReason> {
235 if !self.range.is_allowed_for_peer(circ_id) {
237 return Err(DestroyReason::NONE);
238 }
239
240 let circ_ent = CircEnt::OpenRelay {
241 _circ: circ,
242 cell_sender: sink,
243 padding_ctrl,
244 };
245
246 if let Entry::Vacant(ent) = self.m.entry(circ_id) {
247 ent.insert(circ_ent);
248 self.open_count += 1;
249 Ok(())
250 } else {
251 Err(DestroyReason::NONE)
252 }
253 }
254
255 #[cfg(test)]
258 pub(super) fn put_unchecked(&mut self, id: CircId, ent: CircEnt) {
259 self.m.insert(id, ent);
260 }
261
262 pub(super) fn get_mut(&mut self, id: CircId) -> Option<MutCircEnt> {
264 let open_count = &mut self.open_count;
265 self.m.get_mut(&id).map(move |ent| MutCircEnt {
266 open_count,
267 was_open: !matches!(ent, CircEnt::DestroySent(_)),
268 value: ent,
269 })
270 }
271
272 pub(super) fn is_open(&self, id: CircId) -> bool {
278 let Some(entry) = self.m.get(&id) else {
279 return false;
280 };
281
282 match entry {
283 CircEnt::Opening { .. } | CircEnt::OpenOrigin { .. } => true,
284 #[cfg(feature = "relay")]
285 CircEnt::OpenRelay { .. } => true,
286 CircEnt::DestroySent(..) => false,
287 }
288 }
289
290 pub(super) fn note_cell_flushed(&mut self, id: CircId, info: QueuedCellPaddingInfo) {
292 let padding_ctrl = match self.m.get(&id) {
293 Some(CircEnt::Opening { padding_ctrl, .. }) => padding_ctrl,
294 Some(CircEnt::OpenOrigin { padding_ctrl, .. }) => padding_ctrl,
295 #[cfg(feature = "relay")]
296 Some(CircEnt::OpenRelay { padding_ctrl, .. }) => padding_ctrl,
297 Some(CircEnt::DestroySent(..)) | None => return,
298 };
299 padding_ctrl.flushed_relay_cell(info);
300 }
301
302 pub(super) fn advance_from_opening(
308 &mut self,
309 id: CircId,
310 ) -> Option<oneshot::Sender<CreateResponse>> {
311 let ok = matches!(self.m.get(&id), Some(CircEnt::Opening { .. }));
316 if ok {
317 if let Some(CircEnt::Opening {
318 create_response_sender: oneshot,
319 cell_sender: sink,
320 padding_ctrl,
321 }) = self.m.remove(&id)
322 {
323 self.m.insert(
324 id,
325 CircEnt::OpenOrigin {
326 cell_sender: sink,
327 padding_ctrl,
328 },
329 );
330 Some(oneshot)
331 } else {
332 panic!("internal error: inconsistent circuit state");
333 }
334 } else {
335 None
336 }
337 }
338
339 pub(super) fn destroy_sent(&mut self, id: CircId, hc: HalfCirc) {
343 if let Some(replaced) = self.m.insert(id, CircEnt::DestroySent(hc)) {
344 if !matches!(replaced, CircEnt::DestroySent(_)) {
345 self.open_count = self.open_count.saturating_sub(1);
347 }
348 }
349 }
350
351 pub(super) fn remove(&mut self, id: CircId) -> Option<CircEnt> {
353 self.m.remove(&id).map(|removed| {
354 if !matches!(removed, CircEnt::DestroySent(_)) {
355 self.open_count = self.open_count.saturating_sub(1);
356 }
357 removed
358 })
359 }
360
361 pub(super) fn open_ent_count(&self) -> usize {
363 self.open_count
364 }
365}
366
367#[cfg(test)]
368mod test {
369 #![allow(clippy::bool_assert_comparison)]
371 #![allow(clippy::clone_on_copy)]
372 #![allow(clippy::dbg_macro)]
373 #![allow(clippy::mixed_attributes_style)]
374 #![allow(clippy::print_stderr)]
375 #![allow(clippy::print_stdout)]
376 #![allow(clippy::single_char_pattern)]
377 #![allow(clippy::unwrap_used)]
378 #![allow(clippy::unchecked_time_subtraction)]
379 #![allow(clippy::useless_vec)]
380 #![allow(clippy::needless_pass_by_value)]
381 #![allow(clippy::string_slice)] use super::*;
384 use crate::circuit::test::fake_mpsc;
385 use crate::client::circuit::padding::new_padding;
386 use tor_basic_utils::test_rng::testing_rng;
387 use tor_rtcompat::DynTimeProvider;
388
389 #[test]
390 fn circmap_basics() {
391 let mut map_low = CircMap::new(CircIdRange::Low);
392 let mut map_high = CircMap::new(CircIdRange::High);
393 let mut ids_low: Vec<CircId> = Vec::new();
394 let mut ids_high: Vec<CircId> = Vec::new();
395 let mut rng = testing_rng();
396 tor_rtcompat::test_with_one_runtime!(|runtime| async {
397 let (padding_ctrl, _padding_stream) = new_padding(DynTimeProvider::new(runtime));
398
399 assert!(map_low.get_mut(CircId::new(77).unwrap()).is_none());
400
401 for _ in 0..128 {
402 let (csnd, _) = oneshot::channel();
403 let (snd, _) = fake_mpsc(8);
404 let id_low = map_low
405 .add_origin_ent(&mut rng, csnd, snd, padding_ctrl.clone())
406 .unwrap();
407 assert!(u32::from(id_low) > 0);
408 assert!(u32::from(id_low) < 0x80000000);
409 assert!(!ids_low.contains(&id_low));
410 ids_low.push(id_low);
411
412 assert!(matches!(
413 *map_low.get_mut(id_low).unwrap(),
414 CircEnt::Opening { .. }
415 ));
416
417 let (csnd, _) = oneshot::channel();
418 let (snd, _) = fake_mpsc(8);
419 let id_high = map_high
420 .add_origin_ent(&mut rng, csnd, snd, padding_ctrl.clone())
421 .unwrap();
422 assert!(u32::from(id_high) >= 0x80000000);
423 assert!(!ids_high.contains(&id_high));
424 ids_high.push(id_high);
425 }
426
427 assert_eq!(128, map_low.open_ent_count());
429 assert_eq!(128, map_high.open_ent_count());
430
431 assert!(map_low.get_mut(ids_low[0]).is_some());
433 map_low.remove(ids_low[0]);
434 assert!(map_low.get_mut(ids_low[0]).is_none());
435 assert_eq!(127, map_low.open_ent_count());
436
437 map_low.destroy_sent(CircId::new(256).unwrap(), HalfCirc::new(1));
439 assert_eq!(127, map_low.open_ent_count());
440
441 assert!(map_high.get_mut(ids_high[0]).is_some());
445 assert!(matches!(
446 *map_high.get_mut(ids_high[0]).unwrap(),
447 CircEnt::Opening { .. }
448 ));
449 let adv = map_high.advance_from_opening(ids_high[0]);
450 assert!(adv.is_some());
451 assert!(matches!(
452 *map_high.get_mut(ids_high[0]).unwrap(),
453 CircEnt::OpenOrigin { .. }
454 ));
455
456 let adv = map_high.advance_from_opening(ids_high[0]);
458 assert!(adv.is_none());
459
460 let adv = map_high.advance_from_opening(CircId::new(77).unwrap());
464 assert!(adv.is_none());
465 });
466 }
467}