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, 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/// 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
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    /// Resolves one logical socket endpoint.
136    #[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    /// Returns the reconnect handle when exactly one owner registered `endpoint`.
160    #[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(&registry),
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(&registry, 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(&registry);
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(&registry),
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(&registry, 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(&registry, 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/// Creates socket controls which share client identity and registry ownership.
383#[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    /// Creates a new [`SocketControlFactory`] instance.
394    #[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    /// Creates a factory which registers with an explicit live socket registry.
401    #[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    /// Creates a control for one logical socket endpoint.
411    #[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/// State publisher and reconnect registration for one logical socket endpoint.
497///
498/// A clone starts without ownership of the source control's active transport generation.
499#[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    /// Creates a new [`SocketControl`] instance.
520    #[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    /// Creates a control which registers with an explicit live socket registry.
526    #[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    /// Returns a sink which publishes transport state changes as system events.
537    #[must_use]
538    pub fn sink(&self) -> SocketStateSink {
539        self.sink_with(|_| {})
540    }
541
542    /// Returns a state sink which calls `on_state` before publishing each change.
543    ///
544    /// `on_state` must not synchronously trigger another state change for the same endpoint.
545    #[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    /// Registers a reconnect request function for this endpoint's active generation.
573    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    /// Removes the reconnect handle owned by this control generation.
600    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(&registry);
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(&registry);
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(&registry);
723        let _sink = control.sink();
724        control.register(move || expected);
725
726        assert_eq!(handle(&registry).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            &registry,
736        );
737        let first = factory.control(ENDPOINT);
738        let _first_sink = first.sink();
739        first.register(|| ReconnectRequestOutcome::Accepted);
740        let stale_handle = handle(&registry);
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(&registry).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            &registry,
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(&registry);
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(&registry));
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(&registry);
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(&registry).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(&registry);
903        let _sink = control.sink();
904        control.register(|| ReconnectRequestOutcome::Accepted);
905        let retained = handle(&registry);
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(&registry);
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(&registry);
936        let second = control(&registry);
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(&registry);
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}