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
14pub const DEFAULT_CHANNEL_SIZE: usize = 1024;
16
17type SubscriptionRequest = (
19 Box<dyn TraceSinkHandle>,
20 oneshot::Sender<Result<Vec<IpcMessageWithId>>>,
21);
22
23pub struct TraceRouter {
25 sender: Sender,
27
28 subscription_sender: flume::Sender<SubscriptionRequest>,
30}
31
32impl TraceRouter {
33 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 pub fn new_with_store(
43 store: Arc<dyn Store>,
44 cancellation_token: CancellationToken,
45 ) -> (Arc<Self>, impl Future<Output = Result<()>>) {
46 let (sender, receiver) = flume::bounded(DEFAULT_CHANNEL_SIZE);
48
49 let (subscription_sender, subscription_receiver) = flume::bounded(1);
51
52 let router = TraceRouter {
53 sender,
54 subscription_sender,
55 };
56
57 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 if let Err(e) = store.update(&msg) {
70 tracing::error!("Error while updating the store: {}", e);
71 }
72
73 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 closed_sinks.push(idx);
84 }
85 }
86 }
87
88 if !closed_sinks.is_empty() {
90 closed_sinks.sort_unstable_by(|a, b| b.cmp(a));
93
94 for idx in closed_sinks {
95 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 let mut sinks = Vec::new();
122
123 loop {
124 tokio::select! {
125 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 metrics::counter!("messages_received", "task" => "router").increment(1);
142 metrics::gauge!("receiver_len", "task" => "router").set(receiver.len() as f64);
143
144 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 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 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 let metadata = sub_response_receiver
189 .await
190 .map_err(|_| anyhow::anyhow!("Response channel closed"))??;
191
192 Ok((receiver, metadata))
193 }
194
195 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 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}