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}
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 let result = Pin::new(&mut self.inner).poll_read(cx, buf);
172 if result.is_ready() {
173 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 Poll::Pending
182 } else {
183 result
184 }
185 }
186}
187
188pub 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
231pub struct FragmentedStream {
235 inner: Box<dyn AsyncStream>,
236 fragment_size: usize,
237 write_buf: Vec<u8>,
238}
239
240pub 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 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 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 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, );
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}