1pub 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
29pub 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
35pub 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
68pub 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
100pub 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 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 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
150pub struct SlowReader {
154 inner: tokio::net::tcp::OwnedReadHalf,
155 delay: Duration,
156 primed: bool,
159}
160
161impl SlowReader {
162 pub fn new(inner: tokio::net::tcp::OwnedReadHalf, delay: Duration) -> Self {
163 Self {
164 inner,
165 delay,
166 primed: false,
167 }
168 }
169}
170
171impl AsyncRead for SlowReader {
172 fn poll_read(
173 mut self: Pin<&mut Self>,
174 cx: &mut Context<'_>,
175 buf: &mut tokio::io::ReadBuf<'_>,
176 ) -> Poll<std::io::Result<()>> {
177 if !self.primed {
182 self.primed = true;
183 let waker = cx.waker().clone();
184 let delay = self.delay;
185 tokio::spawn(async move {
186 tokio::time::sleep(delay).await;
187 waker.wake();
188 });
189 return Poll::Pending;
190 }
191 self.primed = false;
192 Pin::new(&mut self.inner).poll_read(cx, buf)
193 }
194}
195
196pub struct SlowWriter {
200 inner: tokio::net::tcp::OwnedWriteHalf,
201 delay: Duration,
202 primed: bool,
205}
206
207impl SlowWriter {
208 pub fn new(inner: tokio::net::tcp::OwnedWriteHalf, delay: Duration) -> Self {
209 Self {
210 inner,
211 delay,
212 primed: false,
213 }
214 }
215}
216
217impl AsyncWrite for SlowWriter {
218 fn poll_write(
219 mut self: Pin<&mut Self>,
220 cx: &mut Context<'_>,
221 buf: &[u8],
222 ) -> Poll<std::io::Result<usize>> {
223 if !self.primed {
226 self.primed = true;
227 let waker = cx.waker().clone();
228 let delay = self.delay;
229 tokio::spawn(async move {
230 tokio::time::sleep(delay).await;
231 waker.wake();
232 });
233 return Poll::Pending;
234 }
235 self.primed = false;
236 Pin::new(&mut self.inner).poll_write(cx, buf)
237 }
238
239 fn poll_flush(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<std::io::Result<()>> {
240 Pin::new(&mut self.inner).poll_flush(cx)
241 }
242
243 fn poll_shutdown(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<std::io::Result<()>> {
244 Pin::new(&mut self.inner).poll_shutdown(cx)
245 }
246}
247
248pub struct FragmentedStream {
252 inner: Box<dyn AsyncStream>,
253 fragment_size: usize,
254 write_buf: Vec<u8>,
255}
256
257pub trait AsyncStream: AsyncRead + AsyncWrite + Send + Unpin {}
259impl<T: AsyncRead + AsyncWrite + Send + Unpin> AsyncStream for T {}
260
261impl FragmentedStream {
262 pub fn new(inner: Box<dyn AsyncStream>, fragment_size: usize) -> Self {
263 Self {
264 inner,
265 fragment_size,
266 write_buf: Vec::new(),
267 }
268 }
269}
270
271impl AsyncRead for FragmentedStream {
272 fn poll_read(
273 mut self: Pin<&mut Self>,
274 cx: &mut Context<'_>,
275 buf: &mut tokio::io::ReadBuf<'_>,
276 ) -> Poll<std::io::Result<()>> {
277 Pin::new(&mut self.inner).poll_read(cx, buf)
278 }
279}
280
281impl AsyncWrite for FragmentedStream {
282 fn poll_write(
283 mut self: Pin<&mut Self>,
284 _cx: &mut Context<'_>,
285 buf: &[u8],
286 ) -> Poll<std::io::Result<usize>> {
287 self.write_buf.extend_from_slice(buf);
290 Poll::Ready(Ok(buf.len()))
291 }
292
293 fn poll_flush(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<std::io::Result<()>> {
294 while !self.write_buf.is_empty() {
296 let chunk_len = self.write_buf.len().min(self.fragment_size);
297 let chunk = self.write_buf[..chunk_len].to_vec();
298
299 match Pin::new(&mut self.inner).poll_write(cx, &chunk) {
300 Poll::Ready(Ok(n)) => {
301 self.write_buf.drain(..n);
302 }
303 Poll::Ready(Err(e)) => return Poll::Ready(Err(e)),
304 Poll::Pending => return Poll::Pending,
305 }
306 }
307 Pin::new(&mut self.inner).poll_flush(cx)
308 }
309
310 fn poll_shutdown(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<std::io::Result<()>> {
311 if !self.write_buf.is_empty() {
313 let _ = self.as_mut().poll_flush(cx);
314 }
315 Pin::new(&mut self.inner).poll_shutdown(cx)
316 }
317}
318
319#[cfg(test)]
320mod tests {
321 use super::*;
322
323 #[tokio::test]
324 async fn test_echo_server() {
325 let (addr, jh) = start_echo_server().await;
326
327 let mut stream = tokio::net::TcpStream::connect(addr).await.unwrap();
328 stream.write_all(b"hello").await.unwrap();
329 stream.shutdown().await.unwrap();
330
331 let mut buf = [0u8; 5];
332 stream.read_exact(&mut buf).await.unwrap();
333 assert_eq!(&buf, b"hello");
334
335 jh.abort();
336 }
337
338 #[tokio::test]
339 async fn test_half_close_server() {
340 let (addr, jh) = start_half_close_server().await;
341
342 let mut stream = tokio::net::TcpStream::connect(addr).await.unwrap();
343 stream.write_all(b"request").await.unwrap();
344 stream.shutdown().await.unwrap();
345
346 let mut buf = [0u8; 7];
347 stream.read_exact(&mut buf).await.unwrap();
348 assert_eq!(&buf, b"request");
349
350 jh.abort();
351 }
352
353 #[tokio::test]
354 async fn test_http_origin_server() {
355 let (addr, jh) = start_http_origin_server().await;
356
357 let mut stream = tokio::net::TcpStream::connect(addr).await.unwrap();
358 stream
359 .write_all(b"GET / HTTP/1.1\r\nHost: example.com\r\n\r\n")
360 .await
361 .unwrap();
362
363 let mut buf = Vec::new();
364 stream.read_to_end(&mut buf).await.unwrap();
365 let response = String::from_utf8_lossy(&buf);
366 assert!(response.contains("200 OK"));
367 assert!(response.contains("hello from origin"));
368
369 jh.abort();
370 }
371
372 #[tokio::test]
373 async fn test_get_free_port() {
374 let port = get_free_port().await;
375 assert!(port > 0);
376 }
377
378 #[tokio::test]
379 async fn test_fragmented_stream() {
380 let (addr, jh) = start_echo_server().await;
381
382 let stream = tokio::net::TcpStream::connect(addr).await.unwrap();
383 let (read_half, write_half) = stream.into_split();
384 let fragmented = FragmentedStream::new(
385 Box::new(tokio::io::join(read_half, write_half)),
386 3, );
388
389 let mut stream = fragmented;
390 stream.write_all(b"hello world").await.unwrap();
391 stream.shutdown().await.unwrap();
392
393 let mut buf = Vec::new();
394 stream.read_to_end(&mut buf).await.unwrap();
395 assert_eq!(&buf, b"hello world");
396
397 jh.abort();
398 }
399}