Skip to main content

basilisk_rust_client/
bus_client.rs

1use crate::error::{ClientError, ClientResult};
2use crate::protocol::{
3    ServiceBusEventEnvelope, ServiceBusForwardRequest, ServiceBusForwardResponse,
4    ServiceBusProtocolMessage, protocol_types,
5};
6use chrono::Utc;
7use futures_util::{FutureExt, future::BoxFuture};
8use std::collections::{HashMap, VecDeque};
9use std::sync::Arc;
10use std::time::{SystemTime, UNIX_EPOCH};
11use tokio::io::{AsyncBufReadExt, AsyncWriteExt, BufReader};
12use tokio::net::TcpStream;
13use tokio::net::tcp::{OwnedReadHalf, OwnedWriteHalf};
14use tokio::sync::{Mutex, RwLock, oneshot};
15use tokio::time::{Duration, timeout};
16use tracing::{debug, error, info, warn};
17
18const COMMAND_RESPONSE_TIMEOUT: Duration = Duration::from_secs(10);
19const METRICS_INTERVAL: Duration = Duration::from_secs(20);
20const CONNECT_RETRY_BASE_DELAY: Duration = Duration::from_millis(500);
21const CONNECT_RETRY_MAX_DELAY: Duration = Duration::from_secs(30);
22const CONNECT_RETRY_MAX_JITTER_MS: u64 = 500;
23const CONNECT_RETRY_MAX_ATTEMPTS: usize = 8;
24const METRICS_TOPIC: &str = "basilisk.metrics.distribution";
25const METRICS_MESSAGE_TYPE: &str = "basilisk.internal";
26
27/// Async event-handler function signature used by `BusClient::on_event`.
28pub type EventHandler =
29    Arc<dyn Fn(ServiceBusEventEnvelope) -> BoxFuture<'static, ()> + Send + Sync>;
30/// Async request-handler function signature used by `BusClient::on_request`.
31pub type RequestHandler = Arc<
32    dyn Fn(ServiceBusRequest, RequestResponder) -> BoxFuture<'static, ClientResult<()>>
33        + Send
34        + Sync,
35>;
36
37/// Low-level TCP service-bus client.
38#[derive(Clone)]
39pub struct BusClient {
40    inner: Arc<Inner>,
41}
42
43struct Inner {
44    service_id: String,
45    instance_id: String,
46    writer: Mutex<OwnedWriteHalf>,
47    pending: Mutex<VecDeque<oneshot::Sender<ServiceBusProtocolMessage>>>,
48    event_handlers: RwLock<HashMap<String, Vec<EventHandler>>>,
49    request_handlers: RwLock<HashMap<String, RequestHandler>>,
50}
51
52#[derive(Debug, Clone)]
53/// Wrapper around an incoming request-style event.
54pub struct ServiceBusRequest {
55    /// Incoming event envelope.
56    pub event: ServiceBusEventEnvelope,
57}
58
59impl ServiceBusRequest {
60    /// Returns the reply topic from the payload field ` reply_to `, if present.
61    pub fn reply_to(&self) -> Option<&str> {
62        self.event.payload.get("reply_to")?.as_str()
63    }
64}
65
66/// Helper used by request handlers to publish replies.
67#[derive(Clone)]
68pub struct RequestResponder {
69    client: BusClient,
70    reply_to_topic: String,
71    causation_id: String,
72    correlation_id: i64,
73    default_message_type: String,
74}
75
76impl RequestResponder {
77    /// Sends a typed reply event to the request's reply topic.
78    pub async fn respond(
79        &self,
80        message_type: impl Into<String>,
81        payload: HashMap<String, serde_json::Value>,
82    ) -> ClientResult<i32> {
83        let event = ServiceBusEventEnvelope {
84            event_id: String::new(),
85            emitted_at_utc: Utc::now(),
86            service_id: String::new(),
87            instance_id: String::new(),
88            topic: self.reply_to_topic.clone(),
89            message_type: message_type.into(),
90            correlation_id: self.correlation_id,
91            causation_id: Some(self.causation_id.clone()),
92            payload,
93        };
94        self.client.publish_event(event).await
95    }
96
97    /// Sends a reply using the request's original message type.
98    pub async fn respond_ok(
99        &self,
100        payload: HashMap<String, serde_json::Value>,
101    ) -> ClientResult<i32> {
102        self.respond(self.default_message_type.clone(), payload)
103            .await
104    }
105}
106
107#[derive(Debug, Clone)]
108/// Input payload for `BusClient::forward`.
109pub struct ForwardRequest {
110    /// Target service id that should process the request.
111    pub target_service_id: String,
112    /// Message type to execute at the target.
113    pub message_type: String,
114    /// Request payload.
115    pub payload: HashMap<String, serde_json::Value>,
116    /// Optional timeout in milliseconds.
117    pub timeout_ms: Option<u64>,
118}
119
120impl BusClient {
121    /// Opens a TCP connection and authenticates with a `connect` protocol frame.
122    pub async fn connect(
123        host: &str,
124        port: u16,
125        service_id: impl Into<String>,
126        instance_id: impl Into<String>,
127        token: impl Into<String>,
128    ) -> ClientResult<Self> {
129        let service_id = service_id.into();
130        let instance_id = instance_id.into();
131        let token = token.into();
132        let connection_key = format!("{}:{}", service_id, instance_id);
133
134        let mut attempt = 0usize;
135        loop {
136            let stream = match TcpStream::connect((host, port)).await {
137                Ok(stream) => stream,
138                Err(err) => {
139                    warn!(%connection_key, %err, "TCP connection failed");
140                    if attempt >= CONNECT_RETRY_MAX_ATTEMPTS - 1 {
141                        error!(%connection_key, "Giving up after connection retries");
142                        return Err(err.into());
143                    }
144                    let delay = connect_retry_delay(attempt);
145                    info!(%connection_key, ?delay, "Connection retry scheduled");
146                    tokio::time::sleep(delay).await;
147                    attempt += 1;
148                    continue;
149                }
150            };
151
152            let _ = stream.set_nodelay(true);
153            let (reader, writer) = stream.into_split();
154
155            let inner = Arc::new(Inner {
156                service_id: service_id.clone(),
157                instance_id: instance_id.clone(),
158                writer: Mutex::new(writer),
159                pending: Mutex::new(VecDeque::new()),
160                event_handlers: RwLock::new(HashMap::new()),
161                request_handlers: RwLock::new(HashMap::new()),
162            });
163
164            tokio::spawn(read_loop(Arc::clone(&inner), reader));
165
166            info!(%connection_key, "TCP socket established; sending connect handshake");
167
168            let client = Self {
169                inner: Arc::clone(&inner),
170            };
171
172            let connect_result = client
173                .send_command(ServiceBusProtocolMessage {
174                    r#type: protocol_types::CONNECT.to_string(),
175                    service_id: Some(service_id.clone()),
176                    instance_id: Some(instance_id.clone()),
177                    token: Some(token.clone()),
178                    ..Default::default()
179                })
180                .await;
181
182            if let Err(err) = connect_result {
183                warn!(%connection_key, %err, "Connect handshake failed");
184                if attempt >= CONNECT_RETRY_MAX_ATTEMPTS - 1 {
185                    return Err(err);
186                }
187                let delay = connect_retry_delay(attempt);
188                info!(%connection_key, ?delay, "Connection retry scheduled");
189                tokio::time::sleep(delay).await;
190                attempt += 1;
191                continue;
192            }
193
194            info!(%connection_key, "Authenticated and connected");
195
196            start_metrics_publisher(client.clone());
197            return Ok(client);
198        }
199    }
200
201    /// Subscribes to one or more topics and waits for an `ack`.
202    pub async fn subscribe(&self, topics: Vec<String>) -> ClientResult<()> {
203        self.send_command(ServiceBusProtocolMessage {
204            r#type: protocol_types::SUBSCRIBE.to_string(),
205            topics: Some(topics),
206            ..Default::default()
207        })
208        .await
209        .map(|_| ())
210    }
211
212    /// Unsubscribes from topics without waiting for a response frame.
213    pub async fn unsubscribe(&self, topics: Vec<String>) -> ClientResult<()> {
214        self.send_fire_and_forget(ServiceBusProtocolMessage {
215            r#type: protocol_types::UNSUBSCRIBE.to_string(),
216            topics: Some(topics),
217            ..Default::default()
218        })
219        .await
220    }
221
222    /// Publishes a topic + message-type payload and returns subscriber count.
223    pub async fn publish(
224        &self,
225        topic: impl Into<String>,
226        message_type: impl Into<String>,
227        payload: HashMap<String, serde_json::Value>,
228    ) -> ClientResult<i32> {
229        let event = ServiceBusEventEnvelope {
230            event_id: String::new(),
231            emitted_at_utc: Utc::now(),
232            service_id: String::new(),
233            instance_id: String::new(),
234            topic: topic.into(),
235            message_type: message_type.into(),
236            correlation_id: 0,
237            causation_id: None,
238            payload,
239        };
240
241        self.publish_event(event).await
242    }
243
244    /// Publishes a fully formed event envelope and returns the subscriber count.
245    pub async fn publish_event(&self, event: ServiceBusEventEnvelope) -> ClientResult<i32> {
246        let msg = self
247            .send_command(ServiceBusProtocolMessage {
248                r#type: protocol_types::PUBLISH.to_string(),
249                event: Some(event),
250                ..Default::default()
251            })
252            .await?;
253
254        Ok(msg.subscriber_count.unwrap_or_default())
255    }
256
257    /// Sends a forward request and validates a `forward_response` frame.
258    pub async fn forward(
259        &self,
260        request: ForwardRequest,
261    ) -> ClientResult<ServiceBusForwardResponse> {
262        let response = self
263            .send_command_expect(ServiceBusProtocolMessage {
264                r#type: protocol_types::FORWARD.to_string(),
265                forward_request: Some(ServiceBusForwardRequest {
266                    target_service_id: request.target_service_id,
267                    message_type: request.message_type,
268                    payload: request.payload,
269                    timeout_ms: request.timeout_ms,
270                }),
271                ..Default::default()
272            })
273            .await?;
274
275        if response.r#type != protocol_types::FORWARD_RESPONSE {
276            return Err(ClientError::UnexpectedMessage(response.r#type));
277        }
278
279        response
280            .forward_response
281            .ok_or(ClientError::MissingField("forwardResponse"))
282    }
283
284    /// Registers an async event handler for a topic and subscribes automatically.
285    pub async fn on_event<F, Fut>(&self, topic: impl Into<String>, handler: F) -> ClientResult<()>
286    where
287        F: Fn(ServiceBusEventEnvelope) -> Fut + Send + Sync + 'static,
288        Fut: Future<Output = ()> + Send + 'static,
289    {
290        let topic = topic.into();
291        self.subscribe(vec![topic.clone()]).await?;
292
293        let boxed: EventHandler = Arc::new(move |event| handler(event).boxed());
294        let mut guard = self.inner.event_handlers.write().await;
295        guard.entry(topic).or_default().push(boxed);
296        Ok(())
297    }
298
299    /// Registers an async request responder by message type.
300    pub async fn on_request<F, Fut>(
301        &self,
302        topic: impl Into<String>,
303        responder: F,
304    ) -> ClientResult<()>
305    where
306        F: Fn(ServiceBusRequest, RequestResponder) -> Fut + Send + Sync + 'static,
307        Fut: Future<Output = ClientResult<()>> + Send + 'static,
308    {
309        let topic = topic.into();
310        let service_topic = format!("service-{}", self.inner.service_id);
311        self.subscribe(vec![service_topic]).await?;
312
313        let handler: RequestHandler = Arc::new(move |req, resp| responder(req, resp).boxed());
314        let mut guard = self.inner.request_handlers.write().await;
315        guard.insert(topic, handler);
316        Ok(())
317    }
318
319    async fn send_command(
320        &self,
321        msg: ServiceBusProtocolMessage,
322    ) -> ClientResult<ServiceBusProtocolMessage> {
323        let response = self.send_command_expect(msg).await?;
324        if response.r#type != protocol_types::ACK {
325            return Err(ClientError::UnexpectedMessage(response.r#type));
326        }
327        Ok(response)
328    }
329
330    async fn send_command_expect(
331        &self,
332        msg: ServiceBusProtocolMessage,
333    ) -> ClientResult<ServiceBusProtocolMessage> {
334        let (tx, rx) = oneshot::channel();
335        {
336            let mut pending = self.inner.pending.lock().await;
337            pending.push_back(tx);
338        }
339
340        let mut wire = serde_json::to_string(&msg)?;
341        wire.push('\n');
342        let mut writer = self.inner.writer.lock().await;
343        if let Err(err) = writer.write_all(wire.as_bytes()).await {
344            let mut pending = self.inner.pending.lock().await;
345            let _ = pending.pop_back();
346            return Err(err.into());
347        }
348        writer.flush().await?;
349
350        let response = timeout(COMMAND_RESPONSE_TIMEOUT, rx)
351            .await
352            .map_err(|_| ClientError::Protocol {
353                code: "COMMAND_TIMEOUT".to_string(),
354                message: "Timed out waiting for protocol response".to_string(),
355            })?
356            .map_err(|_| ClientError::ChannelClosed)?;
357        if response.r#type == protocol_types::ERROR {
358            return Err(ClientError::Protocol {
359                code: response
360                    .error_code
361                    .unwrap_or_else(|| "UNKNOWN_ERROR".to_string()),
362                message: response
363                    .message
364                    .unwrap_or_else(|| "Service bus protocol error".to_string()),
365            });
366        }
367
368        Ok(response)
369    }
370
371    async fn send_fire_and_forget(&self, msg: ServiceBusProtocolMessage) -> ClientResult<()> {
372        let mut wire = serde_json::to_string(&msg)?;
373        wire.push('\n');
374        let mut writer = self.inner.writer.lock().await;
375        writer.write_all(wire.as_bytes()).await?;
376        writer.flush().await?;
377        Ok(())
378    }
379}
380
381fn start_metrics_publisher(client: BusClient) {
382    let weak_inner = Arc::downgrade(&client.inner);
383    tokio::spawn(async move {
384        let mut ticker = tokio::time::interval(METRICS_INTERVAL);
385        loop {
386            ticker.tick().await;
387            let Some(inner) = weak_inner.upgrade() else {
388                debug!("Metrics task stopping because client was dropped");
389                break;
390            };
391
392            let metrics_client = BusClient { inner };
393            let connection_key = format!(
394                "{}:{}",
395                metrics_client.inner.service_id, metrics_client.inner.instance_id
396            );
397            match metrics_client
398                .publish(METRICS_TOPIC, METRICS_MESSAGE_TYPE, build_metrics_payload())
399                .await
400            {
401                Ok(subscribers) => {
402                    debug!(%connection_key, subscribers, "Metric published");
403                }
404                Err(err) => {
405                    warn!(%connection_key, %err, "Metric publish failed");
406                }
407            }
408        }
409    });
410}
411
412fn build_metrics_payload() -> HashMap<String, serde_json::Value> {
413    let mut payload = HashMap::new();
414    payload.insert(
415        "name".to_string(),
416        serde_json::Value::String("memory_usage".to_string()),
417    );
418    payload.insert(
419        "value".to_string(),
420        serde_json::Value::from(current_memory_usage_bytes()),
421    );
422    payload.insert(
423        "unit".to_string(),
424        serde_json::Value::String("bytes".to_string()),
425    );
426    payload
427}
428
429fn current_memory_usage_bytes() -> u64 {
430    #[cfg(target_os = "linux")]
431    {
432        if let Ok(status) = std::fs::read_to_string("/proc/self/statm")
433            && let Some(pages_str) = status.split_whitespace().next()
434            && let Ok(pages) = pages_str.parse::<u64>()
435        {
436            return pages.saturating_mul(4096);
437        }
438    }
439
440    0
441}
442
443fn connect_retry_delay(attempt: usize) -> Duration {
444    let exp_factor = 1u128 << attempt.min(16);
445    let base_ms = CONNECT_RETRY_BASE_DELAY
446        .as_millis()
447        .saturating_mul(exp_factor);
448    let capped_ms = base_ms.min(CONNECT_RETRY_MAX_DELAY.as_millis()) as u64;
449    let jitter = (SystemTime::now()
450        .duration_since(UNIX_EPOCH)
451        .map(|d| d.subsec_nanos() as u64)
452        .unwrap_or(0))
453        % (CONNECT_RETRY_MAX_JITTER_MS + 1);
454    Duration::from_millis(capped_ms.saturating_add(jitter))
455}
456
457async fn read_loop(inner: Arc<Inner>, reader: OwnedReadHalf) {
458    let mut reader = BufReader::new(reader);
459    let mut line = String::new();
460    let connection_key = format!("{}:{}", inner.service_id, inner.instance_id);
461
462    debug!(%connection_key, "Read loop started");
463
464    loop {
465        line.clear();
466        let bytes = match reader.read_line(&mut line).await {
467            Ok(size) => size,
468            Err(err) => {
469                warn!(%connection_key, %err, "Read error");
470                break;
471            }
472        };
473        if bytes == 0 {
474            info!(%connection_key, "Socket closed by peer");
475            break;
476        }
477
478        let message: ServiceBusProtocolMessage = match serde_json::from_str(&line) {
479            Ok(msg) => msg,
480            Err(err) => {
481                warn!(%connection_key, %err, "Failed to parse message");
482                continue;
483            }
484        };
485
486        match message.r#type.as_str() {
487            protocol_types::EVENT => {
488                if let Some(event) = message.event
489                    && event.instance_id == inner.instance_id
490                {
491                    dispatch_event(Arc::clone(&inner), event).await;
492                }
493            }
494            protocol_types::ACK | protocol_types::ERROR | protocol_types::FORWARD_RESPONSE => {
495                let sender = {
496                    let mut pending = inner.pending.lock().await;
497                    pending.pop_front()
498                };
499                if let Some(sender) = sender {
500                    let _ = sender.send(message);
501                }
502            }
503            _ => {}
504        }
505    }
506
507    debug!(%connection_key, "Read loop stopped");
508}
509
510async fn dispatch_event(inner: Arc<Inner>, event: ServiceBusEventEnvelope) {
511    let handlers = {
512        let guard = inner.event_handlers.read().await;
513        let mut collected: Vec<EventHandler> = guard.get(&event.topic).cloned().unwrap_or_default();
514        if let Some(wildcard) = guard.get("*") {
515            collected.extend(wildcard.iter().cloned());
516        }
517        collected
518    };
519
520    for handler in handlers {
521        let event_clone = event.clone();
522        tokio::spawn(async move {
523            handler(event_clone).await;
524        });
525    }
526
527    if event.topic == format!("service-{}", inner.service_id) {
528        let request_handler = {
529            let guard = inner.request_handlers.read().await;
530            guard.get(&event.message_type).cloned()
531        };
532
533        if let Some(handler) = request_handler
534            && let Some(reply_to) = event.payload.get("reply_to").and_then(|v| v.as_str())
535        {
536            let req = ServiceBusRequest {
537                event: event.clone(),
538            };
539            let responder = RequestResponder {
540                client: BusClient {
541                    inner: Arc::clone(&inner),
542                },
543                reply_to_topic: reply_to.to_string(),
544                causation_id: event.event_id.clone(),
545                correlation_id: event.correlation_id,
546                default_message_type: event.message_type.clone(),
547            };
548
549            tokio::spawn(async move {
550                let _ = handler(req, responder).await;
551            });
552        }
553    }
554}