use anyhow::Result;
use bytes::Bytes;
use saorsa_gossip_transport::{GossipStreamType, GossipTransport};
use saorsa_gossip_types::PeerId;
use std::collections::HashMap;
use std::sync::Arc;
use tokio::sync::{RwLock, mpsc};
use tracing::{debug, warn};
use super::sites::{SiteResponse, SitesWire};
use super::sites_listener::SitesListener;
const RESPONSE_CHANNEL_SIZE: usize = 100;
pub struct SitesDispatcher {
transport: Arc<dyn GossipTransport + Send + Sync>,
listener: Arc<SitesListener>,
response_channels: Arc<RwLock<HashMap<u64, mpsc::Sender<SiteResponse>>>>,
shutdown: Arc<tokio::sync::Notify>,
}
impl SitesDispatcher {
pub fn new(
transport: Arc<dyn GossipTransport + Send + Sync>,
listener: Arc<SitesListener>,
) -> Self {
Self {
transport,
listener,
response_channels: Arc::new(RwLock::new(HashMap::new())),
shutdown: Arc::new(tokio::sync::Notify::new()),
}
}
pub fn start(self: Arc<Self>) -> tokio::task::JoinHandle<()> {
tokio::spawn(async move {
loop {
tokio::select! {
_ = self.shutdown.notified() => {
debug!("Sites dispatcher shutting down");
break;
}
result = self.transport.receive_message() => {
match result {
Ok((peer_id, stream_type, data)) => {
if let Err(e) = self.handle_message(peer_id, stream_type, data).await {
warn!("Failed to handle Sites message: {}", e);
}
}
Err(e) => {
let err_str = e.to_string();
if err_str.contains("No messages available") {
tokio::time::sleep(tokio::time::Duration::from_millis(100)).await;
} else {
warn!("Sites transport receive error: {}", e);
tokio::time::sleep(tokio::time::Duration::from_millis(100)).await;
}
}
}
}
}
}
})
}
pub fn stop(&self) {
self.shutdown.notify_one();
}
async fn handle_message(
&self,
peer_id: PeerId,
stream_type: GossipStreamType,
data: Bytes,
) -> Result<()> {
if stream_type != GossipStreamType::Bulk {
return Ok(());
}
let wire_msg: SitesWire = match bincode::deserialize(&data) {
Ok(msg) => msg,
Err(_) => {
return Ok(());
}
};
match wire_msg {
SitesWire::Request { .. } => {
self.listener
.maybe_handle_incoming(peer_id, stream_type, data)
.await;
Ok(())
}
SitesWire::Response { id, body } => {
let channels = self.response_channels.read().await;
if let Some(tx) = channels.get(&id) {
if tx.send(body).await.is_err() {
warn!("Failed to send response {} - receiver dropped", id);
} else {
debug!("Routed response {} to fetcher", id);
}
} else {
warn!("Received response {} with no waiting fetcher", id);
}
Ok(())
}
}
}
pub async fn register_response_channel(&self, request_id: u64) -> mpsc::Receiver<SiteResponse> {
let (tx, rx) = mpsc::channel(RESPONSE_CHANNEL_SIZE);
let mut channels = self.response_channels.write().await;
channels.insert(request_id, tx);
rx
}
pub async fn unregister_response_channel(&self, request_id: u64) {
let mut channels = self.response_channels.write().await;
channels.remove(&request_id);
}
}