use std::{cell::Cell, cell::RefCell};
use bitflags::bitflags;
use crate::client::ClientRawRequest;
use crate::codec::{Decoder, Encoder};
use crate::http::HttpServiceConfig;
use crate::http::error::{DecodeError, EncodeError, PayloadError};
use crate::http::h1::{
Message, MessageType, PayloadDecoder, PayloadItem, PayloadType, decoder, encoder,
};
use crate::http::{
ConnectionType, Method, RequestHead, ResponseHead, StatusCode, Version, body::BodySize,
};
use crate::service::cfg::Cfg;
use crate::util::{BytePages, Bytes, BytesMut};
bitflags! {
#[derive(Copy, Clone, Debug, Eq, PartialEq, Ord, PartialOrd, Hash)]
struct Flags: u8 {
const HEAD = 0b0000_0001;
const CONNECT = 0b0000_0010;
const KEEPALIVE_ENABLED = 0b0000_1000;
const STREAM = 0b0001_0000;
}
}
#[derive(Debug)]
pub(crate) struct ClientCodec {
inner: ClientCodecInner,
}
#[derive(Debug)]
pub(crate) struct ClientPayloadCodec {
inner: ClientCodecInner,
}
#[derive(Debug)]
struct ClientCodecInner {
decoder: decoder::MessageDecoder<ResponseHead>,
payload: RefCell<Option<PayloadDecoder>>,
version: Cell<Version>,
ctype: Cell<ConnectionType>,
flags: Cell<Flags>,
encoder: encoder::MessageEncoder<RequestHead>,
}
impl ClientCodec {
pub(crate) fn new(keep_alive: bool, cfg: Cfg<HttpServiceConfig>) -> Self {
let flags = if keep_alive {
Flags::KEEPALIVE_ENABLED
} else {
Flags::empty()
};
ClientCodec {
inner: ClientCodecInner {
decoder: decoder::MessageDecoder::new(cfg),
payload: RefCell::new(None),
version: Cell::new(Version::HTTP_11),
ctype: Cell::new(ConnectionType::Close),
flags: Cell::new(flags),
encoder: encoder::MessageEncoder::default(),
},
}
}
pub(crate) fn keepalive(&self) -> bool {
self.inner.ctype.get() == ConnectionType::KeepAlive
}
pub(crate) fn set_close(&self) {
self.inner.ctype.set(ConnectionType::Close);
}
pub(crate) fn message_type(&self) -> MessageType {
if self.inner.flags.get().contains(Flags::STREAM) {
MessageType::Stream
} else if self.inner.payload.borrow().is_none() {
MessageType::None
} else {
MessageType::Payload
}
}
pub(crate) fn into_payload_codec(self) -> ClientPayloadCodec {
ClientPayloadCodec { inner: self.inner }
}
}
impl ClientPayloadCodec {
pub(crate) fn keepalive(&self) -> bool {
self.inner.ctype.get() == ConnectionType::KeepAlive
}
pub(crate) fn eof_delimited(&self) -> bool {
self.inner
.payload
.borrow()
.as_ref()
.is_some_and(PayloadDecoder::is_eof)
}
}
impl Decoder for ClientCodec {
type Item = ResponseHead;
type Error = DecodeError;
fn decode(&self, src: &mut BytesMut) -> Result<Option<Self::Item>, Self::Error> {
debug_assert!(
!self.inner.payload.borrow().is_some(),
"Payload decoder is set"
);
loop {
let Some((req, payload)) = self.inner.decoder.decode(src)? else {
return Ok(None);
};
if req.status.is_informational() && req.status != StatusCode::SWITCHING_PROTOCOLS {
log::trace!("Skipping interim response: {}", req.status);
continue;
}
match req.ctype() {
Some(ConnectionType::KeepAlive) => (),
Some(ctype) => self.inner.ctype.set(ctype),
None if req.version < Version::HTTP_11 => {
self.inner.ctype.set(ConnectionType::Close);
}
None => (),
}
let flags = self.inner.flags.get();
if flags.contains(Flags::CONNECT) && req.status.is_success() {
self.inner.ctype.set(ConnectionType::Close);
*self.inner.payload.borrow_mut() = Some(PayloadDecoder::eof());
self.inner.flags.set(flags | Flags::STREAM);
} else if flags.contains(Flags::HEAD) {
self.inner.payload.borrow_mut().take();
} else {
match payload {
PayloadType::None => {
self.inner.payload.borrow_mut().take();
}
PayloadType::Payload(pl) => *self.inner.payload.borrow_mut() = Some(pl),
PayloadType::Stream(pl) => {
*self.inner.payload.borrow_mut() = Some(pl);
let mut flags = self.inner.flags.get();
flags.insert(Flags::STREAM);
self.inner.flags.set(flags);
}
}
}
return Ok(Some(req));
}
}
}
impl Decoder for ClientPayloadCodec {
type Item = Option<Bytes>;
type Error = PayloadError;
fn decode(&self, src: &mut BytesMut) -> Result<Option<Self::Item>, Self::Error> {
debug_assert!(
self.inner.payload.borrow().is_some(),
"Payload decoder is not specified"
);
let item = self
.inner
.payload
.borrow_mut()
.as_mut()
.unwrap()
.decode(src)?;
Ok(match item {
Some(PayloadItem::Chunk(chunk)) => Some(Some(chunk)),
Some(PayloadItem::Trailers(_)) => return self.decode(src),
Some(PayloadItem::Eof) => {
self.inner.payload.borrow_mut().take();
Some(None)
}
None => None,
})
}
}
impl Encoder for ClientCodec {
type Item = Message<ClientRawRequest>;
type Error = EncodeError;
fn encode(&self, item: Self::Item, dst: &mut BytePages) -> Result<(), Self::Error> {
match item {
Message::Item(mut req) => {
let inner = &self.inner;
inner.version.set(req.head.version);
let mut flags = inner.flags.get();
flags.set(Flags::HEAD, req.head.method == Method::HEAD);
flags.set(Flags::CONNECT, req.head.method == Method::CONNECT);
inner.flags.set(flags);
inner.ctype.set(match req.head.connection_type() {
ConnectionType::KeepAlive => {
if inner.flags.get().contains(Flags::KEEPALIVE_ENABLED) {
ConnectionType::KeepAlive
} else {
ConnectionType::Close
}
}
ConnectionType::Upgrade => ConnectionType::Upgrade,
ConnectionType::Close => ConnectionType::Close,
});
let size = if req.size == BodySize::None
&& matches!(req.head.method, Method::POST | Method::PUT | Method::PATCH)
{
BodySize::Empty
} else {
req.size
};
let headers = req.headers.take();
inner.encoder.encode(
dst,
&req.head,
false,
false,
inner.version.get(),
size,
inner.ctype.get(),
headers,
)?;
}
Message::Chunk(Some(bytes)) => {
self.inner.encoder.encode_chunk(bytes, dst);
}
Message::Chunk(None) => {
self.inner.encoder.encode_eof(dst)?;
}
}
Ok(())
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::service::cfg::SharedCfg;
#[test]
fn test_skip_interim_responses() {
let cfg: SharedCfg = SharedCfg::new("DBG").add(HttpServiceConfig::new()).into();
let codec = ClientCodec::new(true, cfg.get());
let mut buf = BytesMut::from(
"HTTP/1.1 100 Continue\r\n\r\n\
HTTP/1.1 103 Early Hints\r\nlink: </style.css>\r\n\r\n\
HTTP/1.1 200 OK\r\ncontent-length: 2\r\n\r\nok",
);
let head = codec.decode(&mut buf).unwrap().unwrap();
assert_eq!(head.status, StatusCode::OK);
assert!(!head.headers.contains_key("link"));
assert_eq!(codec.message_type(), MessageType::Payload);
assert_eq!(&buf[..], b"ok");
let codec = ClientCodec::new(true, cfg.get());
let mut buf = BytesMut::from("HTTP/1.1 103 Early Hints\r\n\r\nHTTP/1.1 200");
assert!(codec.decode(&mut buf).unwrap().is_none());
buf.extend_from_slice(b" OK\r\ncontent-length: 0\r\n\r\n");
let head = codec.decode(&mut buf).unwrap().unwrap();
assert_eq!(head.status, StatusCode::OK);
let codec = ClientCodec::new(true, cfg.get());
let mut buf = BytesMut::from("HTTP/1.1 101 Switching Protocols\r\n\r\n");
let head = codec.decode(&mut buf).unwrap().unwrap();
assert_eq!(head.status, StatusCode::SWITCHING_PROTOCOLS);
}
#[test]
fn test_http10_response_keepalive() {
let cfg: SharedCfg = SharedCfg::new("DBG").add(HttpServiceConfig::new()).into();
for (resp, keepalive) in [
("HTTP/1.0 200 OK\r\ncontent-length: 0\r\n\r\n", false),
(
"HTTP/1.0 200 OK\r\nconnection: keep-alive\r\ncontent-length: 0\r\n\r\n",
true,
),
("HTTP/1.1 200 OK\r\ncontent-length: 0\r\n\r\n", true),
(
"HTTP/1.1 200 OK\r\nconnection: close\r\ncontent-length: 0\r\n\r\n",
false,
),
] {
let codec = ClientCodec::new(true, cfg.get());
codec.inner.ctype.set(ConnectionType::KeepAlive);
let mut buf = BytesMut::from(resp);
codec.decode(&mut buf).unwrap().unwrap();
assert_eq!(codec.keepalive(), keepalive, "{resp:?}");
}
}
#[crate::rt_test]
async fn test_connect_response_is_tunnel() {
let cfg: SharedCfg = SharedCfg::new("DBG").add(HttpServiceConfig::new()).into();
let connect = |codec: &ClientCodec| {
let mut head = crate::http::Message::<RequestHead>::new();
head.method = Method::CONNECT;
head.uri = urly::Url::from_static("http://example.com:443");
let req = ClientRawRequest {
head,
headers: None,
size: crate::http::body::BodySize::None,
};
codec
.encode(Message::Item(req), &mut BytePages::default())
.unwrap();
};
for resp in [
"HTTP/1.1 200 OK\r\ncontent-length: 0\r\n\r\ntunnel",
"HTTP/1.1 200 OK\r\ntransfer-encoding: chunked\r\n\r\ntunnel",
"HTTP/1.1 200 OK\r\n\r\ntunnel",
] {
let codec = ClientCodec::new(true, cfg.get());
connect(&codec);
let mut buf = BytesMut::from(resp);
let head = codec.decode(&mut buf).unwrap().unwrap();
assert_eq!(head.status, StatusCode::OK);
assert_eq!(codec.message_type(), MessageType::Stream, "{resp:?}");
assert!(!codec.keepalive(), "{resp:?}");
let codec = codec.into_payload_codec();
assert!(codec.eof_delimited());
assert_eq!(codec.decode(&mut buf).unwrap(), Some(Some("tunnel".into())));
}
let codec = ClientCodec::new(true, cfg.get());
connect(&codec);
let mut buf = BytesMut::from(
"HTTP/1.1 407 Proxy Authentication Required\r\ncontent-length: 2\r\n\r\nno",
);
let head = codec.decode(&mut buf).unwrap().unwrap();
assert_eq!(head.status, StatusCode::PROXY_AUTHENTICATION_REQUIRED);
assert_eq!(codec.message_type(), MessageType::Payload);
assert!(codec.keepalive());
}
#[crate::rt_test]
async fn test_empty_request_content_length() {
let cfg: SharedCfg = SharedCfg::new("DBG").add(HttpServiceConfig::new()).into();
for (method, size, cl) in [
(Method::POST, BodySize::None, Some("0")),
(Method::PUT, BodySize::None, Some("0")),
(Method::PATCH, BodySize::None, Some("0")),
(Method::POST, BodySize::Sized(3), Some("3")),
(Method::GET, BodySize::None, None),
(Method::DELETE, BodySize::None, None),
(Method::GET, BodySize::Empty, Some("0")),
] {
let codec = ClientCodec::new(true, cfg.get());
let mut head = crate::http::Message::<RequestHead>::new();
head.method = method.clone();
head.uri = urly::Url::from_static("/");
let req = ClientRawRequest {
head,
headers: None,
size,
};
let mut buf = BytePages::default();
codec.encode(Message::Item(req), &mut buf).unwrap();
let data = String::from_utf8(buf.take().unwrap().to_vec()).unwrap();
let expected = cl.map(|cl| format!("content-length: {cl}\r\n"));
match expected {
Some(line) => assert!(data.contains(&line), "{method} {size:?}: {data:?}"),
None => assert!(!data.contains("content-length"), "{method}: {data:?}"),
}
}
}
}