Skip to main content

nautilus_live/
socket.rs

1// -------------------------------------------------------------------------------------------------
2//  Copyright (C) 2015-2026 Nautech Systems Pty Ltd. All rights reserved.
3//  https://nautechsystems.io
4//
5//  Licensed under the GNU Lesser General Public License Version 3.0 (the "License");
6//  You may not use this file except in compliance with the License.
7//  You may obtain a copy of the License at https://www.gnu.org/licenses/lgpl-3.0.en.html
8//
9//  Unless required by applicable law or agreed to in writing, software
10//  distributed under the License is distributed on an "AS IS" BASIS,
11//  WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12//  See the License for the specific language governing permissions and
13//  limitations under the License.
14// -------------------------------------------------------------------------------------------------
15
16//! Socket state publication and reconnect control for live clients.
17
18use 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/// Outcome from resolving a client socket through a live node registry.
46#[derive(Clone, Debug)]
47pub enum SocketReconnectLookup {
48    /// No client with the requested ID belongs to the live node.
49    ClientNotFound,
50    /// The client has no controller-reconnectable sockets.
51    Unsupported,
52    /// The client supports socket reconnects but not for the requested endpoint.
53    EndpointNotFound,
54    /// More than one client surface owns the requested endpoint.
55    AmbiguousEndpoint,
56    /// The endpoint resolved to one reconnect handle.
57    Handle(SocketReconnectHandle),
58}
59
60/// Cloneable control handle for one registered socket endpoint.
61#[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    /// Requests reconnect of the registered endpoint.
77    #[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/// Registry of reconnectable socket endpoints owned by one live node.
114#[derive(Debug, Default)]
115pub struct SocketReconnectRegistry {
116    inner: Arc<Mutex<RegistryInner>>,
117}
118
119impl SocketReconnectRegistry {
120    /// Makes this registry available while live client factories construct their clients.
121    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    /// Resolves one logical socket endpoint.
135    #[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    /// Returns the reconnect handle when exactly one owner registered `endpoint`.
158    #[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(&registry),
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(&registry, 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(&registry);
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(&registry),
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(&registry, 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(&registry, 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/// Creates socket controls which share client identity and registry ownership.
367#[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    /// Creates a new [`SocketControlFactory`] instance.
378    #[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    /// Creates a factory which registers with an explicit live socket registry.
385    #[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    /// Creates a control for one logical socket endpoint.
395    #[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/// State publisher and reconnect registration for one logical socket endpoint.
478///
479/// A clone starts without ownership of the source control's active transport generation.
480#[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    /// Creates a new [`SocketControl`] instance.
501    #[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    /// Creates a control which registers with an explicit live socket registry.
507    #[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    /// Returns a sink which publishes transport state changes as system events.
518    #[must_use]
519    pub fn sink(&self) -> SocketStateSink {
520        self.sink_with(|_| {})
521    }
522
523    /// Returns a state sink which calls `on_state` before publishing each change.
524    ///
525    /// `on_state` must not synchronously trigger another state change for the same endpoint.
526    #[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    /// Registers a reconnect request function for this endpoint's active generation.
552    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    /// Removes the reconnect handle owned by this control generation.
576    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(&registry);
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(&registry);
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(&registry);
697        let _sink = control.sink();
698        control.register(move || expected);
699
700        assert_eq!(handle(&registry).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            &registry,
710        );
711        let first = factory.control(ENDPOINT);
712        let _first_sink = first.sink();
713        first.register(|| ReconnectRequestOutcome::Accepted);
714        let stale_handle = handle(&registry);
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(&registry).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            &registry,
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(&registry);
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(&registry));
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(&registry);
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(&registry).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(&registry);
870        let _sink = control.sink();
871        control.register(|| ReconnectRequestOutcome::Accepted);
872        let retained = handle(&registry);
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(&registry);
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(&registry);
903        let second = control(&registry);
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(&registry);
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}