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