Skip to main content

ankurah_core/peer_subscription/
client_relay.rs

1// TODO: Rename this module from client_relay to remote_subscription for clarity
2use ankurah_proto::{self as proto, CollectionId};
3use anyhow::anyhow;
4use async_trait::async_trait;
5use std::collections::HashMap;
6use std::sync::{Arc, OnceLock};
7use tracing::{debug, warn};
8
9use crate::error::{RequestError, RetrievalError};
10use crate::node::ContextData;
11use crate::util::safeset::SafeSet;
12
13/// Trait for query initialization that can be driven by SubscriptionRelay
14/// Abstracts the relay's interaction with LiveQuery
15#[async_trait::async_trait]
16pub trait RemoteQuerySubscriber: Clone + Send + Sync + 'static {
17    /// Called after remote subscription deltas have been applied
18    /// Dispatches to initialize (version 1) or update_selection_init (version >1) internally
19    /// Handles marking initialization as complete and setting last_error on failure
20    async fn subscription_established(&self, version: u32);
21
22    /// Set the last error for this subscription
23    fn set_last_error(&self, error: RetrievalError);
24}
25
26#[derive(Debug, Clone)]
27pub enum Status {
28    PendingRemote,
29    Requested(proto::EntityId, u32),     // peer_id, version
30    Established(proto::EntityId, u32),   // peer_id, version
31    PendingUpdate(proto::EntityId, u32), // peer_id, version
32    /// Non-retryable
33    Failed,
34}
35
36#[derive(Debug)]
37pub struct Content<CD: ContextData> {
38    pub query_id: proto::QueryId,
39    pub collection_id: CollectionId,
40    pub selection: ankql::ast::Selection,
41    pub context_data: CD,
42    pub version: u32,
43}
44
45pub struct RemoteQueryState<CD: ContextData, Q: RemoteQuerySubscriber> {
46    pub content: Arc<Content<CD>>,
47    pub status: Status,
48    pub livequery: Q,
49}
50
51struct SubscriptionRelayInner<CD: ContextData, Q: RemoteQuerySubscriber> {
52    // All subscription information in one place
53    subscriptions: std::sync::Mutex<HashMap<proto::QueryId, RemoteQueryState<CD, Q>>>,
54    // Track connected durable peers
55    connected_peers: SafeSet<proto::EntityId>,
56    // Node for communicating with remote peers
57    node: OnceLock<Arc<dyn TNode<CD>>>,
58    // Shutdown signal for retry task - when dropped, the task will stop
59    _shutdown_tx: tokio::sync::mpsc::Sender<()>,
60}
61
62/// Manages predicate registration on remote peer reactor subscriptions.
63///
64/// The SubscriptionRelay provides a resilient, event-driven approach to managing which predicates
65/// are registered with remote durable peers. It automatically handles:
66/// - Registering predicates on peer reactor subscriptions when peers connect
67/// - Re-registering predicates when peers disconnect and reconnect
68/// - Retrying failed predicate registration attempts
69/// - Clean teardown when predicates are removed
70/// - Storing ContextData for each predicate to enable proper authorization
71///
72/// This design separates predicate management concerns from the main Node implementation,
73/// making it easier to test and reason about predicate lifecycle management.
74///
75/// # Public API (for Node integration)
76///
77/// - `subscribe_predicate()` - Call when local subscriptions are created (parallel to reactor.subscribe)
78/// - `unsubscribe_predicate()` - Call when local subscriptions are removed (parallel to reactor.unsubscribe)
79/// - `notify_peer_connected()` - Call when durable peers connect (triggers automatic predicate registration)
80/// - `notify_peer_disconnected()` - Call when durable peers disconnect (orphans predicate registrations)
81/// - `get_status()` - Query current state of a predicate registration
82///
83/// # Internal/Testing API
84///
85/// - `setup_remote_subscriptions()` - Internal method for triggering predicate registration with specific peers
86///   (called automatically by notify_peer_connected, but exposed for testing)
87///
88/// The relay will automatically handle predicate registration/teardown asynchronously.
89#[derive(Clone)]
90pub struct SubscriptionRelay<CD: ContextData, Q: RemoteQuerySubscriber> {
91    inner: Arc<SubscriptionRelayInner<CD, Q>>,
92}
93
94impl<CD: ContextData, Q: RemoteQuerySubscriber> Default for SubscriptionRelay<CD, Q> {
95    fn default() -> Self { Self::new() }
96}
97
98impl<CD: ContextData, Q: RemoteQuerySubscriber> SubscriptionRelay<CD, Q> {
99    pub fn new() -> Self {
100        let (shutdown_tx, shutdown_rx) = tokio::sync::mpsc::channel(1);
101
102        let relay = Self {
103            inner: Arc::new(SubscriptionRelayInner {
104                subscriptions: std::sync::Mutex::new(HashMap::new()),
105                connected_peers: SafeSet::new(),
106                node: OnceLock::new(),
107                _shutdown_tx: shutdown_tx,
108            }),
109        };
110
111        // Start background retry task
112        relay.start_retry_task(shutdown_rx);
113
114        relay
115    }
116
117    /// Inject the node (typically a WeakNode for production)
118    ///
119    /// This should be called once during initialization. Returns an error if
120    /// the node has already been set.
121    pub fn set_node(&self, node: Arc<dyn TNode<CD>>) -> Result<(), ()> { self.inner.node.set(node).map_err(|_| ()) }
122
123    /// Notify the relay that a new predicate needs to be registered on remote peer subscriptions
124    ///
125    /// This should be called whenever a local subscription is established. The relay will
126    /// track this predicate and automatically attempt to register it with available durable peers.
127    pub fn subscribe_query(
128        &self,
129        query_id: proto::QueryId,
130        collection_id: CollectionId,
131        selection: ankql::ast::Selection,
132        context_data: CD,
133        version: u32,
134        livequery: Q,
135    ) {
136        debug!("SubscriptionRelay.subscribe_predicate() - New predicate {} needs remote registration", query_id);
137        {
138            self.inner.subscriptions.lock().expect("poisoned lock").insert(
139                query_id,
140                RemoteQueryState {
141                    content: Arc::new(Content { collection_id, selection, context_data, query_id, version }),
142                    status: Status::PendingRemote,
143                    livequery,
144                },
145            );
146        }
147
148        // Immediately attempt setup with available peers
149        if !self.inner.connected_peers.is_empty() {
150            self.setup_remote_subscriptions()
151        }
152    }
153    pub fn update_query(&self, query_id: proto::QueryId, selection: ankql::ast::Selection, version: u32) -> Result<(), anyhow::Error> {
154        debug!("SubscriptionRelay.update_query() - New query {} needs remote registration", query_id);
155
156        let update = {
157            let mut subscriptions = self.inner.subscriptions.lock().expect("poisoned lock");
158            match subscriptions.get_mut(&query_id) {
159                Some(state) => {
160                    // Update the content with new predicate and version
161                    let old_content = &state.content;
162                    state.content = Arc::new(Content {
163                        collection_id: old_content.collection_id.clone(),
164                        selection: selection.clone(),
165                        context_data: old_content.context_data.clone(),
166                        query_id: old_content.query_id,
167                        version,
168                    });
169
170                    match state.status {
171                        Status::Established(peer_id, _old_version) => {
172                            // Update to new version, mark as requested for this peer
173                            state.status = Status::Requested(peer_id, version);
174                            Some((peer_id, state.content.collection_id.clone(), state.content.context_data.clone()))
175                            // Return the peer_id to send update to
176                        }
177                        _ => {
178                            // Not established yet, just update to PendingRemote and setup
179                            state.status = Status::PendingRemote;
180                            None
181                        }
182                    }
183                }
184                None => return Err(anyhow!("Predicate {} not found", query_id)),
185            }
186        };
187
188        match update {
189            Some((peer_id, collection_id, context_data)) => {
190                self.update_query_on_peer(peer_id, query_id, collection_id, selection, version, context_data);
191            }
192            None => {
193                // Not established yet - use setup_remote_subscriptions for initial setup
194                self.setup_remote_subscriptions();
195            }
196        };
197
198        Ok(())
199    }
200
201    fn update_query_on_peer(
202        &self,
203        peer_id: proto::EntityId,
204        query_id: proto::QueryId,
205        collection_id: CollectionId,
206        selection: ankql::ast::Selection,
207        version: u32,
208        context_data: CD,
209    ) {
210        let me = self.clone();
211        crate::task::spawn(async move {
212            if let Some(node) = me.inner.node.get() {
213                // Get the livequery for error handling
214                let livequery = {
215                    me.inner.subscriptions.lock().unwrap_or_else(|e| e.into_inner()).get(&query_id).map(|state| state.livequery.clone())
216                };
217
218                // Send the updated predicate to the peer
219                match node.remote_subscribe(peer_id, query_id, collection_id, selection, &context_data, version).await {
220                    Ok(()) => {
221                        // Deltas applied successfully, now activate the livequery
222                        if let Some(lq) = livequery {
223                            lq.subscription_established(version).await;
224                        }
225
226                        // Mark as established - subscription succeeded even if livequery activation had issues
227                        let mut subscriptions = me.inner.subscriptions.lock().unwrap_or_else(|e| e.into_inner());
228                        if let Some(info) = subscriptions.get_mut(&query_id) {
229                            info.status = Status::Established(peer_id, version);
230                        }
231                        debug!("Successfully updated predicate {} on peer {} subscription", query_id, peer_id);
232                    }
233                    Err(e) => {
234                        // Handle error with retry logic
235                        me.handle_error(query_id, peer_id, e, livequery).await;
236                    }
237                }
238            }
239        });
240    }
241
242    /// Notify the relay that a predicate should be removed from remote peer subscriptions
243    ///
244    /// This will clean up all tracking state and send unsubscribe requests to any
245    /// remote peers that have this predicate registered.
246    pub fn unsubscribe_predicate(&self, query_id: proto::QueryId) {
247        debug!("Unregistering predicate {}", query_id);
248
249        // If subscription was established with a peer, send unsubscribe request
250        {
251            let mut subscriptions = self.inner.subscriptions.lock().unwrap_or_else(|e| e.into_inner());
252            if let Some(info) = subscriptions.remove(&query_id) {
253                if let Status::Established(peer_id, _version) = &info.status {
254                    let node = self.inner.node.get();
255                    if let Some(node) = node {
256                        let node = node.clone();
257                        let peer_id = *peer_id;
258                        crate::task::spawn(async move {
259                            if let Err(e) = node.peer_unsubscribe(peer_id, query_id).await {
260                                warn!("Failed to send unsubscribe message for {}: {}", query_id, e);
261                            } else {
262                                debug!("Successfully sent unsubscribe message for {}", query_id);
263                            }
264                        });
265                    }
266                }
267            }
268        }
269    }
270
271    /// Handle peer disconnection - mark all predicates for that peer as needing re-registration
272    ///
273    /// This should be called when a durable peer disconnects. All predicates registered
274    /// with that peer will be marked as pending and will be automatically re-registered
275    /// when the peer reconnects or another suitable peer becomes available.
276    pub fn notify_peer_disconnected(&self, peer_id: proto::EntityId) {
277        debug!("Peer {} disconnected, orphaning predicate registrations", peer_id);
278
279        // Remove from connected peers
280        self.inner.connected_peers.remove(&peer_id);
281
282        for info in self.inner.subscriptions.lock().expect("poisoned lock").values_mut() {
283            if let Status::Established(established_peer_id, _) | Status::Requested(established_peer_id, _) = &info.status {
284                if *established_peer_id == peer_id {
285                    // Update state to pending
286                    info.status = Status::PendingRemote;
287                    warn!("Predicate {} orphaned due to peer {} disconnect", info.content.query_id, peer_id);
288                }
289            }
290        }
291
292        // Resubscribe any orphaned subscriptions
293        self.setup_remote_subscriptions();
294    }
295
296    /// Handle peer connection - trigger predicate registration on the new peer subscription
297    ///
298    /// This should be called when a new durable peer connects. The relay will automatically
299    /// attempt to register any pending predicates on the newly connected peer's subscription.
300    pub fn notify_peer_connected(&self, peer_id: proto::EntityId) {
301        debug!("SubscriptionRelay.notify_peer_connected() - Peer {} connected, registering predicates on peer subscription", peer_id);
302
303        // Add to connected peers
304        self.inner.connected_peers.insert(peer_id);
305
306        // Trigger setup with all connected peers
307        self.setup_remote_subscriptions();
308    }
309
310    /// Get the current state of a predicate registration
311    pub fn get_status(&self, query_id: proto::QueryId) -> Option<Status> {
312        let subscriptions = self.inner.subscriptions.lock().unwrap_or_else(|e| e.into_inner());
313        subscriptions.get(&query_id).map(|info| info.status.clone())
314    }
315
316    /// Get all unique contexts for predicates established or requested with a specific peer
317    /// TODO: update the data structure to do this via a direct lookup rather than having to scan the entire map
318    pub fn get_contexts_for_peer(&self, peer_id: &proto::EntityId) -> std::collections::HashSet<CD> {
319        let subscriptions = self.inner.subscriptions.lock().unwrap_or_else(|e| e.into_inner());
320        let mut contexts = std::collections::HashSet::new();
321
322        for (_, state) in subscriptions.iter() {
323            match &state.status {
324                Status::Established(established_peer, _) | Status::Requested(established_peer, _) => {
325                    if established_peer == peer_id {
326                        contexts.insert(state.content.context_data.clone());
327                    }
328                }
329                _ => {}
330            }
331        }
332
333        contexts
334    }
335
336    /// Register predicates on available durable peer subscriptions
337    fn setup_remote_subscriptions(&self) {
338        let node = match self.inner.node.get() {
339            Some(node) => node,
340            None => {
341                warn!("No node configured for remote subscription setup");
342                return;
343            }
344        };
345
346        // For now, use the first available peer (could be made smarter)
347        let connected_peers = self.inner.connected_peers.to_vec();
348        if connected_peers.is_empty() {
349            warn!("No durable peers available for remote subscription setup");
350            return;
351        }
352
353        let target_peer = connected_peers[0];
354
355        // Atomically get pending subscriptions and mark them as requested
356        let pending: Vec<_> = {
357            self.inner
358                .subscriptions
359                .lock()
360                .expect("poisoned lock")
361                .values_mut()
362                .filter_map(|info| {
363                    if let Status::PendingRemote = info.status {
364                        info.status = Status::Requested(target_peer, info.content.version);
365                        Some(info.content.clone())
366                    } else {
367                        None
368                    }
369                })
370                .collect()
371        };
372
373        if pending.is_empty() {
374            return;
375        }
376
377        debug!("Registering {} predicates on {} peer subscriptions", pending.len(), self.inner.connected_peers.len());
378
379        for content in pending {
380            crate::task::spawn(self.clone().attempt_subscribe(node.clone(), target_peer, content));
381        }
382    }
383
384    async fn attempt_subscribe(self, node: Arc<dyn TNode<CD>>, target_peer: proto::EntityId, content: Arc<Content<CD>>) {
385        let query_id = content.query_id;
386        let predicate = content.selection.clone();
387        let context_data = content.context_data.clone();
388        let version = content.version;
389
390        // Get the livequery for error handling
391        let livequery =
392            { self.inner.subscriptions.lock().unwrap_or_else(|e| e.into_inner()).get(&query_id).map(|state| state.livequery.clone()) };
393
394        // Call remote_subscribe which fetches known matches, subscribes, applies deltas, and stores events
395        match node.remote_subscribe(target_peer, query_id, content.collection_id.clone(), predicate, &context_data, version).await {
396            Ok(()) => {
397                // Deltas applied successfully, now activate the livequery
398                // The livequery handles its own errors internally
399                if let Some(lq) = livequery {
400                    lq.subscription_established(version).await;
401                }
402
403                // Mark as established - subscription succeeded even if livequery activation had issues
404                let mut subscriptions = self.inner.subscriptions.lock().unwrap_or_else(|e| e.into_inner());
405                if let Some(info) = subscriptions.get_mut(&query_id) {
406                    info.status = Status::Established(target_peer, version);
407                }
408                debug!("Successfully registered predicate {} on peer {} subscription", query_id, target_peer);
409            }
410            Err(e) => {
411                // Handle error with retry logic
412                self.handle_error(query_id, target_peer, e, livequery).await;
413            }
414        }
415    }
416
417    /// Start background task that periodically retries pending subscriptions
418    fn start_retry_task(&self, mut shutdown_rx: tokio::sync::mpsc::Receiver<()>) {
419        let me = self.clone();
420        crate::task::spawn(async move {
421            loop {
422                let delay = futures_timer::Delay::new(std::time::Duration::from_secs(5));
423                tokio::select! {
424                    _ = delay => {
425                        // Attempt to setup any pending subscriptions
426                        me.setup_remote_subscriptions();
427                    }
428                    _ = shutdown_rx.recv() => {
429                        debug!("Retry task shutting down - SubscriptionRelay dropped");
430                        break;
431                    }
432                }
433            }
434        });
435    }
436
437    /// Handle errors with retry logic
438    async fn handle_error(&self, query_id: proto::QueryId, target_peer: proto::EntityId, error: RetrievalError, livequery: Option<Q>) {
439        let error_msg = error.to_string();
440
441        // Evaluate retriability at failure time
442        let is_retryable = match &error {
443            // Retrieval errors from fetching are generally not retryable
444            RetrievalError::RequestError(req_err) => match req_err {
445                RequestError::PeerNotConnected => true,
446                RequestError::ConnectionLost => true,
447                RequestError::SendError(_) => true,
448                RequestError::InternalChannelClosed => true,
449                RequestError::ServerError(_) => false,
450                RequestError::UnexpectedResponse(_) => false,
451                RequestError::AccessDenied(_) => false,
452            },
453            // Other retrieval errors are not retryable
454            _ => false,
455        };
456
457        // Update state based on retriability
458        let mut subscriptions = self.inner.subscriptions.lock().unwrap_or_else(|e| e.into_inner());
459        if let Some(info) = subscriptions.get_mut(&query_id) {
460            if is_retryable {
461                // Retryable errors go back to pending for retry by background task
462                info.status = Status::PendingRemote;
463                warn!("Retryable failure for predicate {} with peer {}: {} - will retry", query_id, target_peer, error_msg);
464            } else {
465                // Non-retryable errors are permanently failed
466                info.status = Status::Failed;
467                tracing::error!("Permanent failure for predicate {} with peer {}: {} - no retry", query_id, target_peer, error_msg);
468
469                // Set error on livequery
470                if let Some(lq) = livequery {
471                    lq.set_last_error(error);
472                }
473            }
474        }
475    }
476}
477
478/// Trait for communicating with remote peers (abstraction over WeakNode for testing)
479#[async_trait]
480pub trait TNode<CD: ContextData>: Send + Sync {
481    /// Send a predicate registration request to a remote peer, fetch known matches,
482    /// apply received deltas, and store used events.
483    /// Returns Ok(()) if subscription was established and deltas applied successfully.
484    async fn remote_subscribe(
485        &self,
486        peer_id: proto::EntityId,
487        query_id: proto::QueryId,
488        collection_id: CollectionId,
489        selection: ankql::ast::Selection,
490        context_data: &CD,
491        version: u32,
492    ) -> Result<(), RetrievalError>;
493
494    /// Send a predicate unregistration message to a remote peer
495    /// This is a one-way message, no response expected
496    async fn peer_unsubscribe(&self, peer_id: proto::EntityId, query_id: proto::QueryId) -> Result<(), anyhow::Error>;
497}
498
499/// Implementation of TNode for WeakNode
500#[async_trait]
501impl<SE, PA> TNode<PA::ContextData> for crate::node::WeakNode<SE, PA>
502where
503    SE: crate::storage::StorageEngine + Send + Sync + 'static,
504    PA: crate::policy::PolicyAgent + Send + Sync + 'static,
505{
506    async fn remote_subscribe(
507        &self,
508        peer_id: proto::EntityId,
509        query_id: proto::QueryId,
510        collection_id: CollectionId,
511        selection: ankql::ast::Selection,
512        context_data: &PA::ContextData,
513        version: u32,
514    ) -> Result<(), RetrievalError> {
515        let node = self.upgrade().ok_or_else(|| RetrievalError::Other("Node has been dropped".to_string()))?;
516
517        // 1. Pre-fetch known_matches from local storage
518        let known_matches: Vec<ankurah_proto::KnownEntity> = node
519            .fetch_entities_from_local(&collection_id, &selection)
520            .await?
521            .into_iter()
522            .map(|entity| ankurah_proto::KnownEntity { entity_id: entity.id(), head: entity.head() })
523            .collect();
524
525        // 2. Send subscribe request with known_matches
526        let deltas = match node
527            .request(
528                peer_id,
529                context_data,
530                ankurah_proto::NodeRequestBody::SubscribeQuery {
531                    query_id,
532                    collection: collection_id.clone(),
533                    selection: selection.clone(),
534                    version,
535                    known_matches,
536                },
537            )
538            .await
539            .map_err(|e| RetrievalError::RequestError(e))?
540        {
541            ankurah_proto::NodeResponseBody::QuerySubscribed { query_id: _response_query_id, deltas } => deltas,
542            ankurah_proto::NodeResponseBody::Error(e) => return Err(RetrievalError::RequestError(RequestError::ServerError(e))),
543            other => return Err(RetrievalError::RequestError(RequestError::UnexpectedResponse(other))),
544        };
545
546        tracing::debug!(
547            "Node.remote_subscribe: query_id: {}, collection_id: {}, received deltas: {}",
548            query_id,
549            collection_id,
550            deltas.len()
551        );
552        // 3. Apply deltas to local node using NodeApplier
553        let collection = node.collections.get(&collection_id).await?;
554        let event_getter = crate::retrieval::CachedEventGetter::new(collection_id, collection.clone(), &node, context_data);
555        let state_getter = crate::retrieval::LocalStateGetter::new(collection);
556        crate::node_applier::NodeApplier::apply_deltas(&node, &peer_id, deltas, &event_getter, &state_getter).await?;
557
558        Ok(())
559    }
560
561    async fn peer_unsubscribe(&self, peer_id: proto::EntityId, query_id: proto::QueryId) -> Result<(), anyhow::Error> {
562        let node = self.upgrade().ok_or_else(|| anyhow!("Node has been dropped"))?;
563
564        // Use the existing request_remote_unsubscribe method
565        node.request_remote_unsubscribe(query_id, vec![peer_id]).await?;
566
567        Ok(())
568    }
569}
570
571#[cfg(test)]
572mod tests {
573    use super::*;
574    use ankql::ast::Predicate;
575    use ankurah_proto::EntityId;
576    use std::sync::{Arc, Mutex};
577
578    // Note: Some tests call setup_remote_subscriptions() directly to test the core
579    // subscription setup logic in isolation, while others use notify_peer_connected()
580    // to test the full event-driven flow. Both approaches are valuable:
581    // - Direct calls test the setup mechanism itself (error handling, state transitions)
582    // - Event-driven calls test the integration and user-facing API
583
584    // For testing, we'll use CollectionId as our ContextData
585    impl ContextData for CollectionId {}
586
587    /// Mock message sender for testing
588    #[derive(Debug)]
589    struct MockMessageSender<CD: ContextData> {
590        next_error: Arc<Mutex<Option<RequestError>>>,
591        sent_requests: Arc<Mutex<Vec<(EntityId, proto::QueryId, CollectionId, ankql::ast::Selection)>>>,
592        should_fail: Arc<Mutex<bool>>,
593        failure_message: Arc<Mutex<String>>,
594        _phantom: std::marker::PhantomData<CD>,
595    }
596
597    impl<CD: ContextData> MockMessageSender<CD> {
598        fn new() -> Self {
599            Self {
600                sent_requests: Arc::new(Mutex::new(Vec::new())),
601                next_error: Arc::new(Mutex::new(None)),
602                should_fail: Arc::new(Mutex::new(false)),
603                failure_message: Arc::new(Mutex::new(String::new())),
604                _phantom: std::marker::PhantomData,
605            }
606        }
607
608        fn set_fail_next(&self, error: RequestError) { *self.next_error.lock().unwrap() = Some(error); }
609
610        fn get_sent_requests(&self) -> Vec<(EntityId, proto::QueryId, CollectionId, ankql::ast::Selection)> {
611            self.sent_requests.lock().unwrap().clone()
612        }
613
614        fn clear_sent_requests(&self) { self.sent_requests.lock().unwrap().clear(); }
615    }
616
617    #[async_trait]
618    impl<CD: ContextData> TNode<CD> for MockMessageSender<CD> {
619        async fn remote_subscribe(
620            &self,
621            peer_id: EntityId,
622            query_id: proto::QueryId,
623            collection_id: CollectionId,
624            selection: ankql::ast::Selection,
625            _context_data: &CD,
626            _version: u32,
627        ) -> Result<(), RetrievalError> {
628            self.sent_requests.lock().unwrap().push((peer_id, query_id, collection_id.clone(), selection.clone()));
629
630            // Check if there's an error to fail with
631            if let Some(error) = self.next_error.lock().unwrap().take() {
632                Err(RetrievalError::RequestError(error))
633            } else {
634                // Mock successful subscription (fetch, subscribe, apply, store all succeeded)
635                Ok(())
636            }
637        }
638
639        async fn peer_unsubscribe(&self, peer_id: EntityId, query_id: proto::QueryId) -> Result<(), anyhow::Error> {
640            self.sent_requests.lock().unwrap().push((
641                peer_id,
642                query_id,
643                CollectionId::from("unsubscribe"),
644                ankql::ast::Selection { predicate: ankql::ast::Predicate::True, order_by: None, limit: None },
645            ));
646
647            // Check if there's an error to fail with
648            if let Some(error) = self.next_error.lock().unwrap().take() {
649                Err(anyhow!(error.to_string()))
650            } else {
651                Ok(())
652            }
653        }
654    }
655
656    // Mock implementation of RemoteQuerySubscriber for tests
657    #[derive(Clone)]
658    struct MockLiveQuery;
659
660    #[async_trait::async_trait]
661    impl RemoteQuerySubscriber for MockLiveQuery {
662        async fn subscription_established(&self, _version: u32) {
663            // Mock - no-op
664        }
665
666        fn set_last_error(&self, _error: RetrievalError) {
667            // For tests, we don't track errors
668        }
669    }
670
671    fn create_test_selection() -> ankql::ast::Selection {
672        // Create a simple test predicate
673        ankql::ast::Selection { predicate: ankql::ast::Predicate::True, order_by: None, limit: None }
674    }
675
676    fn create_test_collection_id() -> CollectionId { CollectionId::from("test_collection") }
677
678    #[tokio::test]
679    async fn test_new_subscription_setup() {
680        let relay = SubscriptionRelay::new();
681        let mock_sender = Arc::new(MockMessageSender::<CollectionId>::new());
682        relay.set_node(mock_sender.clone()).expect("Failed to set message sender");
683
684        let query_id = proto::QueryId::new();
685        let collection_id = create_test_collection_id();
686        let predicate = create_test_selection();
687        let peer_id = EntityId::new();
688
689        // Connect the peer first
690        relay.notify_peer_connected(peer_id);
691
692        // Notify of new subscription
693        relay.subscribe_query(query_id, collection_id.clone(), predicate.clone(), collection_id.clone(), 0, MockLiveQuery);
694
695        // Check initial state - subscription should immediately go to Requested state since peer is connected
696        assert!(matches!(relay.get_status(query_id), Some(Status::Requested(_, _))));
697
698        // Give async task time to complete (setup should happen automatically)
699        futures_timer::Delay::new(std::time::Duration::from_millis(10)).await;
700
701        // Verify request was sent
702        let sent_requests = mock_sender.get_sent_requests();
703        assert_eq!(sent_requests.len(), 1);
704        assert_eq!(sent_requests[0].0, peer_id);
705        assert_eq!(sent_requests[0].1, query_id);
706        assert_eq!(sent_requests[0].2, collection_id);
707
708        // Verify subscription is marked as established
709        assert!(matches!(relay.get_status(query_id), Some(Status::Established(established_peer_id, _)) if established_peer_id == peer_id));
710    }
711
712    #[tokio::test]
713    async fn test_peer_disconnection_orphans_subscriptions() {
714        let relay = SubscriptionRelay::new();
715
716        let mock_sender = Arc::new(MockMessageSender::<CollectionId>::new());
717        relay.set_node(mock_sender.clone()).expect("Failed to set message sender");
718
719        let query_id = proto::QueryId::new();
720        let collection_id = create_test_collection_id();
721        let predicate = create_test_selection();
722        let peer_id = EntityId::new();
723
724        // Connect the peer first
725        relay.notify_peer_connected(peer_id);
726
727        // Setup established subscription by going through the full flow
728        relay.subscribe_query(query_id, collection_id.clone(), predicate, collection_id.clone(), 0, MockLiveQuery);
729
730        // Give async task time to complete
731        futures_timer::Delay::new(std::time::Duration::from_millis(10)).await;
732
733        assert!(matches!(relay.get_status(query_id), Some(Status::Established(established_peer_id, _)) if established_peer_id == peer_id));
734
735        // Simulate peer disconnection
736        relay.notify_peer_disconnected(peer_id);
737
738        // Verify subscription is marked as pending again
739        assert!(matches!(relay.get_status(query_id), Some(Status::PendingRemote)));
740    }
741
742    #[tokio::test]
743    async fn test_peer_connection_triggers_setup() {
744        let relay = SubscriptionRelay::new();
745        let mock_sender = Arc::new(MockMessageSender::<CollectionId>::new());
746        relay.set_node(mock_sender.clone()).expect("Failed to set message sender");
747
748        let query_id = proto::QueryId::new();
749        let collection_id = create_test_collection_id();
750        let predicate = create_test_selection();
751        let peer_id = EntityId::new();
752
753        // Add pending subscription (no peers connected yet)
754        relay.subscribe_query(query_id, collection_id.clone(), predicate.clone(), collection_id.clone(), 0, MockLiveQuery);
755        assert!(matches!(relay.get_status(query_id), Some(Status::PendingRemote)));
756
757        // Clear any previous requests
758        mock_sender.clear_sent_requests();
759
760        // Simulate peer connection (should trigger automatic setup)
761        relay.notify_peer_connected(peer_id);
762
763        // Give async task time to complete
764        futures_timer::Delay::new(std::time::Duration::from_millis(10)).await;
765
766        // Verify request was sent
767        let sent_requests = mock_sender.get_sent_requests();
768        assert_eq!(sent_requests.len(), 1);
769        assert_eq!(sent_requests[0].0, peer_id);
770        assert_eq!(sent_requests[0].1, query_id);
771
772        // Verify subscription is established
773        assert!(matches!(relay.get_status(query_id), Some(Status::Established(established_peer_id, _)) if established_peer_id == peer_id));
774    }
775
776    #[tokio::test]
777    async fn test_failed_subscription_retry() {
778        let relay = SubscriptionRelay::new();
779        let mock_sender = Arc::new(MockMessageSender::<CollectionId>::new());
780        relay.set_node(mock_sender.clone()).expect("Failed to set message sender");
781
782        let query_id = proto::QueryId::new();
783        let collection_id = create_test_collection_id();
784        let predicate = create_test_selection();
785        let peer_id = EntityId::new();
786
787        // Connect peer and add subscription (should succeed initially)
788        relay.notify_peer_connected(peer_id);
789        relay.subscribe_query(query_id, collection_id.clone(), predicate.clone(), collection_id.clone(), 0, MockLiveQuery);
790
791        // Give async task time to complete
792        futures_timer::Delay::new(std::time::Duration::from_millis(10)).await;
793
794        // Verify subscription is marked as established (since no error was set)
795        assert!(matches!(relay.get_status(query_id), Some(Status::Established(established_peer_id, _)) if established_peer_id == peer_id));
796
797        // Now test the retry behavior by disconnecting the peer (puts subscription back to PendingRemote)
798        // then setting up the mock to fail, and reconnecting to trigger the retry
799        relay.notify_peer_disconnected(peer_id);
800
801        // Verify subscription is now in pending state
802        assert!(matches!(relay.get_status(query_id), Some(Status::PendingRemote)));
803
804        // Clear requests and set up mock to fail on the next call
805        mock_sender.clear_sent_requests();
806        mock_sender.set_fail_next(RequestError::ServerError("Invalid predicate".to_string()));
807
808        // Reconnect peer to trigger retry attempt
809        relay.notify_peer_connected(peer_id);
810
811        // Give async task time to complete
812        futures_timer::Delay::new(std::time::Duration::from_millis(10)).await;
813
814        // Verify retry was attempted (the error gets consumed)
815        let sent_requests = mock_sender.get_sent_requests();
816        assert_eq!(sent_requests.len(), 1);
817
818        // Verify subscription remains in failed state (non-retryable error)
819        assert!(matches!(relay.get_status(query_id), Some(Status::Failed)));
820    }
821
822    #[tokio::test]
823    async fn test_retryable_vs_non_retryable_failures() {
824        let relay = SubscriptionRelay::new();
825        let mock_sender = Arc::new(MockMessageSender::<CollectionId>::new());
826        relay.set_node(mock_sender.clone()).expect("Failed to set message sender");
827
828        let retryable_query_id = proto::QueryId::new();
829        let non_retryable_query_id = proto::QueryId::new();
830        let collection_id = create_test_collection_id();
831        let predicate = create_test_selection();
832        let peer_id = EntityId::new();
833
834        // Add subscriptions
835        relay.subscribe_query(retryable_query_id, collection_id.clone(), predicate.clone(), collection_id.clone(), 0, MockLiveQuery);
836        relay.subscribe_query(non_retryable_query_id, collection_id.clone(), predicate.clone(), collection_id.clone(), 0, MockLiveQuery);
837
838        // Manually set different failure types - retryable goes back to pending, non-retryable stays failed
839        {
840            let mut subscriptions = relay.inner.subscriptions.lock().unwrap_or_else(|e| e.into_inner());
841            if let Some(info) = subscriptions.get_mut(&retryable_query_id) {
842                info.status = Status::PendingRemote; // Retryable errors go back to pending
843            }
844            if let Some(info) = subscriptions.get_mut(&non_retryable_query_id) {
845                info.status = Status::Failed; // Non-retryable errors stay failed
846            }
847        }
848
849        // Connect peer and trigger retry
850        relay.notify_peer_connected(peer_id);
851
852        // Give async task time to complete
853        futures_timer::Delay::new(std::time::Duration::from_millis(10)).await;
854
855        // Verify only the retryable subscription was attempted
856        let sent_requests = mock_sender.get_sent_requests();
857        assert_eq!(sent_requests.len(), 1);
858        assert_eq!(sent_requests[0].1, retryable_query_id);
859
860        // Verify states
861        assert!(
862            matches!(relay.get_status(retryable_query_id), Some(Status::Established(established_peer_id, _)) if established_peer_id == peer_id)
863        );
864        assert!(matches!(relay.get_status(non_retryable_query_id), Some(Status::Failed)));
865    }
866
867    #[tokio::test]
868    async fn test_subscription_removal() {
869        let relay = SubscriptionRelay::new();
870        let mock_sender = Arc::new(MockMessageSender::<CollectionId>::new());
871        relay.set_node(mock_sender.clone()).expect("Failed to set message sender");
872
873        let query_id = proto::QueryId::new();
874        let collection_id = create_test_collection_id();
875        let predicate = create_test_selection();
876        let peer_id = EntityId::new();
877
878        // Connect peer and setup established subscription
879        relay.notify_peer_connected(peer_id);
880        relay.subscribe_query(query_id, collection_id.clone(), predicate, collection_id.clone(), 0, MockLiveQuery);
881
882        // Give async task time to complete
883        futures_timer::Delay::new(std::time::Duration::from_millis(10)).await;
884
885        assert!(matches!(relay.get_status(query_id), Some(Status::Established(established_peer_id, _)) if established_peer_id == peer_id));
886
887        // Clear previous requests to focus on unsubscribe
888        mock_sender.clear_sent_requests();
889
890        // Remove subscription
891        relay.unsubscribe_predicate(query_id);
892
893        // Give async task time to complete
894        futures_timer::Delay::new(std::time::Duration::from_millis(10)).await;
895
896        // Verify unsubscribe message was sent
897        let sent_requests = mock_sender.get_sent_requests();
898        assert_eq!(sent_requests.len(), 1);
899        assert_eq!(sent_requests[0].0, peer_id);
900        assert_eq!(sent_requests[0].1, query_id);
901
902        // Verify subscription is gone
903        assert!(matches!(relay.get_status(query_id), None));
904    }
905
906    #[tokio::test]
907    async fn test_edge_cases() {
908        let relay = SubscriptionRelay::new();
909        let mock_sender = Arc::new(MockMessageSender::<CollectionId>::new());
910
911        let query_id = proto::QueryId::new();
912        let collection_id = create_test_collection_id();
913        let predicate = create_test_selection();
914        let peer_id = EntityId::new();
915
916        // Test setup without message sender - should not crash
917        relay.subscribe_query(query_id, collection_id.clone(), predicate.clone(), collection_id.clone(), 0, MockLiveQuery);
918        futures_timer::Delay::new(std::time::Duration::from_millis(10)).await;
919
920        // Should still be pending since no sender
921        assert!(matches!(relay.get_status(query_id), Some(Status::PendingRemote)));
922
923        // Now set sender and test with no connected peers
924        relay.set_node(mock_sender.clone()).expect("Failed to set message sender");
925        futures_timer::Delay::new(std::time::Duration::from_millis(10)).await;
926
927        // Should still be pending since no peers available
928        assert!(matches!(relay.get_status(query_id), Some(Status::PendingRemote)));
929
930        // Verify no requests were sent
931        assert_eq!(mock_sender.get_sent_requests().len(), 0);
932
933        // Now connect a peer (should trigger automatic setup)
934        relay.notify_peer_connected(peer_id);
935        futures_timer::Delay::new(std::time::Duration::from_millis(10)).await;
936
937        // Should now be established
938        assert!(matches!(relay.get_status(query_id), Some(Status::Established(established_peer_id, _)) if established_peer_id == peer_id));
939        assert_eq!(mock_sender.get_sent_requests().len(), 1);
940    }
941
942    #[tokio::test]
943    async fn test_notify_unsubscribe_with_no_established_subscription() {
944        let relay = SubscriptionRelay::new();
945        let mock_sender = Arc::new(MockMessageSender::<CollectionId>::new());
946        relay.set_node(mock_sender.clone()).expect("Failed to set message sender");
947
948        let query_id = proto::QueryId::new();
949        let collection_id = create_test_collection_id();
950        let predicate = create_test_selection();
951
952        // Add subscription but don't establish it
953        relay.subscribe_query(query_id, collection_id.clone(), predicate, collection_id.clone(), 0, MockLiveQuery);
954        assert!(matches!(relay.get_status(query_id), Some(Status::PendingRemote)));
955
956        // Unsubscribe from pending subscription
957        relay.unsubscribe_predicate(query_id);
958
959        // Give async task time to complete (though no request should be sent)
960        futures_timer::Delay::new(std::time::Duration::from_millis(10)).await;
961
962        // Verify no unsubscribe message was sent (since it wasn't established)
963        let sent_requests = mock_sender.get_sent_requests();
964        assert_eq!(sent_requests.len(), 0);
965
966        // Verify subscription is gone
967        assert!(matches!(relay.get_status(query_id), None));
968    }
969}