Skip to main content

goose_http/encode/
mod.rs

1//! HTTP response serialization utilities.
2//!
3//! Responsible for emitting response lines, headers, and bodies according to
4//! RFC 9112 framing rules, including chunked transfer coding and trailers.
5
6use std::fmt::Write;
7
8use futures_util::StreamExt;
9use tokio::io::{AsyncWrite, AsyncWriteExt};
10
11use crate::{
12    common::{Method, StatusCode},
13    date,
14    headers::{Headers, header_keys},
15    response::{BoxBodyStream, Response, ResponseBody},
16};
17
18/// Directive applied to the `Connection` header for a response.
19#[derive(Debug, Clone, Copy, PartialEq, Eq)]
20pub enum ConnectionDirective {
21    /// Explicitly instruct the client to close the TCP connection.
22    Close,
23    /// Leave the connection open for further requests (default for HTTP/1.1).
24    KeepAlive,
25}
26
27/// Errors produced while serialising a response.
28#[derive(Debug, thiserror::Error)]
29pub enum EncodeError {
30    #[error(transparent)]
31    Io(#[from] std::io::Error),
32    #[error("response contains both Content-Length and Transfer-Encoding")]
33    ConflictingLengthAndTransferEncoding,
34}
35
36/// Writer responsible for serialising responses onto the wire.
37pub struct ResponseWriter<'a, W>
38where
39    W: AsyncWrite + Unpin,
40{
41    writer: &'a mut W,
42}
43
44impl<'a, W> ResponseWriter<'a, W>
45where
46    W: AsyncWrite + Unpin,
47{
48    /// Create a new writer referencing the underlying transport.
49    pub fn new(writer: &'a mut W) -> Self {
50        Self { writer }
51    }
52
53    /// Emit an interim `100 Continue` response.
54    pub async fn write_continue(&mut self) -> Result<(), EncodeError> {
55        self.writer
56            .write_all(b"HTTP/1.1 100 Continue\r\n\r\n")
57            .await?;
58        Ok(())
59    }
60
61    /// Serialise the supplied response according to RFC 9112 framing rules.
62    pub async fn write_response(
63        &mut self,
64        response: &mut Response,
65        request_method: &Method,
66        directive: ConnectionDirective,
67    ) -> Result<(), EncodeError> {
68        let version = response.version();
69        let status = response.status();
70        let status_code = status.as_u16();
71        let reason = response.reason_phrase().to_owned();
72
73        let mut body = response.take_body();
74        let mut trailers = response.take_trailers();
75        let headers = response.headers_mut();
76
77        let has_content_length = headers.contains(header_keys::CONTENT_LENGTH);
78        let has_transfer_encoding = headers.contains(header_keys::TRANSFER_ENCODING);
79        if has_content_length && has_transfer_encoding {
80            return Err(EncodeError::ConflictingLengthAndTransferEncoding);
81        }
82
83        if !headers.contains(header_keys::DATE) {
84            headers.insert(header_keys::DATE, date::now());
85        }
86
87        apply_connection_directive(headers, directive);
88
89        let forbid_content_length = matches!(status_code, 100..=199 | 204 | 304);
90        let body_allowed = response_allows_body(status, request_method);
91
92        if !body_allowed {
93            if matches!(request_method, Method::Head) {
94                if let ResponseBody::Full(ref bytes) = body {
95                    if !forbid_content_length && !headers.contains(header_keys::CONTENT_LENGTH) {
96                        headers.insert(header_keys::CONTENT_LENGTH, bytes.len().to_string());
97                    }
98                }
99            }
100            headers.remove(header_keys::TRANSFER_ENCODING);
101            if forbid_content_length {
102                headers.remove(header_keys::CONTENT_LENGTH);
103            } else if !headers.contains(header_keys::CONTENT_LENGTH)
104                && !matches!(request_method, Method::Head)
105            {
106                headers.insert(header_keys::CONTENT_LENGTH, "0");
107            }
108            trailers = None;
109            body = ResponseBody::Empty;
110        } else {
111            match &body {
112                ResponseBody::Empty => {
113                    headers.remove(header_keys::TRANSFER_ENCODING);
114                    if forbid_content_length {
115                        headers.remove(header_keys::CONTENT_LENGTH);
116                    } else if !headers.contains(header_keys::CONTENT_LENGTH) {
117                        headers.insert(header_keys::CONTENT_LENGTH, "0");
118                    }
119                    trailers = None;
120                }
121                ResponseBody::Full(bytes) => {
122                    headers.remove(header_keys::TRANSFER_ENCODING);
123                    if forbid_content_length {
124                        headers.remove(header_keys::CONTENT_LENGTH);
125                    } else if !headers.contains(header_keys::CONTENT_LENGTH) {
126                        headers.insert(header_keys::CONTENT_LENGTH, bytes.len().to_string());
127                    }
128                    trailers = None;
129                }
130                ResponseBody::Stream(_) => {
131                    headers.remove(header_keys::CONTENT_LENGTH);
132                    ensure_chunked_encoding(headers);
133                }
134            }
135        }
136
137        let mut head = String::new();
138        head.push_str(&format!("{} {}", version, status_code));
139        if !reason.is_empty() {
140            head.push(' ');
141            head.push_str(&reason);
142        }
143        head.push_str("\r\n");
144
145        append_headers(&mut head, headers);
146
147        self.writer.write_all(head.as_bytes()).await?;
148
149        match body {
150            ResponseBody::Empty => {}
151            ResponseBody::Full(bytes) => {
152                if body_allowed {
153                    self.writer.write_all(bytes.as_ref()).await?;
154                }
155            }
156            ResponseBody::Stream(mut stream) => {
157                write_chunked_body(self.writer, &mut stream, trailers.as_ref()).await?;
158            }
159        }
160
161        Ok(())
162    }
163
164    /// Flush the underlying writer.
165    pub async fn flush(&mut self) -> Result<(), EncodeError> {
166        self.writer.flush().await.map_err(EncodeError::from)
167    }
168}
169
170fn apply_connection_directive(headers: &mut Headers, directive: ConnectionDirective) {
171    match directive {
172        ConnectionDirective::Close => headers.insert(header_keys::CONNECTION, "close"),
173        ConnectionDirective::KeepAlive => {}
174    }
175}
176
177fn response_allows_body(status: StatusCode, method: &Method) -> bool {
178    let code = status.as_u16();
179    if matches!(code, 100..=199 | 204 | 304) {
180        return false;
181    }
182    !matches!(method, Method::Head)
183}
184
185fn append_headers(buffer: &mut String, headers: &Headers) {
186    for (name, values) in headers.iter() {
187        for value in values {
188            buffer.push_str(name.as_str());
189            buffer.push_str(": ");
190            buffer.push_str(value);
191            buffer.push_str("\r\n");
192        }
193    }
194    buffer.push_str("\r\n");
195}
196
197fn ensure_chunked_encoding(headers: &mut Headers) {
198    let existing = headers
199        .get(header_keys::TRANSFER_ENCODING)
200        .map(|v| v.to_string());
201    if let Some(value) = existing {
202        if value
203            .split(',')
204            .any(|token| token.trim().eq_ignore_ascii_case("chunked"))
205        {
206            return;
207        }
208    }
209    headers.insert(header_keys::TRANSFER_ENCODING, "chunked");
210}
211
212async fn write_chunked_body<W>(
213    writer: &mut W,
214    stream: &mut BoxBodyStream,
215    trailers: Option<&Headers>,
216) -> Result<(), EncodeError>
217where
218    W: AsyncWrite + Unpin,
219{
220    while let Some(chunk) = stream.next().await {
221        let chunk = chunk?;
222        if chunk.is_empty() {
223            continue;
224        }
225        let mut prefix = String::new();
226        write!(&mut prefix, "{:X}\r\n", chunk.len()).unwrap();
227        writer.write_all(prefix.as_bytes()).await?;
228        writer.write_all(chunk.as_ref()).await?;
229        writer.write_all(b"\r\n").await?;
230    }
231
232    writer.write_all(b"0\r\n").await?;
233    if let Some(trailers) = trailers {
234        let mut block = String::new();
235        append_headers(&mut block, trailers);
236        writer.write_all(block.as_bytes()).await?;
237    } else {
238        writer.write_all(b"\r\n").await?;
239    }
240
241    Ok(())
242}