fetch_raw_budget.rs (9781B)
1 use futures::{SinkExt, StreamExt}; 2 use radroots_transport::{ 3 EventSource, FetchRequest, TargetSet, outcome::FetchTargetState, source::FetchBounds, 4 }; 5 use radroots_transport_nostr::{ 6 Config, NostrTransport, RelayAccess, RelayEndpoint, RelayProfile, RelayProfileKind, 7 RelayUrlPolicy, 8 }; 9 use serde_json::Value; 10 use std::time::{Duration, SystemTime, UNIX_EPOCH}; 11 use tokio::net::{TcpListener, TcpStream}; 12 use tokio_tungstenite::{WebSocketStream, accept_async, tungstenite::Message}; 13 14 const WATCHDOG: Duration = Duration::from_secs(10); 15 16 fn transport(urls: &[String]) -> (NostrTransport, FetchRequest) { 17 let profile = RelayProfile::explicit( 18 RelayProfileKind::Simulator, 19 urls.iter().map(|url| { 20 RelayEndpoint::new(url, RelayUrlPolicy::Local, RelayAccess::ReadOnly).unwrap() 21 }), 22 ) 23 .unwrap(); 24 let config = Config::from_profile(profile) 25 .with_timeouts(1000, 5000, 500) 26 .unwrap() 27 .with_max_connections(1) 28 .unwrap(); 29 let targets = TargetSet::new( 30 config 31 .read_relays() 32 .map(|relay| relay.to_target().unwrap()) 33 .collect(), 34 ) 35 .unwrap(); 36 let now = SystemTime::now() 37 .duration_since(UNIX_EPOCH) 38 .unwrap() 39 .as_millis() as u64; 40 let request = FetchRequest::new( 41 "raw-budget-loopback", 42 targets, 43 FetchBounds::new(10, now + 5000).unwrap(), 44 ) 45 .unwrap(); 46 (NostrTransport::new(config), request) 47 } 48 49 async fn listener() -> (TcpListener, String) { 50 let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); 51 let url = format!("ws://{}", listener.local_addr().unwrap()); 52 (listener, url) 53 } 54 55 // A WebSocket pong is produced only after the client's receive stream has 56 // consumed all preceding messages. No scheduling sleep is an admission barrier. 57 async fn barrier(socket: &mut WebSocketStream<TcpStream>, sequence: usize) -> bool { 58 let marker = sequence.to_be_bytes().to_vec(); 59 if socket 60 .send(Message::Ping(marker.clone().into())) 61 .await 62 .is_err() 63 { 64 return false; 65 } 66 while let Some(message) = socket.next().await { 67 match message { 68 Ok(Message::Pong(value)) if value.as_ref() == marker.as_slice() => return true, 69 Ok(Message::Close(_)) | Err(_) => return false, 70 _ => { 71 if socket.flush().await.is_err() { 72 return false; 73 } 74 } 75 } 76 } 77 false 78 } 79 80 async fn serve_malformed(listener: TcpListener, frames: usize) { 81 let (stream, _) = listener.accept().await.unwrap(); 82 let mut socket = accept_async(stream).await.unwrap(); 83 while let Some(message) = socket.next().await { 84 let text = match message { 85 Ok(Message::Text(text)) => text, 86 Ok(Message::Close(_)) | Err(_) => return, 87 _ => { 88 let _ = socket.flush().await; 89 continue; 90 } 91 }; 92 let values: Value = serde_json::from_str(&text).unwrap(); 93 if values[0] != "REQ" { 94 continue; 95 } 96 let malformed = serde_json::to_string(&( 97 "EVENT", 98 &values[1], 99 serde_json::json!({ "invalid_event": "x".repeat(384 * 1024) }), 100 )) 101 .unwrap(); 102 assert!(malformed.len() < 512 * 1024); 103 for sequence in 0..frames { 104 if socket 105 .send(Message::Text(malformed.clone().into())) 106 .await 107 .is_err() 108 || !barrier(&mut socket, sequence).await 109 { 110 return; 111 } 112 } 113 if socket 114 .send(Message::Text( 115 serde_json::to_string(&("EOSE", &values[1])).unwrap().into(), 116 )) 117 .await 118 .is_err() 119 { 120 return; 121 } 122 } 123 } 124 125 async fn fetch_malformed(frames: usize) -> FetchTargetState { 126 let (listener, url) = listener().await; 127 let server = tokio::spawn(serve_malformed(listener, frames)); 128 let (transport, request) = transport(&[url]); 129 let result = tokio::time::timeout(WATCHDOG, transport.fetch(request)).await; 130 server.abort(); 131 let joined = server.await; 132 assert!(joined.is_ok() || joined.unwrap_err().is_cancelled()); 133 let page = result.unwrap().unwrap(); 134 assert!(page.events().is_empty()); 135 assert_eq!(page.target_outcomes().len(), 1); 136 page.target_outcomes()[0].state() 137 } 138 139 #[tokio::test(flavor = "multi_thread")] 140 async fn empty_eose_remains_complete_with_raw_ingress_accounting() { 141 assert_eq!(fetch_malformed(0).await, FetchTargetState::Complete); 142 } 143 144 #[tokio::test(flavor = "multi_thread")] 145 async fn malformed_event_traffic_exhausts_raw_budget_before_eose() { 146 assert_eq!(fetch_malformed(24).await, FetchTargetState::Partial); 147 } 148 149 async fn finish_server(mut server: tokio::task::JoinHandle<()>) { 150 let result = tokio::time::timeout(WATCHDOG, &mut server).await; 151 if result.is_err() { 152 server.abort(); 153 let _ = server.await; 154 } 155 result 156 .expect("client must close the bounded relay connection") 157 .unwrap(); 158 } 159 160 #[tokio::test(flavor = "multi_thread")] 161 async fn completed_batches_do_not_refund_raw_bytes_for_later_relays() { 162 let (one, one_url) = listener().await; 163 let (two, two_url) = listener().await; 164 let servers = [ 165 tokio::spawn(serve_malformed(one, 12)), 166 tokio::spawn(serve_malformed(two, 12)), 167 ]; 168 let (transport, request) = transport(&[one_url, two_url]); 169 let result = tokio::time::timeout(WATCHDOG, transport.fetch(request)).await; 170 for server in servers { 171 finish_server(server).await; 172 } 173 let page = result.unwrap().unwrap(); 174 assert!(page.events().is_empty()); 175 assert_eq!( 176 page.target_outcomes() 177 .iter() 178 .filter(|outcome| outcome.state() == FetchTargetState::Complete) 179 .count(), 180 1 181 ); 182 assert_eq!( 183 page.target_outcomes() 184 .iter() 185 .filter(|outcome| outcome.state() == FetchTargetState::Partial) 186 .count(), 187 1 188 ); 189 } 190 191 #[tokio::test(flavor = "multi_thread")] 192 async fn unsolicited_malformed_traffic_is_bounded_before_the_request() { 193 let (listener, url) = listener().await; 194 let server = tokio::spawn(async move { 195 let (stream, _) = listener.accept().await.unwrap(); 196 let mut socket = accept_async(stream).await.unwrap(); 197 for sequence in 0..24 { 198 let malformed = format!( 199 "[\"EVENT\",\"unsolicited\",{{\"invalid\":\"{}\"}}]", 200 "x".repeat(384 * 1024) 201 ); 202 if socket.send(Message::Text(malformed.into())).await.is_err() 203 || !barrier(&mut socket, sequence).await 204 { 205 return; 206 } 207 } 208 panic!("oversized unsolicited inventory must close before admission"); 209 }); 210 let (transport, request) = transport(&[url]); 211 let result = tokio::time::timeout(WATCHDOG, transport.fetch(request)).await; 212 finish_server(server).await; 213 assert_eq!( 214 result.unwrap().unwrap().target_outcomes()[0].state(), 215 FetchTargetState::Partial 216 ); 217 } 218 219 #[tokio::test(flavor = "multi_thread")] 220 async fn zero_payload_continuation_frames_consume_work_before_message_decoding() { 221 use tokio_tungstenite::tungstenite::protocol::frame::{ 222 Frame, 223 coding::{Data, OpCode}, 224 }; 225 let (listener, url) = listener().await; 226 let server = tokio::spawn(async move { 227 let (stream, _) = listener.accept().await.unwrap(); 228 let mut socket = accept_async(stream).await.unwrap(); 229 while let Some(message) = socket.next().await { 230 let text = match message { 231 Ok(Message::Text(text)) => text, 232 Ok(Message::Close(_)) | Err(_) => return, 233 _ => { 234 let _ = socket.flush().await; 235 continue; 236 } 237 }; 238 let values: Value = serde_json::from_str(&text).unwrap(); 239 if values[0] != "REQ" { 240 continue; 241 } 242 socket 243 .send(Message::Frame(Frame::message( 244 Vec::new(), 245 OpCode::Data(Data::Text), 246 false, 247 ))) 248 .await 249 .unwrap(); 250 for sequence in 0..8193 { 251 if socket 252 .send(Message::Frame(Frame::message( 253 Vec::new(), 254 OpCode::Data(Data::Continue), 255 false, 256 ))) 257 .await 258 .is_err() 259 { 260 return; 261 } 262 if sequence % 128 == 0 && !barrier(&mut socket, sequence).await { 263 return; 264 } 265 } 266 if socket 267 .send(Message::Frame(Frame::message( 268 b"[]".to_vec(), 269 OpCode::Data(Data::Continue), 270 true, 271 ))) 272 .await 273 .is_err() 274 { 275 return; 276 } 277 let _ = socket 278 .send(Message::Text( 279 serde_json::to_string(&("EOSE", &values[1])).unwrap().into(), 280 )) 281 .await; 282 } 283 }); 284 let (transport, request) = transport(&[url]); 285 let result = tokio::time::timeout(WATCHDOG, transport.fetch(request)).await; 286 finish_server(server).await; 287 let page = result.unwrap().unwrap(); 288 assert!(page.events().is_empty()); 289 assert_eq!(page.target_outcomes()[0].state(), FetchTargetState::Partial); 290 }