Skip to main content

eggress_testkit/
lib.rs

1//! Test utilities for eggress.
2//!
3//! Provides async test servers, port allocation helpers, and
4//! protocol test harnesses for fragmented and slow I/O scenarios.
5
6pub mod canonical_manifest;
7pub mod case_model;
8pub mod composition;
9pub mod corpus;
10pub mod differential;
11pub mod eggress_runner;
12pub mod fixtures;
13pub mod manifest;
14pub mod oracle;
15pub mod pproxy_oracle;
16pub mod report;
17pub mod strict_comparators;
18pub mod strict_manifest;
19pub mod strict_observations;
20
21use std::net::SocketAddr;
22use std::pin::Pin;
23use std::task::{Context, Poll};
24use std::time::Duration;
25
26use tokio::io::{AsyncRead, AsyncReadExt, AsyncWrite, AsyncWriteExt};
27use tokio::net::TcpListener;
28
29/// Get a free port by binding to port 0.
30pub async fn get_free_port() -> u16 {
31    let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
32    listener.local_addr().unwrap().port()
33}
34
35/// Start an async echo server that echoes back received bytes.
36///
37/// Returns the address the server is listening on and a join handle.
38pub async fn start_echo_server() -> (SocketAddr, tokio::task::JoinHandle<()>) {
39    let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
40    let addr = listener.local_addr().unwrap();
41
42    let jh = tokio::spawn(async move {
43        loop {
44            let (mut stream, _) = match listener.accept().await {
45                Ok(s) => s,
46                Err(_) => break,
47            };
48            tokio::spawn(async move {
49                let mut buf = [0u8; 4096];
50                loop {
51                    match stream.read(&mut buf).await {
52                        Ok(0) => break,
53                        Ok(n) => {
54                            if stream.write_all(&buf[..n]).await.is_err() {
55                                break;
56                            }
57                        }
58                        Err(_) => break,
59                    }
60                }
61            });
62        }
63    });
64
65    (addr, jh)
66}
67
68/// Start a half-close server: reads until EOF, then sends a response and closes.
69///
70/// Returns the address the server is listening on and a join handle.
71pub async fn start_half_close_server() -> (SocketAddr, tokio::task::JoinHandle<()>) {
72    let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
73    let addr = listener.local_addr().unwrap();
74
75    let jh = tokio::spawn(async move {
76        loop {
77            let (mut stream, _) = match listener.accept().await {
78                Ok(s) => s,
79                Err(_) => break,
80            };
81            tokio::spawn(async move {
82                let mut data = Vec::new();
83                let mut buf = [0u8; 4096];
84                loop {
85                    match stream.read(&mut buf).await {
86                        Ok(0) => break,
87                        Ok(n) => data.extend_from_slice(&buf[..n]),
88                        Err(_) => return,
89                    }
90                }
91                let _ = stream.write_all(&data).await;
92                let _ = stream.shutdown().await;
93            });
94        }
95    });
96
97    (addr, jh)
98}
99
100/// Start a minimal HTTP origin server that responds to any request.
101///
102/// Returns a 200 OK with a fixed body. Useful for testing HTTP forward proxying.
103pub async fn start_http_origin_server() -> (SocketAddr, tokio::task::JoinHandle<()>) {
104    let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
105    let addr = listener.local_addr().unwrap();
106
107    let jh = tokio::spawn(async move {
108        loop {
109            let (mut stream, _) = match listener.accept().await {
110                Ok(s) => s,
111                Err(_) => break,
112            };
113            tokio::spawn(async move {
114                // Read the full request (headers + any body)
115                let mut buf = [0u8; 4096];
116                let mut request_data = Vec::new();
117                loop {
118                    match stream.read(&mut buf).await {
119                        Ok(0) => return,
120                        Ok(n) => {
121                            request_data.extend_from_slice(&buf[..n]);
122                            if request_data.windows(4).any(|w| w == b"\r\n\r\n") {
123                                break;
124                            }
125                        }
126                        Err(_) => return,
127                    }
128                }
129
130                // Check if there's a Content-Length to read body
131                let _response_str = String::from_utf8_lossy(&request_data);
132                let body = b"hello from origin";
133                let response = format!(
134                    "HTTP/1.1 200 OK\r\n\
135                     Content-Length: {}\r\n\
136                     Connection: close\r\n\
137                     \r\n",
138                    body.len()
139                );
140                let _ = stream.write_all(response.as_bytes()).await;
141                let _ = stream.write_all(body).await;
142                let _ = stream.shutdown().await;
143            });
144        }
145    });
146
147    (addr, jh)
148}
149
150/// A wrapper that reads from an inner stream with an artificial delay.
151///
152/// Useful for testing timeout behavior and slow-reader scenarios.
153pub struct SlowReader {
154    inner: tokio::net::tcp::OwnedReadHalf,
155    delay: Duration,
156}
157
158impl SlowReader {
159    pub fn new(inner: tokio::net::tcp::OwnedReadHalf, delay: Duration) -> Self {
160        Self { inner, delay }
161    }
162}
163
164impl AsyncRead for SlowReader {
165    fn poll_read(
166        mut self: Pin<&mut Self>,
167        cx: &mut Context<'_>,
168        buf: &mut tokio::io::ReadBuf<'_>,
169    ) -> Poll<std::io::Result<()>> {
170        // First poll the inner read
171        let result = Pin::new(&mut self.inner).poll_read(cx, buf);
172        if result.is_ready() {
173            // After a successful read, insert a delay by re-registering waker
174            let waker = cx.waker().clone();
175            let delay = self.delay;
176            tokio::spawn(async move {
177                tokio::time::sleep(delay).await;
178                waker.wake();
179            });
180            // Return Pending to simulate slowness
181            Poll::Pending
182        } else {
183            result
184        }
185    }
186}
187
188/// A wrapper that writes to an inner stream with an artificial delay.
189///
190/// Useful for testing timeout behavior and slow-writer scenarios.
191pub struct SlowWriter {
192    inner: tokio::net::tcp::OwnedWriteHalf,
193    delay: Duration,
194}
195
196impl SlowWriter {
197    pub fn new(inner: tokio::net::tcp::OwnedWriteHalf, delay: Duration) -> Self {
198        Self { inner, delay }
199    }
200}
201
202impl AsyncWrite for SlowWriter {
203    fn poll_write(
204        mut self: Pin<&mut Self>,
205        cx: &mut Context<'_>,
206        buf: &[u8],
207    ) -> Poll<std::io::Result<usize>> {
208        let result = Pin::new(&mut self.inner).poll_write(cx, buf);
209        if result.is_ready() {
210            let waker = cx.waker().clone();
211            let delay = self.delay;
212            tokio::spawn(async move {
213                tokio::time::sleep(delay).await;
214                waker.wake();
215            });
216            Poll::Pending
217        } else {
218            result
219        }
220    }
221
222    fn poll_flush(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<std::io::Result<()>> {
223        Pin::new(&mut self.inner).poll_flush(cx)
224    }
225
226    fn poll_shutdown(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<std::io::Result<()>> {
227        Pin::new(&mut self.inner).poll_shutdown(cx)
228    }
229}
230
231/// A stream wrapper that fragments all writes into small chunks.
232///
233/// Useful for testing protocol parsers with fragmented data.
234pub struct FragmentedStream {
235    inner: Box<dyn AsyncStream>,
236    fragment_size: usize,
237    write_buf: Vec<u8>,
238}
239
240/// A trait combining AsyncRead + AsyncWrite + Send + Unpin for test streams.
241pub trait AsyncStream: AsyncRead + AsyncWrite + Send + Unpin {}
242impl<T: AsyncRead + AsyncWrite + Send + Unpin> AsyncStream for T {}
243
244impl FragmentedStream {
245    pub fn new(inner: Box<dyn AsyncStream>, fragment_size: usize) -> Self {
246        Self {
247            inner,
248            fragment_size,
249            write_buf: Vec::new(),
250        }
251    }
252}
253
254impl AsyncRead for FragmentedStream {
255    fn poll_read(
256        mut self: Pin<&mut Self>,
257        cx: &mut Context<'_>,
258        buf: &mut tokio::io::ReadBuf<'_>,
259    ) -> Poll<std::io::Result<()>> {
260        Pin::new(&mut self.inner).poll_read(cx, buf)
261    }
262}
263
264impl AsyncWrite for FragmentedStream {
265    fn poll_write(
266        mut self: Pin<&mut Self>,
267        _cx: &mut Context<'_>,
268        buf: &[u8],
269    ) -> Poll<std::io::Result<usize>> {
270        // Buffer the data and return that we accepted it all
271        // The actual flushing happens in poll_flush
272        self.write_buf.extend_from_slice(buf);
273        Poll::Ready(Ok(buf.len()))
274    }
275
276    fn poll_flush(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<std::io::Result<()>> {
277        // Write fragments from the buffer
278        while !self.write_buf.is_empty() {
279            let chunk_len = self.write_buf.len().min(self.fragment_size);
280            let chunk = self.write_buf[..chunk_len].to_vec();
281
282            match Pin::new(&mut self.inner).poll_write(cx, &chunk) {
283                Poll::Ready(Ok(n)) => {
284                    self.write_buf.drain(..n);
285                }
286                Poll::Ready(Err(e)) => return Poll::Ready(Err(e)),
287                Poll::Pending => return Poll::Pending,
288            }
289        }
290        Pin::new(&mut self.inner).poll_flush(cx)
291    }
292
293    fn poll_shutdown(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<std::io::Result<()>> {
294        // Flush remaining buffered data first
295        if !self.write_buf.is_empty() {
296            let _ = self.as_mut().poll_flush(cx);
297        }
298        Pin::new(&mut self.inner).poll_shutdown(cx)
299    }
300}
301
302#[cfg(test)]
303mod tests {
304    use super::*;
305
306    #[tokio::test]
307    async fn test_echo_server() {
308        let (addr, jh) = start_echo_server().await;
309
310        let mut stream = tokio::net::TcpStream::connect(addr).await.unwrap();
311        stream.write_all(b"hello").await.unwrap();
312        stream.shutdown().await.unwrap();
313
314        let mut buf = [0u8; 5];
315        stream.read_exact(&mut buf).await.unwrap();
316        assert_eq!(&buf, b"hello");
317
318        jh.abort();
319    }
320
321    #[tokio::test]
322    async fn test_half_close_server() {
323        let (addr, jh) = start_half_close_server().await;
324
325        let mut stream = tokio::net::TcpStream::connect(addr).await.unwrap();
326        stream.write_all(b"request").await.unwrap();
327        stream.shutdown().await.unwrap();
328
329        let mut buf = [0u8; 7];
330        stream.read_exact(&mut buf).await.unwrap();
331        assert_eq!(&buf, b"request");
332
333        jh.abort();
334    }
335
336    #[tokio::test]
337    async fn test_http_origin_server() {
338        let (addr, jh) = start_http_origin_server().await;
339
340        let mut stream = tokio::net::TcpStream::connect(addr).await.unwrap();
341        stream
342            .write_all(b"GET / HTTP/1.1\r\nHost: example.com\r\n\r\n")
343            .await
344            .unwrap();
345
346        let mut buf = Vec::new();
347        stream.read_to_end(&mut buf).await.unwrap();
348        let response = String::from_utf8_lossy(&buf);
349        assert!(response.contains("200 OK"));
350        assert!(response.contains("hello from origin"));
351
352        jh.abort();
353    }
354
355    #[tokio::test]
356    async fn test_get_free_port() {
357        let port = get_free_port().await;
358        assert!(port > 0);
359    }
360
361    #[tokio::test]
362    async fn test_fragmented_stream() {
363        let (addr, jh) = start_echo_server().await;
364
365        let stream = tokio::net::TcpStream::connect(addr).await.unwrap();
366        let (read_half, write_half) = stream.into_split();
367        let fragmented = FragmentedStream::new(
368            Box::new(tokio::io::join(read_half, write_half)),
369            3, // 3-byte fragments
370        );
371
372        let mut stream = fragmented;
373        stream.write_all(b"hello world").await.unwrap();
374        stream.shutdown().await.unwrap();
375
376        let mut buf = Vec::new();
377        stream.read_to_end(&mut buf).await.unwrap();
378        assert_eq!(&buf, b"hello world");
379
380        jh.abort();
381    }
382}