1use 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#[derive(Debug, Clone, Copy, PartialEq, Eq)]
20pub enum ConnectionDirective {
21 Close,
23 KeepAlive,
25}
26
27#[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
36pub 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 pub fn new(writer: &'a mut W) -> Self {
50 Self { writer }
51 }
52
53 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 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 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}