lib

Core libraries for Radroots
git clone https://radroots.dev/git/lib.git
Log | Files | Refs | README

socket_write.rs (5258B)


      1 //! One serialized writer for SDK messages and exact signed-event publication.
      2 
      3 use async_wsocket::Message;
      4 use core::{
      5     fmt,
      6     pin::Pin,
      7     task::{Context, Poll},
      8 };
      9 use futures::{Sink, SinkExt, lock::Mutex};
     10 use nostr_relay_pool::transport::{error::TransportError, websocket::WebSocketSink};
     11 use radroots_transport::BoxFuture;
     12 use std::{
     13     collections::BTreeMap,
     14     sync::{
     15         Arc, Mutex as RegistryMutex, Weak,
     16         atomic::{AtomicBool, Ordering},
     17     },
     18 };
     19 
     20 use crate::relay::policy_error;
     21 
     22 pub(crate) struct SocketWriter {
     23     sink: Mutex<WebSocketSink>,
     24     open: AtomicBool,
     25 }
     26 
     27 impl SocketWriter {
     28     pub(crate) fn new(sink: WebSocketSink) -> Arc<Self> {
     29         Arc::new(Self {
     30             sink: Mutex::new(sink),
     31             open: AtomicBool::new(true),
     32         })
     33     }
     34 
     35     pub(crate) async fn send(&self, message: Message) -> Result<(), TransportError> {
     36         let mut sink = self.sink.lock().await;
     37         if !self.open.load(Ordering::Acquire) {
     38             return Err(policy_error("relay writer is closed"));
     39         }
     40         let result = sink.send(message).await;
     41         if result.is_err() {
     42             self.invalidate();
     43         }
     44         result
     45     }
     46 
     47     fn invalidate(&self) {
     48         self.open.store(false, Ordering::Release);
     49     }
     50 
     51     async fn close(&self) -> Result<(), TransportError> {
     52         self.invalidate();
     53         self.sink.lock().await.close().await
     54     }
     55 }
     56 
     57 /// Configured keys only; the registry never owns a connection or event bytes.
     58 #[derive(Clone)]
     59 pub(crate) struct WriterRegistry(Arc<RegistryMutex<BTreeMap<String, Weak<SocketWriter>>>>);
     60 
     61 impl WriterRegistry {
     62     pub(crate) fn new(keys: impl Iterator<Item = String>) -> Self {
     63         Self(Arc::new(RegistryMutex::new(
     64             keys.map(|key| (key, Weak::new())).collect(),
     65         )))
     66     }
     67 
     68     pub(crate) fn install(
     69         &self,
     70         key: &str,
     71         writer: &Arc<SocketWriter>,
     72     ) -> Result<(), TransportError> {
     73         let mut entries = self
     74             .0
     75             .lock()
     76             .map_err(|_| policy_error("relay writer registry unavailable"))?;
     77         let slot = entries
     78             .get_mut(key)
     79             .ok_or_else(|| policy_error("relay writer is not configured"))?;
     80         if let Some(previous) = slot.upgrade() {
     81             previous.invalidate();
     82         }
     83         *slot = Arc::downgrade(writer);
     84         Ok(())
     85     }
     86 
     87     pub(crate) fn get(&self, key: &str) -> Result<Arc<SocketWriter>, TransportError> {
     88         self.0
     89             .lock()
     90             .map_err(|_| policy_error("relay writer registry unavailable"))?
     91             .get(key)
     92             .and_then(Weak::upgrade)
     93             .filter(|writer| writer.open.load(Ordering::Acquire))
     94             .ok_or_else(|| policy_error("relay writer is unavailable"))
     95     }
     96 }
     97 
     98 impl fmt::Debug for WriterRegistry {
     99     fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
    100         formatter.write_str("WriterRegistry([redacted])")
    101     }
    102 }
    103 
    104 /// The SDK retains connection ownership; dropping its sink revokes raw writes.
    105 pub(crate) struct SharedSocketSink {
    106     writer: Arc<SocketWriter>,
    107     pending: Option<BoxFuture<'static, Result<(), TransportError>>>,
    108     closing: bool,
    109 }
    110 
    111 impl SharedSocketSink {
    112     pub(crate) fn new(writer: Arc<SocketWriter>) -> Self {
    113         Self {
    114             writer,
    115             pending: None,
    116             closing: false,
    117         }
    118     }
    119 
    120     fn poll_pending(&mut self, context: &mut Context<'_>) -> Poll<Result<(), TransportError>> {
    121         if let Some(pending) = &mut self.pending {
    122             let result = futures::ready!(pending.as_mut().poll(context));
    123             self.pending = None;
    124             return Poll::Ready(result);
    125         }
    126         Poll::Ready(Ok(()))
    127     }
    128 }
    129 
    130 impl Sink<Message> for SharedSocketSink {
    131     type Error = TransportError;
    132 
    133     fn poll_ready(
    134         mut self: Pin<&mut Self>,
    135         context: &mut Context<'_>,
    136     ) -> Poll<Result<(), Self::Error>> {
    137         if self.closing {
    138             return Poll::Ready(Err(policy_error("relay writer is closing")));
    139         }
    140         self.poll_pending(context)
    141     }
    142 
    143     fn start_send(mut self: Pin<&mut Self>, message: Message) -> Result<(), Self::Error> {
    144         if self.closing || self.pending.is_some() {
    145             return Err(policy_error("relay writer is not ready"));
    146         }
    147         let writer = Arc::clone(&self.writer);
    148         self.pending = Some(Box::pin(async move { writer.send(message).await }));
    149         Ok(())
    150     }
    151 
    152     fn poll_flush(
    153         mut self: Pin<&mut Self>,
    154         context: &mut Context<'_>,
    155     ) -> Poll<Result<(), Self::Error>> {
    156         self.poll_pending(context)
    157     }
    158 
    159     fn poll_close(
    160         mut self: Pin<&mut Self>,
    161         context: &mut Context<'_>,
    162     ) -> Poll<Result<(), Self::Error>> {
    163         futures::ready!(self.poll_pending(context))?;
    164         if !self.closing {
    165             self.closing = true;
    166             self.writer.invalidate();
    167             let writer = Arc::clone(&self.writer);
    168             self.pending = Some(Box::pin(async move { writer.close().await }));
    169         }
    170         self.poll_pending(context)
    171     }
    172 }
    173 
    174 impl Drop for SharedSocketSink {
    175     fn drop(&mut self) {
    176         self.writer.invalidate();
    177     }
    178 }
    179 
    180 #[cfg(test)]
    181 #[path = "socket_write_tests.rs"]
    182 mod tests;