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
26pub type EventHandler =
28 Arc<dyn Fn(ServiceBusEventEnvelope) -> BoxFuture<'static, ()> + Send + Sync>;
29pub type RequestHandler = Arc<
31 dyn Fn(ServiceBusRequest, RequestResponder) -> BoxFuture<'static, ClientResult<()>>
32 + Send
33 + Sync,
34>;
35
36#[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)]
52pub struct ServiceBusRequest {
54 pub event: ServiceBusEventEnvelope,
56}
57
58impl ServiceBusRequest {
59 pub fn reply_to(&self) -> Option<&str> {
61 self.event.payload.get("reply_to")?.as_str()
62 }
63}
64
65#[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 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 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)]
107pub struct ForwardRequest {
109 pub target_service_id: String,
111 pub message_type: String,
113 pub payload: HashMap<String, serde_json::Value>,
115 pub timeout_ms: Option<u64>,
117}
118
119impl BusClient {
120 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 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 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 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 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 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 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 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}