lib

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

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 }