1use std::{cell::Cell, fmt};
2
3use bitflags::bitflags;
4
5use crate::codec::{Decoder, Encoder};
6use crate::http::body::BodySize;
7use crate::http::config::{DateService, HttpServiceConfig};
8use crate::http::error::{DecodeError, EncodeError};
9use crate::http::message::ConnectionType;
10use crate::http::{Method, StatusCode, Version, request::Request, response::Response};
11use crate::{Cfg, util::BytePages, util::BytesMut};
12
13use super::{Message, decoder, decoder::PayloadType, encoder};
14
15bitflags! {
16 #[derive(Copy, Clone, Debug, Eq, PartialEq, Ord, PartialOrd, Hash)]
17 struct Flags: u8 {
18 const HEAD = 0b0000_0001;
19 const STREAM = 0b0000_0010;
20 const KEEPALIVE_ENABLED = 0b0000_0100;
21 }
22}
23
24pub struct Codec {
57 con_id: usize,
58 decoder: decoder::MessageDecoder<Request>,
59 version: Cell<Version>,
60 ctype: Cell<ConnectionType>,
61 pub(super) cfg: Cfg<HttpServiceConfig>,
62
63 flags: Cell<Flags>,
65 encoder: encoder::MessageEncoder<Response<()>>,
66}
67
68impl Clone for Codec {
69 fn clone(&self) -> Self {
70 Codec {
71 con_id: self.con_id,
72 decoder: self.decoder.clone(),
73 version: self.version.clone(),
74 cfg: self.cfg.clone(),
75 ctype: self.ctype.clone(),
76 flags: self.flags.clone(),
77 encoder: self.encoder.clone(),
78 }
79 }
80}
81
82impl fmt::Debug for Codec {
83 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
84 f.debug_struct("h1::Codec")
85 .field("con_id", &self.con_id)
86 .field("version", &self.version)
87 .field("flags", &self.flags)
88 .field("ctype", &self.ctype)
89 .field("encoder", &self.encoder)
90 .field("decoder", &self.decoder)
91 .finish()
92 }
93}
94
95impl Codec {
96 pub fn new(con_id: usize, cfg: Cfg<HttpServiceConfig>) -> Self {
101 let flags = if cfg.ka_enabled {
102 Flags::KEEPALIVE_ENABLED
103 } else {
104 Flags::empty()
105 };
106 let ctype = if cfg.ka_enabled {
107 ConnectionType::KeepAlive
108 } else {
109 ConnectionType::Close
110 };
111 let decoder = decoder::MessageDecoder::new(cfg.clone());
112
113 Codec {
114 cfg,
115 con_id,
116 decoder,
117 flags: Cell::new(flags),
118 version: Cell::new(Version::HTTP_11),
119 ctype: Cell::new(ctype),
120 encoder: encoder::MessageEncoder::default(),
121 }
122 }
123
124 pub(super) fn is_reading_hdrs(&self) -> bool {
125 self.decoder.is_reading_hdrs()
126 }
127
128 pub(super) fn is_body_complete(&self) -> bool {
130 self.encoder.is_body_complete()
131 }
132
133 #[inline]
134 pub fn keepalive(&self) -> bool {
140 self.ctype.get() == ConnectionType::KeepAlive
141 }
142
143 #[inline]
144 #[doc(hidden)]
145 pub fn set_date_header(&self, dst: &mut BytesMut) {
146 DateService.set_date_header(dst);
147 }
148
149 fn insert_flags(&self, f: Flags) {
150 let mut flags = self.flags.get();
151 flags.insert(f);
152 self.flags.set(flags);
153 }
154
155 pub(super) fn reset_upgrade(&self) {
156 let mut flags = self.flags.get();
157 flags.remove(Flags::STREAM);
158 self.flags.set(flags);
159 self.ctype.set(ConnectionType::Close);
160 }
161}
162
163impl Decoder for Codec {
164 type Item = (Request, PayloadType);
165 type Error = DecodeError;
166
167 fn decode(&self, src: &mut BytesMut) -> Result<Option<Self::Item>, Self::Error> {
168 if let Some((mut req, payload)) = self.decoder.decode(src)? {
169 let head = req.head_mut();
170 head.id = self.con_id;
171 let mut flags = self.flags.get();
172 flags.set(Flags::HEAD, head.method == Method::HEAD);
173 self.flags.set(flags);
174 self.version.set(head.version);
175
176 let ctype = head.connection_type();
177 if ctype == ConnectionType::KeepAlive && !flags.contains(Flags::KEEPALIVE_ENABLED) {
178 self.ctype.set(ConnectionType::Close);
179 } else {
180 self.ctype.set(ctype);
181 }
182
183 if let PayloadType::Stream(_) = payload {
184 self.insert_flags(Flags::STREAM);
185 }
186 Ok(Some((req, payload)))
187 } else {
188 Ok(None)
189 }
190 }
191}
192
193impl Encoder for Codec {
194 type Item = Message<(Response<()>, BodySize)>;
195 type Error = EncodeError;
196
197 fn encode(&self, item: Self::Item, dst: &mut BytePages) -> Result<(), Self::Error> {
198 match item {
199 Message::Item((mut res, length)) => {
200 res.head_mut().version = self.version.get();
202
203 if res.status() == StatusCode::SWITCHING_PROTOCOLS {
205 self.ctype.set(ConnectionType::Upgrade);
206 } else if let Some(ct) = res.head().ctype()
207 && ct != ConnectionType::KeepAlive
208 {
209 self.ctype.set(ct);
210 }
211
212 let ctype = self.encoder.encode(
214 dst,
215 &res,
216 self.flags.get().contains(Flags::HEAD),
217 self.flags.get().contains(Flags::STREAM),
218 self.version.get(),
219 length,
220 self.ctype.get(),
221 None,
222 )?;
223 self.ctype.set(ctype);
224 }
225 Message::Chunk(Some(bytes)) => {
226 self.encoder.encode_chunk(bytes, dst);
227 }
228 Message::Chunk(None) => {
229 self.encoder.encode_eof(dst)?;
230 }
231 }
232 Ok(())
233 }
234}
235
236#[cfg(test)]
237mod tests {
238 use super::*;
239 use crate::{
240 SharedCfg,
241 http::{HttpMessage, KeepAlive, h1::PayloadItem},
242 util::Bytes,
243 };
244
245 #[crate::rt_test]
247 async fn test_unknown_status_reason() {
248 use crate::http::StatusCode;
249
250 let cfg: SharedCfg = SharedCfg::new("DBG").add(HttpServiceConfig::new()).into();
251 let codec = Codec::new(0, cfg.get());
252 let mut buf = BytesMut::from("GET / HTTP/1.1\r\nhost: localhost\r\n\r\n");
253 codec.decode(&mut buf).unwrap().unwrap();
254
255 let status = StatusCode::from_u16(599).unwrap();
256 let res = Response::with_body(status, ());
257 assert_eq!(res.head().reason(), "");
258 let mut out = BytePages::default();
259 codec
260 .encode(Message::Item((res, BodySize::Empty)), &mut out)
261 .unwrap();
262 let data = out.take().unwrap();
263 assert!(data.starts_with(b"HTTP/1.1 599 \r\n"), "{data:?}");
264
265 let mut res = Response::with_body(status, ());
266 res.head_mut().reason = Some("Custom");
267 let mut out = BytePages::default();
268 codec
269 .encode(Message::Item((res, BodySize::Empty)), &mut out)
270 .unwrap();
271 let data = out.take().unwrap();
272 assert!(data.starts_with(b"HTTP/1.1 599 Custom\r\n"), "{data:?}");
273 }
274
275 #[crate::rt_test]
277 async fn test_bodyless_status_has_no_body() {
278 use crate::http::StatusCode;
279
280 let cfg: SharedCfg = SharedCfg::new("DBG").add(HttpServiceConfig::new()).into();
281 for status in [
282 StatusCode::CONTINUE,
283 StatusCode::from_u16(103).unwrap(),
284 StatusCode::NO_CONTENT,
285 StatusCode::NOT_MODIFIED,
286 ] {
287 for size in [BodySize::Sized(3), BodySize::Stream] {
288 let codec = Codec::new(0, cfg.get());
289 let mut buf = BytesMut::from("GET / HTTP/1.1\r\nhost: localhost\r\n\r\n");
290 codec.decode(&mut buf).unwrap().unwrap();
291
292 let mut out = BytePages::default();
293 let res = Response::with_body(status, ());
294 codec.encode(Message::Item((res, size)), &mut out).unwrap();
295 codec
296 .encode(Message::Chunk(Some(Bytes::from_static(b"abc"))), &mut out)
297 .unwrap();
298 codec.encode(Message::Chunk(None), &mut out).unwrap();
299
300 let mut data = Vec::new();
301 while let Some(chunk) = out.take() {
302 data.extend_from_slice(&chunk);
303 }
304 let data = String::from_utf8(data).unwrap();
305 assert!(data.ends_with("\r\n\r\n"), "{status} {size:?}: {data:?}");
306 assert!(
307 !data.contains("content-length"),
308 "{status} {size:?}: {data:?}"
309 );
310 assert!(
311 !data.contains("transfer-encoding"),
312 "{status} {size:?}: {data:?}"
313 );
314 }
315 }
316 }
317
318 #[crate::rt_test]
320 async fn test_switching_protocols_ends_http1() {
321 let cfg: SharedCfg = SharedCfg::new("DBG").add(HttpServiceConfig::new()).into();
322 let codec = Codec::new(0, cfg.get());
323 let mut buf = BytesMut::from("GET / HTTP/1.1\r\nhost: a\r\n\r\n");
324 codec.decode(&mut buf).unwrap().unwrap();
325 assert!(codec.keepalive());
326
327 let mut out = BytePages::default();
329 codec
330 .encode(
331 Message::Item((
332 Response::with_body(StatusCode::SWITCHING_PROTOCOLS, ()),
333 BodySize::None,
334 )),
335 &mut out,
336 )
337 .unwrap();
338 let mut data = Vec::new();
339 while let Some(chunk) = out.take() {
340 data.extend_from_slice(&chunk);
341 }
342 let data = String::from_utf8(data).unwrap();
343 assert!(data.contains("connection: upgrade\r\n"), "{data:?}");
344 assert!(!codec.keepalive());
345 }
346
347 #[crate::rt_test]
349 async fn test_switching_protocols_body_is_not_framed() {
350 let cfg: SharedCfg = SharedCfg::new("DBG").add(HttpServiceConfig::new()).into();
351 for size in [BodySize::Sized(3), BodySize::Stream] {
352 let codec = Codec::new(0, cfg.get());
354 let mut buf = BytesMut::from("GET / HTTP/1.1\r\nhost: a\r\n\r\n");
355 codec.decode(&mut buf).unwrap().unwrap();
356
357 let mut out = BytePages::default();
358 let res = Response::with_body(StatusCode::SWITCHING_PROTOCOLS, ());
359 codec.encode(Message::Item((res, size)), &mut out).unwrap();
360 for chunk in [&b"abc"[..], b"defg"] {
361 codec
362 .encode(
363 Message::Chunk(Some(Bytes::copy_from_slice(chunk))),
364 &mut out,
365 )
366 .unwrap();
367 }
368 codec.encode(Message::Chunk(None), &mut out).unwrap();
369
370 let mut data = Vec::new();
371 while let Some(chunk) = out.take() {
372 data.extend_from_slice(&chunk);
373 }
374 let data = String::from_utf8(data).unwrap();
375 assert!(data.ends_with("\r\n\r\nabcdefg"), "{size:?}: {data:?}");
376 assert!(!data.contains("content-length"), "{size:?}: {data:?}");
377 assert!(!data.contains("transfer-encoding"), "{size:?}: {data:?}");
378 }
379 }
380
381 #[crate::rt_test]
382 async fn test_response_without_body_has_length() {
383 use crate::http::{StatusCode, header};
384
385 let encode = |req: &str, res: Response<()>| {
386 let cfg: SharedCfg = SharedCfg::new("DBG").add(HttpServiceConfig::new()).into();
387 let codec = Codec::new(0, cfg.get());
388 let mut buf = BytesMut::from(req);
389 codec.decode(&mut buf).unwrap().unwrap();
390
391 let mut out = BytePages::default();
392 codec
393 .encode(Message::Item((res, BodySize::None)), &mut out)
394 .unwrap();
395 let mut data = Vec::new();
396 while let Some(chunk) = out.take() {
397 data.extend_from_slice(&chunk);
398 }
399 (String::from_utf8(data).unwrap(), codec.keepalive())
400 };
401 let get = "GET / HTTP/1.1\r\nhost: a\r\n\r\n";
402
403 let (data, keepalive) = encode(get, Response::with_body(StatusCode::OK, ()));
404 assert!(data.contains("\r\ncontent-length: 0\r\n"), "{data:?}");
405 assert!(keepalive);
406
407 let mut res = Response::with_body(StatusCode::NOT_FOUND, ());
409 res.headers_mut().insert(
410 header::CONTENT_LENGTH,
411 header::HeaderValue::from_static("10"),
412 );
413 let (data, _) = encode(get, res);
414 assert_eq!(data.matches("content-length").count(), 1, "{data:?}");
415 assert!(data.contains("\r\ncontent-length: 0\r\n"), "{data:?}");
416
417 let (data, keepalive) = encode(
418 "GET / HTTP/1.0\r\nconnection: keep-alive\r\n\r\n",
419 Response::with_body(StatusCode::OK, ()),
420 );
421 assert!(data.contains("\r\ncontent-length: 0\r\n"), "{data:?}");
422 assert!(data.contains("connection: keep-alive\r\n"), "{data:?}");
423 assert!(keepalive);
424
425 for (req, status) in [
427 ("HEAD / HTTP/1.1\r\nhost: a\r\n\r\n", StatusCode::OK),
428 (get, StatusCode::NO_CONTENT),
429 (get, StatusCode::NOT_MODIFIED),
430 (
431 "GET / HTTP/1.1\r\nhost: a\r\nconnection: upgrade\r\nupgrade: websocket\r\n\r\n",
432 StatusCode::SWITCHING_PROTOCOLS,
433 ),
434 ] {
435 let (data, _) = encode(req, Response::with_body(status, ()));
436 assert!(!data.contains("content-length"), "{status} {data:?}");
437 assert!(!data.contains("transfer-encoding"), "{status} {data:?}");
438 }
439 }
440
441 fn encode_stream(req: &str, res: Response<()>) -> (String, bool) {
442 let cfg: SharedCfg = SharedCfg::new("DBG").add(HttpServiceConfig::new()).into();
443 let codec = Codec::new(0, cfg.get());
444 let mut buf = BytesMut::from(req);
445 codec.decode(&mut buf).unwrap().unwrap();
446
447 let mut out = BytePages::default();
448 codec
449 .encode(Message::Item((res, BodySize::Stream)), &mut out)
450 .unwrap();
451 codec
452 .encode(Message::Chunk(Some(Bytes::from_static(b"abc"))), &mut out)
453 .unwrap();
454 codec.encode(Message::Chunk(None), &mut out).unwrap();
455
456 let mut data = Vec::new();
457 while let Some(chunk) = out.take() {
458 data.extend_from_slice(&chunk);
459 }
460 (String::from_utf8(data).unwrap(), codec.keepalive())
461 }
462
463 #[crate::rt_test]
466 async fn test_http10_stream_response_is_not_chunked() {
467 use crate::http::StatusCode;
468
469 let (data, keepalive) = encode_stream(
470 "GET / HTTP/1.0\r\nconnection: keep-alive\r\n\r\n",
471 Response::with_body(StatusCode::OK, ()),
472 );
473 assert!(data.starts_with("HTTP/1.1 200 OK\r\n"), "{data:?}");
474 assert!(!data.contains("transfer-encoding"), "{data:?}");
475 assert!(!data.contains("keep-alive"), "{data:?}");
476 assert!(data.contains("connection: close\r\n"), "{data:?}");
477 assert!(data.ends_with("\r\n\r\nabc"), "{data:?}");
478 assert!(!keepalive);
479
480 let (data, keepalive) = encode_stream(
481 "GET / HTTP/1.1\r\nhost: localhost\r\n\r\n",
482 Response::with_body(StatusCode::OK, ()),
483 );
484 assert!(data.contains("transfer-encoding: chunked\r\n"), "{data:?}");
485 assert!(data.ends_with("3\r\nabc\r\n0\r\n\r\n"), "{data:?}");
486 assert!(keepalive);
487
488 let mut res = Response::with_body(StatusCode::OK, ());
489 res.head_mut().no_chunking(true);
490 let (data, keepalive) = encode_stream("GET / HTTP/1.1\r\nhost: localhost\r\n\r\n", res);
491 assert!(data.contains("connection: close\r\n"), "{data:?}");
492 assert!(data.ends_with("\r\n\r\nabc"), "{data:?}");
493 assert!(!keepalive);
494 }
495
496 #[test]
497 fn test_http_request_chunked_payload_and_next_message() {
498 let cfg: SharedCfg = SharedCfg::new("DBG").add(HttpServiceConfig::new()).into();
499
500 let codec = Codec::new(0, cfg.get());
501 assert!(format!("{codec:?}").contains("h1::Codec"));
502
503 let mut buf = BytesMut::from(
504 "GET /test HTTP/1.1\r\nhost: localhost\r\n\
505 transfer-encoding: chunked\r\n\r\n",
506 );
507 let (req, pl) = codec.decode(&mut buf).unwrap().unwrap();
508 let PayloadType::Payload(pl) = pl else { panic!() };
509
510 assert_eq!(req.method(), Method::GET);
511 assert!(req.chunked().unwrap());
512
513 buf.extend(
514 b"4\r\ndata\r\n4\r\nline\r\n0\r\n\r\n\
515 POST /test2 HTTP/1.1\r\nhost: localhost\r\n\
516 transfer-encoding: chunked\r\n\r\n"
517 .iter(),
518 );
519
520 let msg = pl.decode(&mut buf).unwrap().unwrap();
522 assert_eq!(msg, PayloadItem::Chunk(Bytes::from_static(b"dataline")));
523
524 let msg = pl.decode(&mut buf).unwrap().unwrap();
525 assert_eq!(msg, PayloadItem::Eof);
526
527 let (req, _pl) = codec.decode(&mut buf).unwrap().unwrap();
529 assert_eq!(*req.method(), Method::POST);
530 assert!(req.chunked().unwrap());
531
532 let codec = Codec::new(0, cfg.get());
533 let mut buf = BytesMut::from(
534 "GET /test HTTP/1.1\r\nhost: localhost\r\n\
535 connection: upgrade\r\nupgrade: websocket\r\n\r\n",
536 );
537 let (req, _) = codec.decode(&mut buf).unwrap().unwrap();
538 assert!(req.upgrade());
539 assert!(!codec.keepalive());
540 codec.reset_upgrade();
541 assert!(!codec.keepalive());
542
543 let codec = Codec::new(0, cfg.get());
545 let mut buf = BytesMut::from(
546 "GET /test HTTP/1.1\r\nhost: localhost\r\n\
547 connection: keep-alive, Upgrade\r\nupgrade: websocket\r\n\r\n",
548 );
549 let (req, _) = codec.decode(&mut buf).unwrap().unwrap();
550 assert!(req.upgrade());
551
552 let codec = Codec::new(0, cfg.get());
553 let mut buf = BytesMut::from("GET /test HTTP/1.1\r\nhost: localhost\r\n\r\n");
554 let (req, _) = codec.decode(&mut buf).unwrap().unwrap();
555 assert!(!req.upgrade());
556
557 let cfg: SharedCfg = SharedCfg::new("DBG")
558 .add(HttpServiceConfig::new().set_keepalive(KeepAlive::Disabled))
559 .into();
560 let codec = Codec::new(0, cfg.get());
561 assert!(!codec.keepalive());
562 }
563}