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