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;