use std::{cell::Cell, fmt, io, rc::Rc};
use crate::http::message::CurrentIo;
use crate::http::{Request, Response, ResponseError, body::Body, h1::Codec};
use crate::io::{Filter, Io, IoBoxed, IoRef};
pub enum Control<F, Err> {
Connect(Connection<F>),
Request(NewRequest),
Upgrade(Upgrade<F>),
Expect(Expect),
Disconnect(Reason<Err>),
}
#[derive(Debug)]
pub enum Reason<Err> {
Service(ServiceDisconnect),
Error(Error<Err>),
ProtocolError(ProtocolError),
PeerGone(PeerGone),
KeepAlive(KeepAlive),
}
#[derive(Copy, Clone, Debug, PartialEq, Eq)]
pub enum ServiceDisconnectReason {
Shutdown,
UpgradeHandled,
UpgradeFailed,
ExpectFailed,
PayloadDropped,
}
#[derive(Debug)]
pub struct ControlAck<F> {
pub(super) result: ControlResult<F>,
}
#[derive(Debug)]
pub(super) enum ControlResult<F> {
Connect(Io<F>),
Continue(Request),
Expect(Request),
Upgrade(Request),
UpgradeAck(Request),
UpgradeHandled,
Publish(Request),
Response(Response<()>, Body),
Error(Response<()>, Body),
ProtocolError(Response<()>, Body),
UpgradeFailed(Response<()>, Body),
ExpectFailed(Response<()>, Body),
Stop,
}
impl<F, Err> Control<F, Err> {
pub(super) fn connect(id: usize, io: Io<F>) -> Self {
Control::Connect(Connection { id, io })
}
pub(super) fn request(req: Request) -> Self {
Control::Request(NewRequest(req))
}
pub(super) fn upgrade(req: Request, io: Rc<Io<F>>, codec: Codec) -> Self {
Control::Upgrade(Upgrade { req, io, codec })
}
pub(super) fn expect(req: Request) -> Self {
Control::Expect(Expect(req))
}
pub(super) fn err(err: Err) -> Self
where
Err: ResponseError,
{
Control::Disconnect(Reason::Error(Error::new(err)))
}
pub(super) fn peer_gone(err: Option<io::Error>) -> Self {
Control::Disconnect(Reason::PeerGone(PeerGone(err)))
}
pub(super) fn proto_err(err: super::ProtocolError) -> Self {
Control::Disconnect(Reason::ProtocolError(ProtocolError(err)))
}
pub(super) fn keepalive(enabled: bool) -> Self {
Control::Disconnect(Reason::KeepAlive(KeepAlive::new(enabled)))
}
pub(super) fn svc_disconnect(reason: ServiceDisconnectReason) -> Self {
Control::Disconnect(Reason::Service(ServiceDisconnect::new(reason)))
}
#[inline]
pub fn ack(self) -> ControlAck<F>
where
F: Filter,
Err: ResponseError,
{
match self {
Control::Connect(msg) => msg.ack(),
Control::Request(msg) => msg.ack(),
Control::Upgrade(msg) => msg.ack(),
Control::Expect(msg) => msg.ack(),
Control::Disconnect(msg) => msg.ack(),
}
}
}
impl<F, Err> fmt::Debug for Control<F, Err>
where
Err: fmt::Debug,
{
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
Control::Connect(_) => f.debug_tuple("Control::Connect").finish(),
Control::Request(msg) => f.debug_tuple("Control::Request").field(msg).finish(),
Control::Upgrade(msg) => f.debug_tuple("Control::Upgrade").field(msg).finish(),
Control::Expect(msg) => f.debug_tuple("Control::Expect").field(msg).finish(),
Control::Disconnect(msg) => f.debug_tuple("Control::Disconnect").field(msg).finish(),
}
}
}
impl<Err: ResponseError> Reason<Err> {
pub fn ack<F>(self) -> ControlAck<F> {
match self {
Reason::Error(msg) => msg.ack(),
Reason::ProtocolError(msg) => msg.ack(),
Reason::PeerGone(msg) => msg.ack(),
Reason::KeepAlive(msg) => msg.ack(),
Reason::Service(msg) => msg.ack(),
}
}
}
#[derive(Debug)]
pub struct Connection<F> {
id: usize,
io: Io<F>,
}
impl<F> Connection<F> {
#[inline]
pub fn id(&self) -> usize {
self.id
}
#[inline]
pub fn get_ref(&self) -> &Io<F> {
&self.io
}
#[inline]
pub fn get_mut(&mut self) -> &mut Io<F> {
&mut self.io
}
#[inline]
pub fn ack(self) -> ControlAck<F> {
ControlAck {
result: ControlResult::Connect(self.io),
}
}
}
#[derive(Debug)]
pub struct NewRequest(Request);
impl NewRequest {
#[inline]
pub fn get_ref(&self) -> &Request {
&self.0
}
#[inline]
pub fn get_mut(&mut self) -> &mut Request {
&mut self.0
}
#[inline]
pub fn ack<F>(self) -> ControlAck<F> {
let result = if self.0.head().expect() {
ControlResult::Expect(self.0)
} else if self.0.upgrade() {
ControlResult::Upgrade(self.0)
} else {
ControlResult::Publish(self.0)
};
ControlAck { result }
}
#[inline]
pub fn fail<E: ResponseError, F>(self, err: E) -> ControlAck<F> {
let res: Response = (&err).into();
let (res, body) = res.into_parts();
ControlAck {
result: ControlResult::Response(res, body.into()),
}
}
#[inline]
pub fn fail_with<F>(self, res: Response) -> ControlAck<F> {
let (res, body) = res.into_parts();
ControlAck {
result: ControlResult::Response(res, body.into()),
}
}
}
pub struct Upgrade<F> {
req: Request,
io: Rc<Io<F>>,
codec: Codec,
}
struct RequestIoAccess<F> {
io: Rc<Io<F>>,
ioref: IoRef,
codec: Codec,
taken: Cell<bool>,
}
impl<F> fmt::Debug for RequestIoAccess<F> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("RequestIoAccess")
.field("io", &self.ioref)
.field("codec", &self.codec)
.finish()
}
}
impl<F: Filter> crate::http::message::IoAccess for RequestIoAccess<F> {
fn get(&self) -> Option<&IoRef> {
if self.taken.get() { None } else { Some(&self.ioref) }
}
fn take(&self) -> Option<(IoBoxed, Codec)> {
if self.taken.replace(true) {
None
} else {
let io = unsafe { self.io.take() };
Some((io.into(), self.codec.clone()))
}
}
}
impl<F: Filter> Upgrade<F> {
#[inline]
pub fn io(&self) -> &Io<F> {
&self.io
}
#[inline]
pub fn get_ref(&self) -> &Request {
&self.req
}
#[inline]
pub fn get_mut(&mut self) -> &mut Request {
&mut self.req
}
#[inline]
pub fn ack(mut self) -> ControlAck<F> {
let io = Rc::new(RequestIoAccess {
ioref: self.io.get_ref(),
io: self.io,
codec: self.codec,
taken: Cell::new(false),
});
self.req.head_mut().io = CurrentIo::new(io);
ControlAck {
result: ControlResult::UpgradeAck(self.req),
}
}
#[inline]
pub fn handle(self) -> (ControlAck<F>, Io<F>, Request, Codec) {
let io = unsafe { self.io.take() };
(
ControlAck {
result: ControlResult::UpgradeHandled,
},
io,
self.req,
self.codec,
)
}
#[inline]
pub fn fail<E: ResponseError>(self, err: E) -> ControlAck<F> {
let res: Response = (&err).into();
let (res, body) = res.into_parts();
ControlAck {
result: ControlResult::UpgradeFailed(res, body.into()),
}
}
#[inline]
pub fn fail_with(self, res: Response) -> ControlAck<F> {
let (res, body) = res.into_parts();
ControlAck {
result: ControlResult::UpgradeFailed(res, body.into()),
}
}
}
impl<F> fmt::Debug for Upgrade<F> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("Upgrade")
.field("req", &self.req)
.field("io", &self.io)
.field("codec", &self.codec)
.finish()
}
}
#[derive(Debug)]
pub struct ServiceDisconnect(ServiceDisconnectReason);
impl ServiceDisconnect {
fn new(reason: ServiceDisconnectReason) -> Self {
Self(reason)
}
#[inline]
pub fn reason(&self) -> ServiceDisconnectReason {
self.0
}
#[inline]
pub fn ack<F>(self) -> ControlAck<F> {
ControlAck {
result: ControlResult::Stop,
}
}
}
#[derive(Debug)]
pub struct KeepAlive {
enabled: bool,
}
impl KeepAlive {
pub(super) fn new(enabled: bool) -> Self {
Self { enabled }
}
#[inline]
pub fn is_enabled(&self) -> bool {
self.enabled
}
#[inline]
pub fn ack<F>(self) -> ControlAck<F> {
ControlAck {
result: ControlResult::Stop,
}
}
}
#[derive(Debug)]
pub struct Error<Err> {
err: Err,
pkt: Response,
}
impl<Err: ResponseError> Error<Err> {
fn new(err: Err) -> Self {
Self {
pkt: err.error_response(),
err,
}
}
#[inline]
pub fn get_ref(&self) -> &Err {
&self.err
}
#[inline]
pub fn get_mut(&mut self) -> &mut Err {
&mut self.err
}
#[inline]
pub fn ack<F>(self) -> ControlAck<F> {
let (res, body) = self.pkt.into_parts();
ControlAck {
result: ControlResult::Error(res, body.into()),
}
}
#[inline]
pub fn fail<E: ResponseError, F>(self, err: E) -> ControlAck<F> {
let res: Response = (&err).into();
let (res, body) = res.into_parts();
ControlAck {
result: ControlResult::Error(res, body.into()),
}
}
#[inline]
pub fn fail_with<F>(self, res: Response) -> ControlAck<F> {
let (res, body) = res.into_parts();
ControlAck {
result: ControlResult::Error(res, body.into()),
}
}
}
#[derive(Debug)]
pub struct ProtocolError(super::ProtocolError);
impl ProtocolError {
#[inline]
pub fn get_ref(&self) -> &super::ProtocolError {
&self.0
}
#[inline]
pub fn ack<F>(self) -> ControlAck<F> {
let (res, body) = self.0.error_response().into_parts();
ControlAck {
result: ControlResult::ProtocolError(res, body.into()),
}
}
#[inline]
pub fn fail<E: ResponseError, F>(self, err: E) -> ControlAck<F> {
let res: Response = (&err).into();
let (res, body) = res.into_parts();
ControlAck {
result: ControlResult::ProtocolError(res, body.into()),
}
}
#[inline]
pub fn fail_with<F>(self, res: Response) -> ControlAck<F> {
let (res, body) = res.into_parts();
ControlAck {
result: ControlResult::ProtocolError(res, body.into()),
}
}
}
#[derive(Debug)]
pub struct PeerGone(Option<io::Error>);
impl PeerGone {
#[inline]
pub fn get_ref(&self) -> Option<&io::Error> {
self.0.as_ref()
}
#[inline]
pub fn get_mut(&mut self) -> Option<&mut io::Error> {
self.0.as_mut()
}
#[inline]
pub fn take(&mut self) -> Option<io::Error> {
self.0.take()
}
#[inline]
pub fn ack<F>(self) -> ControlAck<F> {
ControlAck {
result: ControlResult::Stop,
}
}
}
#[derive(Debug)]
pub struct Expect(Request);
impl Expect {
#[inline]
pub fn get_ref(&self) -> &Request {
&self.0
}
#[inline]
pub fn get_mut(&mut self) -> &mut Request {
&mut self.0
}
#[inline]
pub fn ack<F>(self) -> ControlAck<F> {
ControlAck {
result: ControlResult::Continue(self.0),
}
}
#[inline]
pub fn fail<E: ResponseError, F>(self, err: E) -> ControlAck<F> {
let res: Response = (&err).into();
let (res, body) = res.into_parts();
ControlAck {
result: ControlResult::ExpectFailed(res, body.into()),
}
}
#[inline]
pub fn fail_with<F>(self, res: Response) -> ControlAck<F> {
let (res, body) = res.into_parts();
ControlAck {
result: ControlResult::ExpectFailed(res, body.into()),
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::http::HttpServiceConfig;
use crate::http::message::IoAccess;
use crate::service::cfg::SharedCfg;
use crate::testing::IoTest;
#[crate::rt_test]
async fn request_io_access_is_one_shot() {
let (_, server) = IoTest::create();
let cfg: SharedCfg = SharedCfg::new("TEST").add(HttpServiceConfig::new()).into();
let io = Rc::new(Io::new(server, cfg.clone()));
let access = RequestIoAccess {
ioref: io.get_ref(),
io,
codec: Codec::new(1, cfg.get()),
taken: Cell::new(false),
};
let ioref = access.get().unwrap();
let (io, _) = access.take().unwrap();
drop(io);
assert_eq!(ioref.tag(), "TEST");
assert!(access.get().is_none());
assert!(access.take().is_none());
}
type Ctl = Control<crate::io::Base, io::Error>;
fn io() -> Io {
let (_, server) = IoTest::create();
let cfg: SharedCfg = SharedCfg::new("TEST").add(HttpServiceConfig::new()).into();
Io::new(server, cfg)
}
#[test]
fn debug_fmt() {
let s = format!("{:?}", Ctl::request(Request::new()));
assert!(s.starts_with("Control::Request(NewRequest("), "{s}");
let s = format!("{:?}", Ctl::expect(Request::new()));
assert!(s.starts_with("Control::Expect(Expect("), "{s}");
let s = format!("{:?}", Ctl::keepalive(true));
assert!(s.contains("Control::Disconnect(KeepAlive("), "{s}");
let s = format!("{:?}", Ctl::err(io::Error::other("err")));
assert!(s.contains("Control::Disconnect(Error("), "{s}");
}
#[crate::rt_test]
async fn debug_fmt_io() {
let s = format!("{:?}", Ctl::connect(1, io()));
assert_eq!(s, "Control::Connect");
let cfg: SharedCfg = SharedCfg::new("TEST").add(HttpServiceConfig::new()).into();
let msg = Ctl::upgrade(Request::new(), Rc::new(io()), Codec::new(1, cfg.get()));
let s = format!("{msg:?}");
assert!(s.starts_with("Control::Upgrade(Upgrade"), "{s}");
}
#[crate::rt_test]
async fn connection() {
let Control::Connect(mut msg) = Ctl::connect(7, io()) else {
panic!()
};
assert_eq!(msg.id(), 7);
assert_eq!(msg.get_ref().tag(), "TEST");
assert_eq!(msg.get_mut().tag(), "TEST");
assert!(matches!(msg.ack().result, ControlResult::Connect(_)));
}
#[test]
fn new_request() {
let Control::Request(mut msg) = Ctl::request(Request::new()) else {
panic!()
};
msg.get_mut().head_mut().method = crate::http::Method::POST;
assert_eq!(msg.get_ref().method(), &crate::http::Method::POST);
assert!(matches!(
msg.ack::<crate::io::Base>().result,
ControlResult::Publish(_)
));
let Control::Request(msg) = Ctl::request(Request::new()) else {
panic!()
};
let ControlResult::Response(res, _) = msg
.fail::<_, crate::io::Base>(super::super::ProtocolError::SlowRequestTimeout)
.result
else {
panic!()
};
assert_eq!(res.status(), crate::http::StatusCode::REQUEST_TIMEOUT);
}
#[test]
fn expect() {
let Control::Expect(mut msg) = Ctl::expect(Request::new()) else {
panic!()
};
msg.get_mut().head_mut().method = crate::http::Method::PUT;
assert_eq!(msg.get_ref().method(), &crate::http::Method::PUT);
assert!(matches!(
msg.ack::<crate::io::Base>().result,
ControlResult::Continue(_)
));
}
#[test]
fn disconnect_reasons() {
let Control::Disconnect(Reason::Service(msg)) =
Ctl::svc_disconnect(ServiceDisconnectReason::PayloadDropped)
else {
panic!()
};
assert_eq!(msg.reason(), ServiceDisconnectReason::PayloadDropped);
assert!(matches!(
msg.ack::<crate::io::Base>().result,
ControlResult::Stop
));
for enabled in [true, false] {
let Control::Disconnect(Reason::KeepAlive(msg)) = Ctl::keepalive(enabled) else {
panic!()
};
assert_eq!(msg.is_enabled(), enabled);
assert!(matches!(
msg.ack::<crate::io::Base>().result,
ControlResult::Stop
));
}
let Control::Disconnect(Reason::Error(mut msg)) = Ctl::err(io::Error::other("err")) else {
panic!()
};
*msg.get_mut() = io::Error::other("changed");
assert_eq!(msg.get_ref().to_string(), "changed");
let ControlResult::Error(res, _) = msg.ack::<crate::io::Base>().result else {
panic!()
};
assert_eq!(res.status(), crate::http::StatusCode::INTERNAL_SERVER_ERROR);
let Control::Disconnect(Reason::ProtocolError(msg)) =
Ctl::proto_err(super::super::ProtocolError::SlowPayloadTimeout)
else {
panic!()
};
assert!(matches!(
msg.get_ref(),
super::super::ProtocolError::SlowPayloadTimeout
));
let Control::Disconnect(Reason::PeerGone(mut msg)) = Ctl::peer_gone(None) else {
panic!()
};
assert!(msg.get_ref().is_none());
assert!(msg.get_mut().is_none());
assert!(msg.take().is_none());
assert!(matches!(
msg.ack::<crate::io::Base>().result,
ControlResult::Stop
));
}
}