1use std::{
19 cell::RefCell,
20 fmt::Debug,
21 sync::{
22 Arc, Weak,
23 atomic::{AtomicU64, Ordering},
24 },
25};
26
27use ahash::{AHashMap, AHashSet};
28use nautilus_common::{
29 live::runner::try_get_system_event_sender,
30 messages::{
31 SystemEvent,
32 system::{SocketState as SystemSocketState, SocketStateChange},
33 },
34};
35use nautilus_model::identifiers::{ClientId, Venue};
36pub use nautilus_network::mode::ReconnectRequestOutcome as SocketReconnectRequestOutcome;
37use nautilus_network::{SocketState, SocketStateSink, mode::ReconnectRequestOutcome};
38use parking_lot::Mutex;
39use ustr::Ustr;
40
41thread_local! {
42 static SOCKET_REGISTRARS: RefCell<Vec<SocketReconnectRegistrar>> = const { RefCell::new(Vec::new()) };
43}
44
45#[derive(Clone, Debug)]
47pub enum SocketReconnectLookup {
48 ClientNotFound,
50 Unsupported,
52 EndpointNotFound,
54 AmbiguousEndpoint,
56 Handle(SocketReconnectHandle),
58}
59
60#[derive(Clone)]
62pub struct SocketReconnectHandle {
63 request: Arc<dyn Fn() -> ReconnectRequestOutcome + Send + Sync>,
64}
65
66impl SocketReconnectHandle {
67 fn new<F>(request: F) -> Self
68 where
69 F: Fn() -> ReconnectRequestOutcome + Send + Sync + 'static,
70 {
71 Self {
72 request: Arc::new(request),
73 }
74 }
75
76 #[must_use]
78 pub fn request_reconnect(&self) -> ReconnectRequestOutcome {
79 (self.request)()
80 }
81}
82
83impl Debug for SocketReconnectHandle {
84 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
85 f.debug_struct(stringify!(SocketReconnectHandle))
86 .finish_non_exhaustive()
87 }
88}
89
90#[derive(Clone, Copy, Debug, Eq, Hash, PartialEq)]
91struct SocketEndpoint {
92 client_id: ClientId,
93 endpoint: Ustr,
94}
95
96#[derive(Debug)]
97struct RegistryEntry {
98 generation: u64,
99 handle: SocketReconnectHandle,
100 request: Arc<Mutex<Option<SocketReconnectHandle>>>,
101}
102
103#[derive(Debug, Default)]
104struct RegistryInner {
105 clients: AHashSet<ClientId>,
106 supported: AHashSet<ClientId>,
107 entries: AHashMap<SocketEndpoint, AHashMap<u64, RegistryEntry>>,
108 owners: AHashMap<ClientId, AHashSet<u64>>,
109 next_generation: u64,
110 next_owner: u64,
111}
112
113#[derive(Debug, Default)]
115pub struct SocketReconnectRegistry {
116 inner: Arc<Mutex<RegistryInner>>,
117}
118
119impl SocketReconnectRegistry {
120 pub fn scope<T>(&self, f: impl FnOnce() -> T) -> T {
122 SOCKET_REGISTRARS.with(|registrars| {
123 registrars.borrow_mut().push(self.registrar());
124 });
125 let _scope = SocketRegistryScope;
126 f()
127 }
128
129 #[cfg(any(feature = "node", test))]
130 pub(crate) fn register_client(&self, client_id: ClientId) {
131 self.inner.lock().clients.insert(client_id);
132 }
133
134 #[must_use]
136 pub fn get(&self, client_id: ClientId, endpoint: Ustr) -> SocketReconnectLookup {
137 let inner = self.inner.lock();
138 let key = SocketEndpoint {
139 client_id,
140 endpoint,
141 };
142
143 match inner.entries.get(&key) {
144 Some(entries) if entries.len() > 1 => SocketReconnectLookup::AmbiguousEndpoint,
145 Some(entries) => entries
146 .values()
147 .next()
148 .map_or(SocketReconnectLookup::EndpointNotFound, |entry| {
149 SocketReconnectLookup::Handle(entry.handle.clone())
150 }),
151 None if inner.supported.contains(&client_id) => SocketReconnectLookup::EndpointNotFound,
152 None if inner.clients.contains(&client_id) => SocketReconnectLookup::Unsupported,
153 None => SocketReconnectLookup::ClientNotFound,
154 }
155 }
156
157 #[must_use]
159 pub fn handle(&self, client_id: ClientId, endpoint: Ustr) -> Option<SocketReconnectHandle> {
160 match self.get(client_id, endpoint) {
161 SocketReconnectLookup::Handle(handle) => Some(handle),
162 _ => None,
163 }
164 }
165
166 fn registrar(&self) -> SocketReconnectRegistrar {
167 SocketReconnectRegistrar {
168 registry: Arc::downgrade(&self.inner),
169 }
170 }
171}
172
173struct SocketRegistryScope;
174
175impl Drop for SocketRegistryScope {
176 fn drop(&mut self) {
177 SOCKET_REGISTRARS.with(|registrars| {
178 registrars.borrow_mut().pop();
179 });
180 }
181}
182
183#[derive(Clone, Debug)]
184struct SocketReconnectRegistrar {
185 registry: Weak<Mutex<RegistryInner>>,
186}
187
188impl SocketReconnectRegistrar {
189 fn owner(&self, client_id: ClientId) -> Option<SocketReconnectOwner> {
190 let registry = self.registry.upgrade()?;
191 let owner_id = {
192 let mut inner = registry.lock();
193 inner.next_owner = inner.next_owner.wrapping_add(1).max(1);
194 let owner_id = inner.next_owner;
195 inner.supported.insert(client_id);
196 inner.owners.entry(client_id).or_default().insert(owner_id);
197 owner_id
198 };
199
200 Some(SocketReconnectOwner(Arc::new(SocketReconnectOwnerInner {
201 registry: Arc::downgrade(®istry),
202 client_id,
203 owner_id,
204 })))
205 }
206}
207
208#[derive(Clone, Debug)]
209struct SocketReconnectOwner(Arc<SocketReconnectOwnerInner>);
210
211impl SocketReconnectOwner {
212 fn register(
213 &self,
214 endpoint: Ustr,
215 handle: SocketReconnectHandle,
216 ) -> (Option<SocketReconnectRegistration>, Option<RegistryEntry>) {
217 let Some(registry) = self.0.registry.upgrade() else {
218 return (None, None);
219 };
220 let key = SocketEndpoint {
221 client_id: self.0.client_id,
222 endpoint,
223 };
224
225 let replaced = remove_entry(®istry, key, self.0.owner_id, None);
226
227 let request = Arc::new(Mutex::new(Some(handle)));
228 let guarded_request = Arc::clone(&request);
229 let registry_ref = Arc::downgrade(®istry);
230 let guarded_handle = SocketReconnectHandle::new(move || {
231 let Some(_registry) = registry_ref.upgrade() else {
232 return ReconnectRequestOutcome::Closed;
233 };
234 let request = guarded_request.lock();
235 request
236 .as_ref()
237 .map_or(ReconnectRequestOutcome::Closed, |handle| {
238 handle.request_reconnect()
239 })
240 });
241 let generation = {
242 let mut inner = registry.lock();
243 inner.next_generation = inner.next_generation.wrapping_add(1).max(1);
244 let generation = inner.next_generation;
245 inner.entries.entry(key).or_default().insert(
246 self.0.owner_id,
247 RegistryEntry {
248 generation,
249 handle: guarded_handle,
250 request,
251 },
252 );
253 generation
254 };
255
256 (
257 Some(SocketReconnectRegistration {
258 registry: Arc::downgrade(®istry),
259 key,
260 owner_id: self.0.owner_id,
261 generation,
262 }),
263 replaced,
264 )
265 }
266
267 fn remove(&self, endpoint: Ustr) -> Option<RegistryEntry> {
268 let registry = self.0.registry.upgrade()?;
269 let key = SocketEndpoint {
270 client_id: self.0.client_id,
271 endpoint,
272 };
273 remove_entry(®istry, key, self.0.owner_id, None)
274 }
275}
276
277#[derive(Debug)]
278struct SocketReconnectOwnerInner {
279 registry: Weak<Mutex<RegistryInner>>,
280 client_id: ClientId,
281 owner_id: u64,
282}
283
284impl Drop for SocketReconnectOwnerInner {
285 fn drop(&mut self) {
286 let Some(registry) = self.registry.upgrade() else {
287 return;
288 };
289 let removed = {
290 let mut inner = registry.lock();
291 if let Some(owners) = inner.owners.get_mut(&self.client_id) {
292 owners.remove(&self.owner_id);
293 if owners.is_empty() {
294 inner.owners.remove(&self.client_id);
295 }
296 }
297
298 let mut removed = Vec::new();
299 inner.entries.retain(|key, entries| {
300 if key.client_id == self.client_id
301 && let Some(entry) = entries.remove(&self.owner_id)
302 {
303 removed.push(entry);
304 }
305 !entries.is_empty()
306 });
307 removed
308 };
309
310 for entry in removed {
311 deactivate(Some(entry));
312 }
313 }
314}
315
316#[derive(Debug)]
317struct SocketReconnectRegistration {
318 registry: Weak<Mutex<RegistryInner>>,
319 key: SocketEndpoint,
320 owner_id: u64,
321 generation: u64,
322}
323
324impl Drop for SocketReconnectRegistration {
325 fn drop(&mut self) {
326 let Some(registry) = self.registry.upgrade() else {
327 return;
328 };
329 let entry = remove_entry(®istry, self.key, self.owner_id, Some(self.generation));
330 deactivate(entry);
331 }
332}
333
334fn remove_entry(
335 registry: &Arc<Mutex<RegistryInner>>,
336 key: SocketEndpoint,
337 owner_id: u64,
338 generation: Option<u64>,
339) -> Option<RegistryEntry> {
340 let mut inner = registry.lock();
341 let mut entry = None;
342 let mut remove_endpoint = false;
343
344 if let Some(entries) = inner.entries.get_mut(&key) {
345 let matches = entries
346 .get(&owner_id)
347 .is_some_and(|entry| generation.is_none_or(|value| entry.generation == value));
348 if matches {
349 entry = entries.remove(&owner_id);
350 }
351 remove_endpoint = entries.is_empty();
352 }
353
354 if remove_endpoint {
355 inner.entries.remove(&key);
356 }
357 entry
358}
359
360fn deactivate(entry: Option<RegistryEntry>) {
361 if let Some(entry) = entry {
362 *entry.request.lock() = None;
363 }
364}
365
366#[derive(Clone, Debug)]
368pub struct SocketControlFactory {
369 client_id: ClientId,
370 venue: Option<Venue>,
371 sender: Option<tokio::sync::mpsc::UnboundedSender<SystemEvent>>,
372 owner: Option<SocketReconnectOwner>,
373 controls: Arc<Mutex<AHashMap<Ustr, SocketControl>>>,
374}
375
376impl SocketControlFactory {
377 #[must_use]
379 pub fn new(client_id: ClientId, venue: Option<Venue>) -> Self {
380 let registrar = SOCKET_REGISTRARS.with(|registrars| registrars.borrow().last().cloned());
381 Self::from_registrar(client_id, venue, registrar)
382 }
383
384 #[must_use]
386 pub fn with_registry(
387 client_id: ClientId,
388 venue: Option<Venue>,
389 registry: &SocketReconnectRegistry,
390 ) -> Self {
391 Self::from_registrar(client_id, venue, Some(registry.registrar()))
392 }
393
394 #[must_use]
396 pub fn control(&self, endpoint: impl AsRef<str>) -> SocketControl {
397 let endpoint = Ustr::from(endpoint.as_ref());
398 let mut controls = self.controls.lock();
399 if let Some(control) = controls.get(&endpoint) {
400 return control.clone();
401 }
402
403 let control = SocketControl {
404 publisher: SocketStatePublisher {
405 client_id: self.client_id,
406 venue: self.venue,
407 endpoint,
408 sender: self.sender.clone(),
409 active_generation: Arc::new(AtomicU64::new(0)),
410 publish_lock: Arc::new(Mutex::new(())),
411 },
412 owner: self.owner.clone(),
413 generation: AtomicU64::new(0),
414 registration: Mutex::new(None),
415 };
416 controls.insert(endpoint, control.clone());
417 control
418 }
419
420 fn from_registrar(
421 client_id: ClientId,
422 venue: Option<Venue>,
423 registrar: Option<SocketReconnectRegistrar>,
424 ) -> Self {
425 Self {
426 client_id,
427 venue,
428 sender: try_get_system_event_sender(),
429 owner: registrar.and_then(|registrar| registrar.owner(client_id)),
430 controls: Arc::new(Mutex::new(AHashMap::new())),
431 }
432 }
433}
434
435#[derive(Clone, Debug)]
436struct SocketStatePublisher {
437 client_id: ClientId,
438 venue: Option<Venue>,
439 endpoint: Ustr,
440 sender: Option<tokio::sync::mpsc::UnboundedSender<SystemEvent>>,
441 active_generation: Arc<AtomicU64>,
442 publish_lock: Arc<Mutex<()>>,
443}
444
445impl SocketStatePublisher {
446 fn publish(&self, state: SocketState) {
447 let Some(sender) = &self.sender else {
448 return;
449 };
450 let state = match state {
451 SocketState::Connected => SystemSocketState::Connected,
452 SocketState::Disconnected => SystemSocketState::Disconnected,
453 };
454 let change = SocketStateChange::new(self.client_id, self.venue, self.endpoint, state);
455 if let Err(e) = sender.send(SystemEvent::SocketState(change)) {
456 log::error!("Failed to emit socket state change: {e}");
457 }
458 }
459
460 fn is_current(&self, generation: u64) -> bool {
461 self.active_generation.load(Ordering::Acquire) == generation
462 }
463
464 fn publish_if_current<F>(&self, generation: u64, state: SocketState, on_state: &F)
465 where
466 F: Fn(SocketState),
467 {
468 let _guard = self.publish_lock.lock();
469
470 if self.is_current(generation) {
471 on_state(state);
472 self.publish(state);
473 }
474 }
475}
476
477#[derive(Debug)]
481pub struct SocketControl {
482 publisher: SocketStatePublisher,
483 owner: Option<SocketReconnectOwner>,
484 generation: AtomicU64,
485 registration: Mutex<Option<SocketReconnectRegistration>>,
486}
487
488impl Clone for SocketControl {
489 fn clone(&self) -> Self {
490 Self {
491 publisher: self.publisher.clone(),
492 owner: self.owner.clone(),
493 generation: AtomicU64::new(0),
494 registration: Mutex::new(None),
495 }
496 }
497}
498
499impl SocketControl {
500 #[must_use]
502 pub fn new(client_id: ClientId, venue: Option<Venue>, endpoint: impl AsRef<str>) -> Self {
503 SocketControlFactory::new(client_id, venue).control(endpoint)
504 }
505
506 #[must_use]
508 pub fn with_registry(
509 client_id: ClientId,
510 venue: Option<Venue>,
511 endpoint: impl AsRef<str>,
512 registry: &SocketReconnectRegistry,
513 ) -> Self {
514 SocketControlFactory::with_registry(client_id, venue, registry).control(endpoint)
515 }
516
517 #[must_use]
519 pub fn sink(&self) -> SocketStateSink {
520 self.sink_with(|_| {})
521 }
522
523 #[must_use]
527 pub fn sink_with<F>(&self, on_state: F) -> SocketStateSink
528 where
529 F: Fn(SocketState) + Send + Sync + 'static,
530 {
531 let (registration, replaced, generation) = {
532 let _guard = self.publisher.publish_lock.lock();
533 let replaced = self
534 .owner
535 .as_ref()
536 .and_then(|owner| owner.remove(self.publisher.endpoint));
537 let generation = advance_generation(&self.publisher.active_generation);
538 self.generation.store(generation, Ordering::Release);
539 let registration = self.registration.lock().take();
540 (registration, replaced, generation)
541 };
542 deactivate(replaced);
543 drop(registration);
544
545 let publisher = self.publisher.clone();
546 SocketStateSink::new(move |state| {
547 publisher.publish_if_current(generation, state, &on_state);
548 })
549 }
550
551 pub fn register<F>(&self, request: F)
553 where
554 F: Fn() -> ReconnectRequestOutcome + Send + Sync + 'static,
555 {
556 let (replaced, old_registration) = {
557 let _guard = self.publisher.publish_lock.lock();
558 let generation = self.generation.load(Ordering::Acquire);
559 if generation == 0 || !self.publisher.is_current(generation) {
560 return;
561 }
562 let (registration, replaced) = if let Some(owner) = &self.owner {
563 owner.register(self.publisher.endpoint, SocketReconnectHandle::new(request))
564 } else {
565 (None, None)
566 };
567 let mut current = self.registration.lock();
568 let old_registration = std::mem::replace(&mut *current, registration);
569 (replaced, old_registration)
570 };
571 deactivate(replaced);
572 drop(old_registration);
573 }
574
575 pub fn deregister(&self) {
577 let registration = {
578 let _guard = self.publisher.publish_lock.lock();
579 self.generation.store(0, Ordering::Release);
580 self.registration.lock().take()
581 };
582 drop(registration);
583 }
584}
585
586impl Drop for SocketControl {
587 fn drop(&mut self) {
588 self.deregister();
589 }
590}
591
592fn advance_generation(counter: &AtomicU64) -> u64 {
593 let mut current = counter.load(Ordering::Relaxed);
594 loop {
595 let next = current.wrapping_add(1).max(1);
596 match counter.compare_exchange_weak(current, next, Ordering::Release, Ordering::Relaxed) {
597 Ok(_) => return next,
598 Err(actual) => current = actual,
599 }
600 }
601}
602
603#[cfg(test)]
604mod tests {
605 use std::{
606 sync::{
607 Arc, Barrier,
608 atomic::{AtomicUsize, Ordering},
609 },
610 thread,
611 time::{Duration, Instant},
612 };
613
614 use rstest::rstest;
615
616 use super::*;
617
618 const ENDPOINT: &str = "test-streams";
619
620 fn control(registry: &SocketReconnectRegistry) -> SocketControl {
621 SocketControl::with_registry(
622 ClientId::from("TEST"),
623 Some(Venue::from("TEST")),
624 ENDPOINT,
625 registry,
626 )
627 }
628
629 fn handle(registry: &SocketReconnectRegistry) -> SocketReconnectHandle {
630 let SocketReconnectLookup::Handle(handle) =
631 registry.get(ClientId::from("TEST"), Ustr::from(ENDPOINT))
632 else {
633 panic!("test socket endpoint should be registered");
634 };
635 handle
636 }
637
638 #[rstest]
639 fn publishes_endpoint_state() {
640 let registry = SocketReconnectRegistry::default();
641 let (sender, mut receiver) = tokio::sync::mpsc::unbounded_channel();
642 let mut control = control(®istry);
643 control.publisher.sender = Some(sender);
644 let _sink = control.sink();
645
646 control.publisher.publish(SocketState::Connected);
647 control.publisher.publish(SocketState::Disconnected);
648
649 let SystemEvent::SocketState(connected) = receiver.try_recv().unwrap();
650 let SystemEvent::SocketState(disconnected) = receiver.try_recv().unwrap();
651 assert_eq!(connected.client_id, ClientId::from("TEST"));
652 assert_eq!(connected.venue, Some(Venue::from("TEST")));
653 assert_eq!(connected.endpoint, Ustr::from(ENDPOINT));
654 assert_eq!(connected.state, SystemSocketState::Connected);
655 assert_eq!(disconnected.client_id, ClientId::from("TEST"));
656 assert_eq!(disconnected.venue, Some(Venue::from("TEST")));
657 assert_eq!(disconnected.endpoint, Ustr::from(ENDPOINT));
658 assert_eq!(disconnected.state, SystemSocketState::Disconnected);
659 }
660
661 #[rstest]
662 fn replacement_sink_suppresses_stale_transport_state() {
663 let registry = SocketReconnectRegistry::default();
664 let (sender, mut receiver) = tokio::sync::mpsc::unbounded_channel();
665 let mut first = control(®istry);
666 first.publisher.sender = Some(sender.clone());
667 let mut replacement = first.clone();
668 replacement.publisher.sender = Some(sender);
669 let _stale_sink = first.sink();
670 let stale_generation = first.generation.load(Ordering::Acquire);
671 let _current_sink = replacement.sink();
672 let current_generation = replacement.generation.load(Ordering::Acquire);
673
674 first
675 .publisher
676 .publish_if_current(stale_generation, SocketState::Disconnected, &|_| {});
677 replacement.publisher.publish_if_current(
678 current_generation,
679 SocketState::Connected,
680 &|_| {},
681 );
682
683 let SystemEvent::SocketState(change) = receiver.try_recv().unwrap();
684 assert_eq!(change.state, SystemSocketState::Connected);
685 assert!(receiver.try_recv().is_err());
686 }
687
688 #[rstest]
689 #[case(ReconnectRequestOutcome::Accepted)]
690 #[case(ReconnectRequestOutcome::AlreadyReconnecting)]
691 #[case(ReconnectRequestOutcome::Disconnected)]
692 #[case(ReconnectRequestOutcome::Closed)]
693 #[case(ReconnectRequestOutcome::Unsupported)]
694 fn register_preserves_transport_outcome(#[case] expected: ReconnectRequestOutcome) {
695 let registry = SocketReconnectRegistry::default();
696 let control = control(®istry);
697 let _sink = control.sink();
698 control.register(move || expected);
699
700 assert_eq!(handle(®istry).request_reconnect(), expected);
701 }
702
703 #[rstest]
704 fn replacement_invalidates_stale_handle() {
705 let registry = SocketReconnectRegistry::default();
706 let factory = SocketControlFactory::with_registry(
707 ClientId::from("TEST"),
708 Some(Venue::from("TEST")),
709 ®istry,
710 );
711 let first = factory.control(ENDPOINT);
712 let _first_sink = first.sink();
713 first.register(|| ReconnectRequestOutcome::Accepted);
714 let stale_handle = handle(®istry);
715 let replacement = factory.control(ENDPOINT);
716 let _replacement_sink = replacement.sink();
717 replacement.register(|| ReconnectRequestOutcome::AlreadyReconnecting);
718
719 first.register(|| ReconnectRequestOutcome::Disconnected);
720 first.deregister();
721
722 assert_eq!(
723 stale_handle.request_reconnect(),
724 ReconnectRequestOutcome::Closed
725 );
726 assert_eq!(
727 handle(®istry).request_reconnect(),
728 ReconnectRequestOutcome::AlreadyReconnecting
729 );
730 }
731
732 #[rstest]
733 fn replacement_waits_for_an_inflight_request_that_publishes_state() {
734 let registry = SocketReconnectRegistry::default();
735 let factory = SocketControlFactory::with_registry(
736 ClientId::from("TEST"),
737 Some(Venue::from("TEST")),
738 ®istry,
739 );
740 let first = factory.control(ENDPOINT);
741 let _first_sink = first.sink();
742 let generation = first.generation.load(Ordering::Acquire);
743 let publisher = first.publisher.clone();
744 let entered = Arc::new(Barrier::new(2));
745 let release = Arc::new(Barrier::new(2));
746 let request_entered = Arc::clone(&entered);
747 let request_release = Arc::clone(&release);
748 first.register(move || {
749 request_entered.wait();
750 request_release.wait();
751 publisher.publish_if_current(generation, SocketState::Disconnected, &|_| {});
752 ReconnectRequestOutcome::Accepted
753 });
754 let stale_handle = handle(®istry);
755 let request_handle = stale_handle.clone();
756 let request = thread::spawn(move || request_handle.request_reconnect());
757 entered.wait();
758
759 let replacement = factory.control(ENDPOINT);
760 let (replaced_tx, replaced_rx) = std::sync::mpsc::channel();
761
762 let replace = thread::spawn(move || {
763 let _replacement_sink = replacement.sink();
764 replaced_tx.send(()).unwrap();
765 });
766 let deadline = Instant::now() + Duration::from_secs(1);
767
768 while !matches!(
769 registry.get(ClientId::from("TEST"), Ustr::from(ENDPOINT)),
770 SocketReconnectLookup::EndpointNotFound
771 ) {
772 assert!(
773 Instant::now() < deadline,
774 "replacement did not remove the stale entry"
775 );
776 thread::yield_now();
777 }
778
779 assert!(replaced_rx.try_recv().is_err());
780 release.wait();
781 assert_eq!(request.join().unwrap(), ReconnectRequestOutcome::Accepted);
782 replace.join().unwrap();
783 assert!(replaced_rx.try_recv().is_ok());
784 assert_eq!(
785 stale_handle.request_reconnect(),
786 ReconnectRequestOutcome::Closed
787 );
788 }
789
790 #[rstest]
791 fn reregistration_waits_for_an_inflight_request_that_publishes_state() {
792 let registry = SocketReconnectRegistry::default();
793 let control = Arc::new(control(®istry));
794 let _sink = control.sink();
795 let generation = control.generation.load(Ordering::Acquire);
796 let publisher = control.publisher.clone();
797 let entered = Arc::new(Barrier::new(2));
798 let release = Arc::new(Barrier::new(2));
799 let request_entered = Arc::clone(&entered);
800 let request_release = Arc::clone(&release);
801 control.register(move || {
802 request_entered.wait();
803 request_release.wait();
804 publisher.publish_if_current(generation, SocketState::Disconnected, &|_| {});
805 ReconnectRequestOutcome::Accepted
806 });
807 let stale_handle = handle(®istry);
808 let key = SocketEndpoint {
809 client_id: ClientId::from("TEST"),
810 endpoint: Ustr::from(ENDPOINT),
811 };
812 let old_generation = registry
813 .inner
814 .lock()
815 .entries
816 .get(&key)
817 .and_then(|entries| entries.values().next())
818 .map(|entry| entry.generation)
819 .unwrap();
820 let request_handle = stale_handle.clone();
821 let request = thread::spawn(move || request_handle.request_reconnect());
822 entered.wait();
823
824 let replacement = Arc::clone(&control);
825 let (registered_tx, registered_rx) = std::sync::mpsc::channel();
826
827 let register = thread::spawn(move || {
828 replacement.register(|| ReconnectRequestOutcome::AlreadyReconnecting);
829 registered_tx.send(()).unwrap();
830 });
831 let deadline = Instant::now() + Duration::from_secs(1);
832
833 loop {
834 let old_entry_is_registered = registry
835 .inner
836 .lock()
837 .entries
838 .get(&key)
839 .and_then(|entries| entries.values().next())
840 .is_some_and(|entry| entry.generation == old_generation);
841 if !old_entry_is_registered {
842 break;
843 }
844 assert!(
845 Instant::now() < deadline,
846 "replacement reconnect handle was not registered"
847 );
848 thread::yield_now();
849 }
850
851 assert!(registered_rx.try_recv().is_err());
852 release.wait();
853 assert_eq!(request.join().unwrap(), ReconnectRequestOutcome::Accepted);
854 register.join().unwrap();
855 assert!(registered_rx.try_recv().is_ok());
856 assert_eq!(
857 stale_handle.request_reconnect(),
858 ReconnectRequestOutcome::Closed
859 );
860 assert_eq!(
861 handle(®istry).request_reconnect(),
862 ReconnectRequestOutcome::AlreadyReconnecting
863 );
864 }
865
866 #[rstest]
867 fn registry_drop_revokes_a_retained_handle() {
868 let registry = SocketReconnectRegistry::default();
869 let control = control(®istry);
870 let _sink = control.sink();
871 control.register(|| ReconnectRequestOutcome::Accepted);
872 let retained = handle(®istry);
873
874 drop(registry);
875
876 assert_eq!(
877 retained.request_reconnect(),
878 ReconnectRequestOutcome::Closed
879 );
880 }
881
882 #[rstest]
883 fn deregister_prevents_delayed_registration() {
884 let registry = SocketReconnectRegistry::default();
885 let control = control(®istry);
886 let _sink = control.sink();
887
888 control.deregister();
889 control.register(|| ReconnectRequestOutcome::Accepted);
890
891 assert!(matches!(
892 registry.get(ClientId::from("TEST"), Ustr::from(ENDPOINT)),
893 SocketReconnectLookup::EndpointNotFound
894 ));
895 }
896
897 #[rstest]
898 fn ambiguous_endpoint_does_not_invoke_either_owner() {
899 let registry = SocketReconnectRegistry::default();
900 let first_count = Arc::new(AtomicUsize::new(0));
901 let second_count = Arc::new(AtomicUsize::new(0));
902 let first = control(®istry);
903 let second = control(®istry);
904 let _first_sink = first.sink();
905 let _second_sink = second.sink();
906 let first_callback = Arc::clone(&first_count);
907 first.register(move || {
908 first_callback.fetch_add(1, Ordering::SeqCst);
909 ReconnectRequestOutcome::Accepted
910 });
911 let second_callback = Arc::clone(&second_count);
912 second.register(move || {
913 second_callback.fetch_add(1, Ordering::SeqCst);
914 ReconnectRequestOutcome::Accepted
915 });
916
917 assert!(matches!(
918 registry.get(ClientId::from("TEST"), Ustr::from(ENDPOINT)),
919 SocketReconnectLookup::AmbiguousEndpoint
920 ));
921 assert_eq!(first_count.load(Ordering::SeqCst), 0);
922 assert_eq!(second_count.load(Ordering::SeqCst), 0);
923 }
924
925 #[rstest]
926 fn distinguishes_unknown_unsupported_and_missing_endpoint() {
927 let registry = SocketReconnectRegistry::default();
928 registry.register_client(ClientId::from("UNSUPPORTED"));
929 let _control = control(®istry);
930
931 assert!(matches!(
932 registry.get(ClientId::from("UNKNOWN"), Ustr::from(ENDPOINT)),
933 SocketReconnectLookup::ClientNotFound
934 ));
935 assert!(matches!(
936 registry.get(ClientId::from("UNSUPPORTED"), Ustr::from(ENDPOINT)),
937 SocketReconnectLookup::Unsupported
938 ));
939 assert!(matches!(
940 registry.get(ClientId::from("TEST"), Ustr::from("unknown")),
941 SocketReconnectLookup::EndpointNotFound
942 ));
943 }
944
945 #[rstest]
946 fn scope_restores_the_prior_registry() {
947 let outer = SocketReconnectRegistry::default();
948 let inner = SocketReconnectRegistry::default();
949 let (outer_control, inner_control) = outer.scope(|| {
950 let outer_control =
951 SocketControl::new(ClientId::from("OUTER"), Some(Venue::from("TEST")), ENDPOINT);
952 let inner_control = inner.scope(|| {
953 SocketControl::new(ClientId::from("INNER"), Some(Venue::from("TEST")), ENDPOINT)
954 });
955 let _sink = inner_control.sink();
956 inner_control.register(|| ReconnectRequestOutcome::Accepted);
957 (outer_control, inner_control)
958 });
959 let _sink = outer_control.sink();
960 outer_control.register(|| ReconnectRequestOutcome::Accepted);
961
962 assert!(matches!(
963 outer.get(ClientId::from("OUTER"), Ustr::from(ENDPOINT)),
964 SocketReconnectLookup::Handle(_)
965 ));
966 assert!(matches!(
967 inner.get(ClientId::from("INNER"), Ustr::from(ENDPOINT)),
968 SocketReconnectLookup::Handle(_)
969 ));
970 drop(inner_control);
971 }
972}