Skip to main content

zelos_trace/
router.rs

1use std::{future::Future, sync::Arc};
2
3use anyhow::Result;
4use tokio::{sync::oneshot, time::Instant};
5use tokio_stream::{Stream, StreamExt};
6use tokio_util::sync::CancellationToken;
7use zelos_trace_types::ipc::{IpcMessageWithId, Receiver, Sender};
8
9use crate::{
10    MetadataOnlyStore, Store, TraceSink,
11    sink::{TraceSinkHandle, TraceSinkHandleAllBlocking},
12};
13
14// TODO(tkeairns): Ground this constant into some relationship with # msgs/sec
15pub const DEFAULT_CHANNEL_SIZE: usize = 1024;
16
17/// Subscription requests sent to the router's main task
18type SubscriptionRequest = (
19    Box<dyn TraceSinkHandle>,
20    oneshot::Sender<Result<Vec<IpcMessageWithId>>>,
21);
22
23/// Pub-sub router for trace data
24pub struct TraceRouter {
25    // Channel for broadcasting trace streams to subscribers
26    sender: Sender,
27
28    // Channel for subscription requests
29    subscription_sender: flume::Sender<SubscriptionRequest>,
30}
31
32impl TraceRouter {
33    /// Create a new trace router with the default metadata-only store
34    pub fn new(
35        cancellation_token: CancellationToken,
36    ) -> (Arc<Self>, impl Future<Output = Result<()>>) {
37        let store = Arc::new(MetadataOnlyStore::new());
38        Self::new_with_store(store, cancellation_token)
39    }
40
41    /// Create a new trace router with a specific store implementation
42    pub fn new_with_store(
43        store: Arc<dyn Store>,
44        cancellation_token: CancellationToken,
45    ) -> (Arc<Self>, impl Future<Output = Result<()>>) {
46        // Initialize the channel for receiving trace streams.
47        let (sender, receiver) = flume::bounded(DEFAULT_CHANNEL_SIZE);
48
49        // Initialize the channel for subscription requests
50        let (subscription_sender, subscription_receiver) = flume::bounded(1);
51
52        let router = TraceRouter {
53            sender,
54            subscription_sender,
55        };
56
57        // Spawn the router's main task
58        let run = TraceRouter::run(receiver, subscription_receiver, store, cancellation_token);
59
60        (Arc::new(router), run)
61    }
62
63    async fn forward_message(
64        store: &Arc<dyn Store>,
65        sinks: &mut Vec<Box<dyn TraceSinkHandle>>,
66        msg: IpcMessageWithId,
67    ) {
68        // Update the store
69        if let Err(e) = store.update(&msg) {
70            tracing::error!("Error while updating the store: {}", e);
71        }
72
73        // Forward this message to all subscribers
74        let mut closed_sinks = Vec::new();
75        {
76            metrics::gauge!("router_sinks", "task" => "router").set(sinks.len() as f64);
77
78            for (idx, sink) in sinks.iter().enumerate() {
79                if let Err(e) = sink.send_async(&msg).await {
80                    tracing::trace!("Error when sending on sink: {}", e);
81                    // If we have an error here, this means that the sink is no longer
82                    // available, so we add it to the list of sinks to remove
83                    closed_sinks.push(idx);
84                }
85            }
86        }
87
88        // Remove all closed sinks
89        if !closed_sinks.is_empty() {
90            // Sort in reverse order so we can remove from highest index to lowest
91            // without affecting the validity of the remaining indices
92            closed_sinks.sort_unstable_by(|a, b| b.cmp(a));
93
94            for idx in closed_sinks {
95                // Remove the sink at the index
96                sinks.remove(idx);
97            }
98        }
99    }
100
101    async fn handle_subscribe(
102        store: &Arc<dyn Store>,
103        sinks: &mut Vec<Box<dyn TraceSinkHandle>>,
104        handle: Box<dyn TraceSinkHandle>,
105        sub_response_sender: oneshot::Sender<Result<Vec<IpcMessageWithId>>>,
106    ) {
107        sinks.push(handle);
108
109        if let Err(e) = sub_response_sender.send(store.metadata_as_ipc()) {
110            tracing::error!("Failed to send metadata to new subscriber: {:?}", e);
111        }
112    }
113
114    async fn run(
115        receiver: Receiver,
116        subscription_receiver: flume::Receiver<SubscriptionRequest>,
117        store: Arc<dyn Store>,
118        cancellation_token: CancellationToken,
119    ) -> Result<()> {
120        // Construct task-local state
121        let mut sinks = Vec::new();
122
123        loop {
124            tokio::select! {
125                // Handle subscription requests
126                sub_req = subscription_receiver.recv_async() => {
127                    match sub_req {
128                        Ok((handle, sub_response_sender)) => {
129                            Self::handle_subscribe(&store, &mut sinks, handle, sub_response_sender).await;
130                        }
131                        Err(_) => {
132                            break;
133                        }
134                    }
135                }
136
137                msg = receiver.recv_async() => {
138                    let msg = msg?;
139
140                    // Update our metrics
141                    metrics::counter!("messages_received", "task" => "router").increment(1);
142                    metrics::gauge!("receiver_len", "task" => "router").set(receiver.len() as f64);
143
144                    // Update our state and forward
145                    let start = Instant::now();
146                    TraceRouter::forward_message(&store, &mut sinks, msg).await;
147                    let elapsed = start.elapsed();
148
149                    metrics::histogram!("update_store_duration_ns", "task" => "router")
150                        .record(elapsed.as_nanos() as f64);
151                }
152                _ = cancellation_token.cancelled() => {
153                    tracing::debug!("Shutting down...");
154
155                    // Drain the receiver
156                    let start = Instant::now();
157                    let mut count: usize = 0;
158                    for msg in receiver.drain() {
159                        TraceRouter::forward_message(&store, &mut sinks, msg).await;
160                        count += 1;
161                    }
162                    let elapsed = start.elapsed();
163
164                    tracing::debug!("Shut down complete, took {:?} processed {} messages", elapsed, count);
165                    return Ok(());
166                }
167            }
168        }
169
170        Ok(())
171    }
172
173    pub fn sender(&self) -> Sender {
174        self.sender.clone()
175    }
176
177    /// Subscribe to all data, applying backpressure when needed
178    pub async fn subscribe_all_blocking(&self) -> Result<(Receiver, Vec<IpcMessageWithId>)> {
179        let (handle, receiver) = TraceSinkHandleAllBlocking::new();
180        let (sub_response_sender, sub_response_receiver) = oneshot::channel();
181
182        self.subscription_sender
183            .send_async((Box::new(handle), sub_response_sender))
184            .await
185            .map_err(|_| anyhow::anyhow!("Router subscription channel closed"))?;
186
187        // Wait for the router to return metadata
188        let metadata = sub_response_receiver
189            .await
190            .map_err(|_| anyhow::anyhow!("Response channel closed"))??;
191
192        Ok((receiver, metadata))
193    }
194
195    /// Subscribe to trace streams
196    pub async fn subscribe(&self) -> Result<(TraceSink, Receiver, Vec<IpcMessageWithId>)> {
197        let (sink, receiver, handle) = TraceSink::new();
198        let (sub_response_sender, sub_response_receiver) = oneshot::channel();
199
200        self.subscription_sender
201            .send_async((Box::new(handle), sub_response_sender))
202            .await
203            .map_err(|_| anyhow::anyhow!("Router subscription channel closed"))?;
204
205        // Wait for the router to return metadata
206        let metadata = sub_response_receiver
207            .await
208            .map_err(|_| anyhow::anyhow!("Response channel closed"))??;
209
210        Ok((sink, receiver, metadata))
211    }
212
213    pub async fn subscribe_all_blocking_stream(
214        &self,
215    ) -> Result<impl Stream<Item = IpcMessageWithId> + use<>> {
216        let (receiver, metadata) = self.subscribe_all_blocking().await?;
217        Ok(tokio_stream::iter(metadata).chain(receiver.into_stream()))
218    }
219
220    pub async fn subscribe_stream(
221        &self,
222    ) -> Result<(TraceSink, impl Stream<Item = IpcMessageWithId> + use<>)> {
223        let (sink, receiver, metadata) = self.subscribe().await?;
224        let stream = tokio_stream::iter(metadata).chain(receiver.into_stream());
225        Ok((sink, stream))
226    }
227}