aster_forge_http/
response_body.rs

1use futures::StreamExt;
2
3/// Reads a reqwest response body while enforcing a strict byte limit.
4///
5/// A declared `Content-Length` above `max_bytes` is rejected before streaming. The same limit is
6/// enforced while reading so missing or incorrect length headers cannot bypass it. `map_error`
7/// keeps this helper independent from each caller's error boundary.
8///
9/// # Errors
10///
11/// Returns the caller-provided error when the declared or observed body exceeds `max_bytes`,
12/// when the response stream fails, or when the accumulated size would overflow `usize`.
13pub async fn read_reqwest_body_limited<E>(
14    response: reqwest::Response,
15    context: &str,
16    max_bytes: usize,
17    map_error: impl Fn(String) -> E,
18) -> Result<Vec<u8>, E> {
19    if response.content_length().is_some_and(|content_length| {
20        usize::try_from(content_length).map_or(true, |length| length > max_bytes)
21    }) {
22        return Err(map_error(format!(
23            "{context} exceeds {max_bytes} bytes limit"
24        )));
25    }
26    let mut body = Vec::with_capacity(max_bytes.min(4096));
27    let mut stream = response.bytes_stream();
28    while let Some(chunk) = stream.next().await {
29        let chunk = chunk.map_err(|error| map_error(format!("{context}: {error}")))?;
30        extend_body_limited(&mut body, &chunk, context, max_bytes, &map_error)?;
31    }
32    Ok(body)
33}
34
35fn extend_body_limited<E>(
36    body: &mut Vec<u8>,
37    chunk: &[u8],
38    context: &str,
39    max_bytes: usize,
40    map_error: &impl Fn(String) -> E,
41) -> Result<(), E> {
42    let next_len = body
43        .len()
44        .checked_add(chunk.len())
45        .ok_or_else(|| map_error(format!("{context} size overflow")))?;
46    if next_len > max_bytes {
47        return Err(map_error(format!(
48            "{context} exceeds {max_bytes} bytes limit"
49        )));
50    }
51    body.extend_from_slice(chunk);
52    Ok(())
53}
54
55#[cfg(test)]
56mod tests {
57    use super::{extend_body_limited, read_reqwest_body_limited};
58
59    fn message_error(message: String) -> String {
60        message
61    }
62
63    async fn response_with_body(body: &'static [u8], chunked: bool) -> reqwest::Response {
64        use tokio::io::{AsyncReadExt, AsyncWriteExt};
65
66        let listener = tokio::net::TcpListener::bind(("127.0.0.1", 0))
67            .await
68            .expect("test listener should bind");
69        let addr = listener
70            .local_addr()
71            .expect("test listener should expose address");
72        let server = tokio::spawn(async move {
73            let (mut socket, _) = listener
74                .accept()
75                .await
76                .expect("test server should accept request");
77            let mut request = Vec::new();
78            let mut buffer = [0_u8; 1024];
79            loop {
80                let read = socket
81                    .read(&mut buffer)
82                    .await
83                    .expect("test server should read request");
84                if read == 0 {
85                    break;
86                }
87                request.extend_from_slice(&buffer[..read]);
88                if request.windows(4).any(|window| window == b"\r\n\r\n") {
89                    break;
90                }
91            }
92            if chunked {
93                socket
94                    .write_all(
95                        b"HTTP/1.1 200 OK\r\nTransfer-Encoding: chunked\r\nConnection: close\r\n\r\n",
96                    )
97                    .await
98                    .expect("test server should write chunked headers");
99                socket
100                    .write_all(format!("{:x}\r\n", body.len()).as_bytes())
101                    .await
102                    .expect("test server should write chunk size");
103                socket
104                    .write_all(body)
105                    .await
106                    .expect("test server should write chunk");
107                socket
108                    .write_all(b"\r\n0\r\n\r\n")
109                    .await
110                    .expect("test server should finish chunked body");
111            } else {
112                socket
113                    .write_all(
114                        format!(
115                            "HTTP/1.1 200 OK\r\nContent-Length: {}\r\nConnection: close\r\n\r\n",
116                            body.len()
117                        )
118                        .as_bytes(),
119                    )
120                    .await
121                    .expect("test server should write headers");
122                socket
123                    .write_all(body)
124                    .await
125                    .expect("test server should write body");
126            }
127        });
128        let response = reqwest::get(format!("http://{addr}/"))
129            .await
130            .expect("test request should succeed");
131        server.await.expect("test server should finish");
132        response
133    }
134
135    #[test]
136    fn limited_body_accumulation_accepts_exact_limit_and_rejects_one_byte_over() {
137        let mut body = Vec::new();
138        extend_body_limited(&mut body, b"123", "test response body", 4, &message_error)
139            .expect("body below limit should be accepted");
140        extend_body_limited(&mut body, b"4", "test response body", 4, &message_error)
141            .expect("body at exact limit should be accepted");
142        let error = extend_body_limited(&mut body, b"5", "test response body", 4, &message_error)
143            .expect_err("body over limit should be rejected");
144
145        assert_eq!(body, b"1234");
146        assert!(error.contains("exceeds 4 bytes limit"));
147    }
148
149    #[tokio::test]
150    async fn reqwest_body_reader_enforces_exact_network_boundary() {
151        let exact = read_reqwest_body_limited(
152            response_with_body(b"1234", false).await,
153            "test network body",
154            4,
155            message_error,
156        )
157        .await
158        .expect("body at exact network limit should be accepted");
159        assert_eq!(exact, b"1234");
160
161        let error = read_reqwest_body_limited(
162            response_with_body(b"12345", true).await,
163            "test network body",
164            4,
165            message_error,
166        )
167        .await
168        .expect_err("body over network limit should be rejected");
169        assert!(error.contains("exceeds 4 bytes limit"));
170    }
171
172    #[tokio::test]
173    async fn reqwest_body_reader_rejects_declared_length_over_limit() {
174        let error = read_reqwest_body_limited(
175            response_with_body(b"12345", false).await,
176            "declared-length body",
177            4,
178            message_error,
179        )
180        .await
181        .expect_err("declared content length over the limit should fail");
182
183        assert_eq!(error, "declared-length body exceeds 4 bytes limit");
184    }
185}