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
27pub type EventHandler =
29 Arc<dyn Fn(ServiceBusEventEnvelope) -> BoxFuture<'static, ()> + Send + Sync>;
30pub type RequestHandler = Arc<
32 dyn Fn(ServiceBusRequest, RequestResponder) -> BoxFuture<'static, ClientResult<()>>
33 + Send
34 + Sync,
35>;
36
37#[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)]
53pub struct ServiceBusRequest {
55 pub event: ServiceBusEventEnvelope,
57}
58
59impl ServiceBusRequest {
60 pub fn reply_to(&self) -> Option<&str> {
62 self.event.payload.get("reply_to")?.as_str()
63 }
64}
65
66#[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 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 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)]
108pub struct ForwardRequest {
110 pub target_service_id: String,
112 pub message_type: String,
114 pub payload: HashMap<String, serde_json::Value>,
116 pub timeout_ms: Option<u64>,
118}
119
120impl BusClient {
121 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 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 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 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 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 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 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 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}