1use asupersync::stream::Stream;
4use fastapi_core::{BodyStream, Response, ResponseBody, StatusCode};
5use std::borrow::Cow;
6use std::pin::Pin;
7use std::task::{Context, Poll};
8
9pub enum ResponseWrite {
11 Full(Vec<u8>),
13 Stream(ChunkedEncoder),
15}
16
17#[derive(Debug, Clone, Default)]
33pub struct Trailers {
34 headers: Vec<(String, String)>,
35}
36
37impl Trailers {
38 #[must_use]
40 pub fn new() -> Self {
41 Self::default()
42 }
43
44 #[must_use]
46 pub fn add(mut self, name: impl Into<String>, value: impl Into<String>) -> Self {
47 self.headers.push((name.into(), value.into()));
48 self
49 }
50
51 #[must_use]
53 pub fn is_empty(&self) -> bool {
54 self.headers.is_empty()
55 }
56
57 #[must_use]
60 pub fn trailer_header_value(&self) -> String {
61 self.headers
62 .iter()
63 .map(|(n, _)| n.as_str())
64 .collect::<Vec<_>>()
65 .join(", ")
66 }
67
68 fn encode(&self) -> Vec<u8> {
72 let mut out = Vec::new();
73 for (name, value) in &self.headers {
74 write_header_line(&mut out, name, value.as_bytes());
75 }
76 out
77 }
78}
79
80pub struct ChunkedEncoder {
82 head: Option<Vec<u8>>,
83 body: BodyStream,
84 finished: bool,
85 trailers: Option<Trailers>,
86}
87
88impl ChunkedEncoder {
89 fn new(head: Vec<u8>, body: BodyStream) -> Self {
90 Self {
91 head: Some(head),
92 body,
93 finished: false,
94 trailers: None,
95 }
96 }
97
98 #[must_use]
100 pub fn with_trailers(mut self, trailers: Trailers) -> Self {
101 self.trailers = Some(trailers);
102 self
103 }
104
105 fn encode_chunk(chunk: &[u8]) -> Vec<u8> {
106 use std::io::Write as _;
109 let mut out = Vec::with_capacity(20 + chunk.len() + 4);
110 write!(out, "{:x}\r\n", chunk.len()).expect("write to Vec cannot fail");
111 out.extend_from_slice(chunk);
112 out.extend_from_slice(b"\r\n");
113 out
114 }
115
116 fn encode_final_chunk(&self) -> Vec<u8> {
122 let mut out = Vec::new();
123 out.extend_from_slice(b"0\r\n");
124 if let Some(ref trailers) = self.trailers {
125 out.extend_from_slice(&trailers.encode());
126 }
127 out.extend_from_slice(b"\r\n");
128 out
129 }
130}
131
132impl Stream for ChunkedEncoder {
133 type Item = Vec<u8>;
134
135 fn poll_next(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
136 if let Some(head) = self.head.take() {
137 return Poll::Ready(Some(head));
138 }
139
140 if self.finished {
141 return Poll::Ready(None);
142 }
143
144 loop {
145 match self.body.as_mut().poll_next(cx) {
146 Poll::Pending => return Poll::Pending,
147 Poll::Ready(Some(chunk)) => {
148 if chunk.is_empty() {
149 continue;
150 }
151 return Poll::Ready(Some(Self::encode_chunk(&chunk)));
152 }
153 Poll::Ready(None) => {
154 self.finished = true;
155 return Poll::Ready(Some(self.encode_final_chunk()));
156 }
157 }
158 }
159 }
160}
161
162pub struct ResponseWriter {
164 buffer: Vec<u8>,
165}
166
167impl ResponseWriter {
168 #[must_use]
170 pub fn new() -> Self {
171 Self {
172 buffer: Vec::with_capacity(4096),
173 }
174 }
175
176 #[must_use]
178 pub fn write(&mut self, response: Response) -> ResponseWrite {
179 let (status, headers, body) = response.into_parts();
180 match body {
181 ResponseBody::Empty => {
182 let bytes = self.write_full(status, &headers, &[]);
183 ResponseWrite::Full(bytes)
184 }
185 ResponseBody::Bytes(body) => {
186 let bytes = self.write_full(status, &headers, &body);
187 ResponseWrite::Full(bytes)
188 }
189 ResponseBody::Stream(body) => {
190 let head = self.write_stream_head(status, &headers);
191 ResponseWrite::Stream(ChunkedEncoder::new(head, body))
192 }
193 }
194 }
195
196 fn write_full(
197 &mut self,
198 status: StatusCode,
199 headers: &[(String, Vec<u8>)],
200 body: &[u8],
201 ) -> Vec<u8> {
202 self.buffer.clear();
203
204 self.buffer.extend_from_slice(b"HTTP/1.1 ");
206 self.write_status(status);
207 self.buffer.extend_from_slice(b"\r\n");
208
209 for (name, value) in headers {
211 if is_content_length(name) || is_transfer_encoding(name) {
212 continue;
213 }
214 write_header_line(&mut self.buffer, name, value);
215 }
216
217 self.buffer.extend_from_slice(b"content-length: ");
219 self.buffer
220 .extend_from_slice(body.len().to_string().as_bytes());
221 self.buffer.extend_from_slice(b"\r\n");
222
223 self.buffer.extend_from_slice(b"\r\n");
225
226 self.buffer.extend_from_slice(body);
228
229 self.take_buffer()
230 }
231
232 fn write_stream_head(&mut self, status: StatusCode, headers: &[(String, Vec<u8>)]) -> Vec<u8> {
233 self.buffer.clear();
234
235 self.buffer.extend_from_slice(b"HTTP/1.1 ");
237 self.write_status(status);
238 self.buffer.extend_from_slice(b"\r\n");
239
240 for (name, value) in headers {
242 if is_content_length(name) || is_transfer_encoding(name) {
243 continue;
244 }
245 write_header_line(&mut self.buffer, name, value);
246 }
247
248 self.buffer
250 .extend_from_slice(b"transfer-encoding: chunked\r\n");
251
252 self.buffer.extend_from_slice(b"\r\n");
254
255 self.take_buffer()
256 }
257
258 fn write_status(&mut self, status: StatusCode) {
259 let code = status.as_u16();
260 self.buffer.extend_from_slice(code.to_string().as_bytes());
261 self.buffer.extend_from_slice(b" ");
262 self.buffer
263 .extend_from_slice(status.canonical_reason().as_bytes());
264 }
265
266 fn take_buffer(&mut self) -> Vec<u8> {
267 let mut out = Vec::new();
268 std::mem::swap(&mut out, &mut self.buffer);
269 self.buffer = Vec::with_capacity(out.capacity());
270 out
271 }
272}
273
274fn is_content_length(name: &str) -> bool {
275 name.eq_ignore_ascii_case("content-length")
276}
277
278fn is_transfer_encoding(name: &str) -> bool {
279 name.eq_ignore_ascii_case("transfer-encoding")
280}
281
282fn write_header_line(buffer: &mut Vec<u8>, name: &str, value: &[u8]) {
283 if !is_valid_header_name(name) {
284 return;
285 }
286 buffer.extend_from_slice(name.as_bytes());
287 buffer.extend_from_slice(b": ");
288 buffer.extend_from_slice(sanitize_header_value(value).as_ref());
289 buffer.extend_from_slice(b"\r\n");
290}
291
292fn sanitize_header_value(value: &[u8]) -> Cow<'_, [u8]> {
293 if value
294 .iter()
295 .all(|&byte| byte != b'\r' && byte != b'\n' && byte != 0)
296 {
297 return Cow::Borrowed(value);
298 }
299 Cow::Owned(
300 value
301 .iter()
302 .copied()
303 .filter(|&byte| byte != b'\r' && byte != b'\n' && byte != 0)
304 .collect(),
305 )
306}
307
308fn is_valid_header_name(name: &str) -> bool {
309 !name.is_empty()
310 && name.bytes().all(|byte| {
311 matches!(
312 byte,
313 b'!' | b'#'
314 | b'$'
315 | b'%'
316 | b'&'
317 | b'\''
318 | b'*'
319 | b'+'
320 | b'-'
321 | b'.'
322 | b'0'..=b'9'
323 | b'A'..=b'Z'
324 | b'^'
325 | b'_'
326 | b'`'
327 | b'a'..=b'z'
328 | b'|'
329 | b'~'
330 )
331 })
332}
333
334impl Default for ResponseWriter {
335 fn default() -> Self {
336 Self::new()
337 }
338}
339
340#[cfg(test)]
341mod tests {
342 use super::*;
343 use asupersync::stream::iter;
344 use std::task::Waker;
345
346 fn noop_waker() -> Waker {
347 Waker::noop().clone()
348 }
349
350 fn collect_stream<S: Stream<Item = Vec<u8>> + Unpin>(mut stream: S) -> Vec<u8> {
351 let waker = noop_waker();
352 let mut cx = Context::from_waker(&waker);
353 let mut out = Vec::new();
354
355 loop {
356 match Pin::new(&mut stream).poll_next(&mut cx) {
357 Poll::Ready(Some(chunk)) => out.extend_from_slice(&chunk),
358 Poll::Ready(None) => break,
359 Poll::Pending => panic!("unexpected pending stream"),
360 }
361 }
362
363 out
364 }
365
366 #[test]
367 fn write_full_sets_content_length() {
368 let response = Response::ok()
369 .header("content-type", b"text/plain".to_vec())
370 .body(ResponseBody::Bytes(b"hello".to_vec()));
371 let mut writer = ResponseWriter::new();
372 let bytes = match writer.write(response) {
373 ResponseWrite::Full(bytes) => bytes,
374 ResponseWrite::Stream(_) => panic!("expected full response"),
375 };
376 let text = String::from_utf8_lossy(&bytes);
377 assert!(text.starts_with("HTTP/1.1 200 OK\r\n"));
378 assert!(text.contains("content-length: 5\r\n"));
379 assert!(text.contains("\r\n\r\nhello"));
380 }
381
382 #[test]
383 fn write_stream_uses_chunked_encoding() {
384 let stream = iter(vec![b"hello".to_vec(), b"world".to_vec()]);
385 let response = Response::ok()
386 .header("content-type", b"text/plain".to_vec())
387 .body(ResponseBody::stream(stream));
388 let mut writer = ResponseWriter::new();
389 let bytes = match writer.write(response) {
390 ResponseWrite::Stream(stream) => collect_stream(stream),
391 ResponseWrite::Full(_) => panic!("expected stream response"),
392 };
393
394 let expected = b"HTTP/1.1 200 OK\r\ncontent-type: text/plain\r\ntransfer-encoding: chunked\r\n\r\n5\r\nhello\r\n5\r\nworld\r\n0\r\n\r\n";
395 assert_eq!(bytes, expected);
396 }
397
398 #[test]
403 fn trailers_empty() {
404 let t = Trailers::new();
405 assert!(t.is_empty());
406 assert_eq!(t.trailer_header_value(), "");
407 }
408
409 #[test]
410 fn trailers_encode() {
411 let t = Trailers::new()
412 .add("Content-MD5", "abc123")
413 .add("Server-Timing", "total;dur=50");
414 assert!(!t.is_empty());
415 assert_eq!(t.trailer_header_value(), "Content-MD5, Server-Timing");
416 let encoded = t.encode();
417 let s = std::str::from_utf8(&encoded).unwrap();
418 assert!(s.contains("Content-MD5: abc123\r\n"));
419 assert!(s.contains("Server-Timing: total;dur=50\r\n"));
420 }
421
422 #[test]
423 fn chunked_encoder_with_trailers() {
424 let stream = iter(vec![b"data".to_vec()]);
425 let body = Box::pin(stream) as BodyStream;
426 let head = b"HTTP/1.1 200 OK\r\n\r\n".to_vec();
427 let trailers = Trailers::new().add("Checksum", "deadbeef");
428 let encoder = ChunkedEncoder::new(head, body).with_trailers(trailers);
429 let bytes = collect_stream(encoder);
430 let s = std::str::from_utf8(&bytes).unwrap();
431 assert!(s.contains("0\r\nChecksum: deadbeef\r\n\r\n"));
433 }
434
435 #[test]
436 fn chunked_encoder_without_trailers_unchanged() {
437 let stream = iter(vec![b"hi".to_vec()]);
438 let body = Box::pin(stream) as BodyStream;
439 let head = b"HTTP/1.1 200 OK\r\n\r\n".to_vec();
440 let encoder = ChunkedEncoder::new(head, body);
441 let bytes = collect_stream(encoder);
442 assert!(bytes.ends_with(b"0\r\n\r\n"));
443 }
444
445 #[test]
446 fn final_chunk_format_with_multiple_trailers() {
447 let t = Trailers::new()
448 .add("Digest", "sha-256=abc")
449 .add("Signature", "sig123");
450 let encoder = ChunkedEncoder {
451 head: None,
452 body: Box::pin(iter(Vec::<Vec<u8>>::new())),
453 finished: false,
454 trailers: Some(t),
455 };
456 let final_chunk = encoder.encode_final_chunk();
457 let s = std::str::from_utf8(&final_chunk).unwrap();
458 assert_eq!(s, "0\r\nDigest: sha-256=abc\r\nSignature: sig123\r\n\r\n");
459 }
460
461 #[test]
462 fn write_full_drops_invalid_header_names_and_sanitizes_values() {
463 let mut writer = ResponseWriter::new();
464 let headers = vec![
465 ("x-ok".to_string(), b"safe".to_vec()),
466 ("bad\r\nname".to_string(), b"ignored".to_vec()),
467 ("x-test".to_string(), b"hello\r\nx-injected: yes".to_vec()),
468 ];
469
470 let bytes = writer.write_full(StatusCode::OK, &headers, b"body");
471 let text = String::from_utf8_lossy(&bytes);
472
473 assert!(text.contains("x-ok: safe\r\n"));
474 assert!(!text.contains("bad\r\nname:"));
475 assert!(text.contains("x-test: hellox-injected: yes\r\n"));
476 assert!(!text.contains("\r\nx-injected: yes\r\n"));
477 }
478
479 #[test]
480 fn write_stream_head_drops_invalid_header_names_and_sanitizes_values() {
481 let mut writer = ResponseWriter::new();
482 let headers = vec![
483 ("content-type".to_string(), b"text/plain".to_vec()),
484 ("bad\nname".to_string(), b"ignored".to_vec()),
485 ("x-test".to_string(), b"hello\r\nx-injected: yes".to_vec()),
486 ];
487
488 let bytes = writer.write_stream_head(StatusCode::OK, &headers);
489 let text = String::from_utf8_lossy(&bytes);
490
491 assert!(text.contains("content-type: text/plain\r\n"));
492 assert!(!text.contains("bad\nname:"));
493 assert!(text.contains("x-test: hellox-injected: yes\r\n"));
494 assert!(!text.contains("\r\nx-injected: yes\r\n"));
495 }
496
497 #[test]
498 fn trailers_encode_drops_invalid_names_and_sanitizes_values() {
499 let encoded = Trailers::new()
500 .add("Checksum", "abc123")
501 .add("Bad\r\nName", "ignored")
502 .add("Signature", "sig\r\nInjected: yes")
503 .encode();
504 let text = std::str::from_utf8(&encoded).unwrap();
505
506 assert!(text.contains("Checksum: abc123\r\n"));
507 assert!(!text.contains("Bad\r\nName"));
508 assert!(text.contains("Signature: sigInjected: yes\r\n"));
509 assert!(!text.contains("\r\nInjected: yes\r\n"));
510 }
511}