1use std::{cell::Cell, fmt, io, rc::Rc};
2
3use crate::http::message::CurrentIo;
4use crate::http::{Request, Response, ResponseError, body::Body, h1::Codec};
5use crate::io::{Filter, Io, IoBoxed, IoRef};
6
7pub enum Control<F, Err> {
12 Connect(Connection<F>),
14 Request(NewRequest),
16 Upgrade(Upgrade<F>),
18 Expect(Expect),
20 Disconnect(Reason<Err>),
22}
23
24#[derive(Debug)]
25pub enum Reason<Err> {
27 Service(ServiceDisconnect),
29 Error(Error<Err>),
31 ProtocolError(ProtocolError),
33 PeerGone(PeerGone),
35 KeepAlive(KeepAlive),
37}
38
39#[derive(Copy, Clone, Debug, PartialEq, Eq)]
41pub enum ServiceDisconnectReason {
42 Shutdown,
44 UpgradeHandled,
47 UpgradeFailed,
49 ExpectFailed,
51 PayloadDropped,
54}
55
56#[derive(Debug)]
58pub struct ControlAck<F> {
59 pub(super) result: ControlResult<F>,
60}
61
62#[derive(Debug)]
63pub(super) enum ControlResult<F> {
64 Connect(Io<F>),
66 Continue(Request),
68 Expect(Request),
70 Upgrade(Request),
72 UpgradeAck(Request),
74 UpgradeHandled,
76 Publish(Request),
78 Response(Response<()>, Body),
80 Error(Response<()>, Body),
82 ProtocolError(Response<()>, Body),
84 UpgradeFailed(Response<()>, Body),
86 ExpectFailed(Response<()>, Body),
88 Stop,
90}
91
92impl<F, Err> Control<F, Err> {
93 pub(super) fn connect(id: usize, io: Io<F>) -> Self {
94 Control::Connect(Connection { id, io })
95 }
96
97 pub(super) fn request(req: Request) -> Self {
98 Control::Request(NewRequest(req))
99 }
100
101 pub(super) fn upgrade(req: Request, io: Rc<Io<F>>, codec: Codec) -> Self {
102 Control::Upgrade(Upgrade { req, io, codec })
103 }
104
105 pub(super) fn expect(req: Request) -> Self {
106 Control::Expect(Expect(req))
107 }
108
109 pub(super) fn err(err: Err) -> Self
110 where
111 Err: ResponseError,
112 {
113 Control::Disconnect(Reason::Error(Error::new(err)))
114 }
115
116 pub(super) fn peer_gone(err: Option<io::Error>) -> Self {
117 Control::Disconnect(Reason::PeerGone(PeerGone(err)))
118 }
119
120 pub(super) fn proto_err(err: super::ProtocolError) -> Self {
121 Control::Disconnect(Reason::ProtocolError(ProtocolError(err)))
122 }
123
124 pub(super) fn keepalive(enabled: bool) -> Self {
125 Control::Disconnect(Reason::KeepAlive(KeepAlive::new(enabled)))
126 }
127
128 pub(super) fn svc_disconnect(reason: ServiceDisconnectReason) -> Self {
129 Control::Disconnect(Reason::Service(ServiceDisconnect::new(reason)))
130 }
131
132 #[inline]
133 pub fn ack(self) -> ControlAck<F>
135 where
136 F: Filter,
137 Err: ResponseError,
138 {
139 match self {
140 Control::Connect(msg) => msg.ack(),
141 Control::Request(msg) => msg.ack(),
142 Control::Upgrade(msg) => msg.ack(),
143 Control::Expect(msg) => msg.ack(),
144 Control::Disconnect(msg) => msg.ack(),
145 }
146 }
147}
148
149impl<F, Err> fmt::Debug for Control<F, Err>
150where
151 Err: fmt::Debug,
152{
153 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
154 match self {
155 Control::Connect(_) => f.debug_tuple("Control::Connect").finish(),
156 Control::Request(msg) => f.debug_tuple("Control::Request").field(msg).finish(),
157 Control::Upgrade(msg) => f.debug_tuple("Control::Upgrade").field(msg).finish(),
158 Control::Expect(msg) => f.debug_tuple("Control::Expect").field(msg).finish(),
159 Control::Disconnect(msg) => f.debug_tuple("Control::Disconnect").field(msg).finish(),
160 }
161 }
162}
163
164impl<Err: ResponseError> Reason<Err> {
165 pub fn ack<F>(self) -> ControlAck<F> {
167 match self {
168 Reason::Error(msg) => msg.ack(),
169 Reason::ProtocolError(msg) => msg.ack(),
170 Reason::PeerGone(msg) => msg.ack(),
171 Reason::KeepAlive(msg) => msg.ack(),
172 Reason::Service(msg) => msg.ack(),
173 }
174 }
175}
176
177#[derive(Debug)]
179pub struct Connection<F> {
180 id: usize,
181 io: Io<F>,
182}
183
184impl<F> Connection<F> {
185 #[inline]
186 pub fn id(&self) -> usize {
188 self.id
189 }
190
191 #[inline]
192 pub fn get_ref(&self) -> &Io<F> {
194 &self.io
195 }
196
197 #[inline]
198 pub fn get_mut(&mut self) -> &mut Io<F> {
200 &mut self.io
201 }
202
203 #[inline]
204 pub fn ack(self) -> ControlAck<F> {
206 ControlAck {
207 result: ControlResult::Connect(self.io),
208 }
209 }
210}
211
212#[derive(Debug)]
214pub struct NewRequest(Request);
215
216impl NewRequest {
217 #[inline]
218 pub fn get_ref(&self) -> &Request {
220 &self.0
221 }
222
223 #[inline]
224 pub fn get_mut(&mut self) -> &mut Request {
226 &mut self.0
227 }
228
229 #[inline]
230 pub fn ack<F>(self) -> ControlAck<F> {
233 let result = if self.0.head().expect() {
234 ControlResult::Expect(self.0)
235 } else if self.0.upgrade() {
236 ControlResult::Upgrade(self.0)
237 } else {
238 ControlResult::Publish(self.0)
239 };
240 ControlAck { result }
241 }
242
243 #[inline]
244 pub fn fail<E: ResponseError, F>(self, err: E) -> ControlAck<F> {
246 let res: Response = (&err).into();
247 let (res, body) = res.into_parts();
248
249 ControlAck {
250 result: ControlResult::Response(res, body.into()),
251 }
252 }
253
254 #[inline]
255 pub fn fail_with<F>(self, res: Response) -> ControlAck<F> {
257 let (res, body) = res.into_parts();
258
259 ControlAck {
260 result: ControlResult::Response(res, body.into()),
261 }
262 }
263}
264
265pub struct Upgrade<F> {
267 req: Request,
268 io: Rc<Io<F>>,
269 codec: Codec,
270}
271
272struct RequestIoAccess<F> {
273 io: Rc<Io<F>>,
274 ioref: IoRef,
276 codec: Codec,
277 taken: Cell<bool>,
278}
279
280impl<F> fmt::Debug for RequestIoAccess<F> {
281 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
282 f.debug_struct("RequestIoAccess")
283 .field("io", &self.ioref)
284 .field("codec", &self.codec)
285 .finish()
286 }
287}
288
289impl<F: Filter> crate::http::message::IoAccess for RequestIoAccess<F> {
290 fn get(&self) -> Option<&IoRef> {
291 if self.taken.get() { None } else { Some(&self.ioref) }
292 }
293
294 fn take(&self) -> Option<(IoBoxed, Codec)> {
295 if self.taken.replace(true) {
296 None
297 } else {
298 let io = unsafe { self.io.take() };
301 Some((io.into(), self.codec.clone()))
302 }
303 }
304}
305
306impl<F: Filter> Upgrade<F> {
307 #[inline]
308 pub fn io(&self) -> &Io<F> {
310 &self.io
311 }
312
313 #[inline]
314 pub fn get_ref(&self) -> &Request {
316 &self.req
317 }
318
319 #[inline]
320 pub fn get_mut(&mut self) -> &mut Request {
322 &mut self.req
323 }
324
325 #[inline]
326 pub fn ack(mut self) -> ControlAck<F> {
331 let io = Rc::new(RequestIoAccess {
333 ioref: self.io.get_ref(),
334 io: self.io,
335 codec: self.codec,
336 taken: Cell::new(false),
337 });
338 self.req.head_mut().io = CurrentIo::new(io);
339
340 ControlAck {
341 result: ControlResult::UpgradeAck(self.req),
342 }
343 }
344
345 #[inline]
346 pub fn handle(self) -> (ControlAck<F>, Io<F>, Request, Codec) {
351 let io = unsafe { self.io.take() };
354 (
355 ControlAck {
356 result: ControlResult::UpgradeHandled,
357 },
358 io,
359 self.req,
360 self.codec,
361 )
362 }
363
364 #[inline]
365 pub fn fail<E: ResponseError>(self, err: E) -> ControlAck<F> {
367 let res: Response = (&err).into();
368 let (res, body) = res.into_parts();
369
370 ControlAck {
371 result: ControlResult::UpgradeFailed(res, body.into()),
372 }
373 }
374
375 #[inline]
376 pub fn fail_with(self, res: Response) -> ControlAck<F> {
378 let (res, body) = res.into_parts();
379
380 ControlAck {
381 result: ControlResult::UpgradeFailed(res, body.into()),
382 }
383 }
384}
385
386impl<F> fmt::Debug for Upgrade<F> {
387 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
388 f.debug_struct("Upgrade")
389 .field("req", &self.req)
390 .field("io", &self.io)
391 .field("codec", &self.codec)
392 .finish()
393 }
394}
395
396#[derive(Debug)]
398pub struct ServiceDisconnect(ServiceDisconnectReason);
399
400impl ServiceDisconnect {
401 fn new(reason: ServiceDisconnectReason) -> Self {
402 Self(reason)
403 }
404
405 #[inline]
406 pub fn reason(&self) -> ServiceDisconnectReason {
408 self.0
409 }
410
411 #[inline]
412 pub fn ack<F>(self) -> ControlAck<F> {
414 ControlAck {
415 result: ControlResult::Stop,
416 }
417 }
418}
419
420#[derive(Debug)]
422pub struct KeepAlive {
423 enabled: bool,
424}
425
426impl KeepAlive {
427 pub(super) fn new(enabled: bool) -> Self {
428 Self { enabled }
429 }
430
431 #[inline]
432 pub fn is_enabled(&self) -> bool {
434 self.enabled
435 }
436
437 #[inline]
438 pub fn ack<F>(self) -> ControlAck<F> {
440 ControlAck {
441 result: ControlResult::Stop,
442 }
443 }
444}
445
446#[derive(Debug)]
448pub struct Error<Err> {
449 err: Err,
450 pkt: Response,
451}
452
453impl<Err: ResponseError> Error<Err> {
454 fn new(err: Err) -> Self {
455 Self {
456 pkt: err.error_response(),
457 err,
458 }
459 }
460
461 #[inline]
462 pub fn get_ref(&self) -> &Err {
464 &self.err
465 }
466
467 #[inline]
468 pub fn get_mut(&mut self) -> &mut Err {
470 &mut self.err
471 }
472
473 #[inline]
474 pub fn ack<F>(self) -> ControlAck<F> {
476 let (res, body) = self.pkt.into_parts();
477 ControlAck {
478 result: ControlResult::Error(res, body.into()),
479 }
480 }
481
482 #[inline]
483 pub fn fail<E: ResponseError, F>(self, err: E) -> ControlAck<F> {
485 let res: Response = (&err).into();
486 let (res, body) = res.into_parts();
487
488 ControlAck {
489 result: ControlResult::Error(res, body.into()),
490 }
491 }
492
493 #[inline]
494 pub fn fail_with<F>(self, res: Response) -> ControlAck<F> {
496 let (res, body) = res.into_parts();
497
498 ControlAck {
499 result: ControlResult::Error(res, body.into()),
500 }
501 }
502}
503
504#[derive(Debug)]
506pub struct ProtocolError(super::ProtocolError);
507
508impl ProtocolError {
509 #[inline]
510 pub fn get_ref(&self) -> &super::ProtocolError {
512 &self.0
513 }
514
515 #[inline]
516 pub fn ack<F>(self) -> ControlAck<F> {
518 let (res, body) = self.0.error_response().into_parts();
519
520 ControlAck {
521 result: ControlResult::ProtocolError(res, body.into()),
522 }
523 }
524
525 #[inline]
526 pub fn fail<E: ResponseError, F>(self, err: E) -> ControlAck<F> {
528 let res: Response = (&err).into();
529 let (res, body) = res.into_parts();
530
531 ControlAck {
532 result: ControlResult::ProtocolError(res, body.into()),
533 }
534 }
535
536 #[inline]
537 pub fn fail_with<F>(self, res: Response) -> ControlAck<F> {
539 let (res, body) = res.into_parts();
540
541 ControlAck {
542 result: ControlResult::ProtocolError(res, body.into()),
543 }
544 }
545}
546
547#[derive(Debug)]
549pub struct PeerGone(Option<io::Error>);
550
551impl PeerGone {
552 #[inline]
553 pub fn get_ref(&self) -> Option<&io::Error> {
555 self.0.as_ref()
556 }
557
558 #[inline]
559 pub fn get_mut(&mut self) -> Option<&mut io::Error> {
561 self.0.as_mut()
562 }
563
564 #[inline]
565 pub fn take(&mut self) -> Option<io::Error> {
567 self.0.take()
568 }
569
570 #[inline]
571 pub fn ack<F>(self) -> ControlAck<F> {
573 ControlAck {
574 result: ControlResult::Stop,
575 }
576 }
577}
578
579#[derive(Debug)]
581pub struct Expect(Request);
582
583impl Expect {
584 #[inline]
585 pub fn get_ref(&self) -> &Request {
587 &self.0
588 }
589
590 #[inline]
591 pub fn get_mut(&mut self) -> &mut Request {
593 &mut self.0
594 }
595
596 #[inline]
597 pub fn ack<F>(self) -> ControlAck<F> {
599 ControlAck {
600 result: ControlResult::Continue(self.0),
601 }
602 }
603
604 #[inline]
605 pub fn fail<E: ResponseError, F>(self, err: E) -> ControlAck<F> {
607 let res: Response = (&err).into();
608 let (res, body) = res.into_parts();
609
610 ControlAck {
611 result: ControlResult::ExpectFailed(res, body.into()),
612 }
613 }
614
615 #[inline]
616 pub fn fail_with<F>(self, res: Response) -> ControlAck<F> {
618 let (res, body) = res.into_parts();
619
620 ControlAck {
621 result: ControlResult::ExpectFailed(res, body.into()),
622 }
623 }
624}
625
626#[cfg(test)]
627mod tests {
628 use super::*;
629 use crate::http::HttpServiceConfig;
630 use crate::http::message::IoAccess;
631 use crate::service::cfg::SharedCfg;
632 use crate::testing::IoTest;
633
634 #[crate::rt_test]
635 async fn request_io_access_is_one_shot() {
636 let (_, server) = IoTest::create();
637 let cfg: SharedCfg = SharedCfg::new("TEST").add(HttpServiceConfig::new()).into();
638 let io = Rc::new(Io::new(server, cfg.clone()));
639 let access = RequestIoAccess {
640 ioref: io.get_ref(),
641 io,
642 codec: Codec::new(1, cfg.get()),
643 taken: Cell::new(false),
644 };
645
646 let ioref = access.get().unwrap();
647 let (io, _) = access.take().unwrap();
648 drop(io);
649 assert_eq!(ioref.tag(), "TEST");
651 assert!(access.get().is_none());
652 assert!(access.take().is_none());
653 }
654
655 type Ctl = Control<crate::io::Base, io::Error>;
656
657 fn io() -> Io {
658 let (_, server) = IoTest::create();
659 let cfg: SharedCfg = SharedCfg::new("TEST").add(HttpServiceConfig::new()).into();
660 Io::new(server, cfg)
661 }
662
663 #[test]
664 fn debug_fmt() {
665 let s = format!("{:?}", Ctl::request(Request::new()));
666 assert!(s.starts_with("Control::Request(NewRequest("), "{s}");
667 let s = format!("{:?}", Ctl::expect(Request::new()));
668 assert!(s.starts_with("Control::Expect(Expect("), "{s}");
669 let s = format!("{:?}", Ctl::keepalive(true));
670 assert!(s.contains("Control::Disconnect(KeepAlive("), "{s}");
671 let s = format!("{:?}", Ctl::err(io::Error::other("err")));
672 assert!(s.contains("Control::Disconnect(Error("), "{s}");
673 }
674
675 #[crate::rt_test]
676 async fn debug_fmt_io() {
677 let s = format!("{:?}", Ctl::connect(1, io()));
678 assert_eq!(s, "Control::Connect");
679
680 let cfg: SharedCfg = SharedCfg::new("TEST").add(HttpServiceConfig::new()).into();
681 let msg = Ctl::upgrade(Request::new(), Rc::new(io()), Codec::new(1, cfg.get()));
682 let s = format!("{msg:?}");
683 assert!(s.starts_with("Control::Upgrade(Upgrade"), "{s}");
684 }
685
686 #[crate::rt_test]
687 async fn connection() {
688 let Control::Connect(mut msg) = Ctl::connect(7, io()) else {
689 panic!()
690 };
691 assert_eq!(msg.id(), 7);
692 assert_eq!(msg.get_ref().tag(), "TEST");
693 assert_eq!(msg.get_mut().tag(), "TEST");
694 assert!(matches!(msg.ack().result, ControlResult::Connect(_)));
695 }
696
697 #[test]
698 fn new_request() {
699 let Control::Request(mut msg) = Ctl::request(Request::new()) else {
700 panic!()
701 };
702 msg.get_mut().head_mut().method = crate::http::Method::POST;
703 assert_eq!(msg.get_ref().method(), &crate::http::Method::POST);
704 assert!(matches!(
705 msg.ack::<crate::io::Base>().result,
706 ControlResult::Publish(_)
707 ));
708
709 let Control::Request(msg) = Ctl::request(Request::new()) else {
710 panic!()
711 };
712 let ControlResult::Response(res, _) = msg
713 .fail::<_, crate::io::Base>(super::super::ProtocolError::SlowRequestTimeout)
714 .result
715 else {
716 panic!()
717 };
718 assert_eq!(res.status(), crate::http::StatusCode::REQUEST_TIMEOUT);
719 }
720
721 #[test]
722 fn expect() {
723 let Control::Expect(mut msg) = Ctl::expect(Request::new()) else {
724 panic!()
725 };
726 msg.get_mut().head_mut().method = crate::http::Method::PUT;
727 assert_eq!(msg.get_ref().method(), &crate::http::Method::PUT);
728 assert!(matches!(
729 msg.ack::<crate::io::Base>().result,
730 ControlResult::Continue(_)
731 ));
732 }
733
734 #[test]
735 fn disconnect_reasons() {
736 let Control::Disconnect(Reason::Service(msg)) =
737 Ctl::svc_disconnect(ServiceDisconnectReason::PayloadDropped)
738 else {
739 panic!()
740 };
741 assert_eq!(msg.reason(), ServiceDisconnectReason::PayloadDropped);
742 assert!(matches!(
743 msg.ack::<crate::io::Base>().result,
744 ControlResult::Stop
745 ));
746
747 for enabled in [true, false] {
748 let Control::Disconnect(Reason::KeepAlive(msg)) = Ctl::keepalive(enabled) else {
749 panic!()
750 };
751 assert_eq!(msg.is_enabled(), enabled);
752 assert!(matches!(
753 msg.ack::<crate::io::Base>().result,
754 ControlResult::Stop
755 ));
756 }
757
758 let Control::Disconnect(Reason::Error(mut msg)) = Ctl::err(io::Error::other("err")) else {
759 panic!()
760 };
761 *msg.get_mut() = io::Error::other("changed");
762 assert_eq!(msg.get_ref().to_string(), "changed");
763 let ControlResult::Error(res, _) = msg.ack::<crate::io::Base>().result else {
764 panic!()
765 };
766 assert_eq!(res.status(), crate::http::StatusCode::INTERNAL_SERVER_ERROR);
767
768 let Control::Disconnect(Reason::ProtocolError(msg)) =
769 Ctl::proto_err(super::super::ProtocolError::SlowPayloadTimeout)
770 else {
771 panic!()
772 };
773 assert!(matches!(
774 msg.get_ref(),
775 super::super::ProtocolError::SlowPayloadTimeout
776 ));
777
778 let Control::Disconnect(Reason::PeerGone(mut msg)) = Ctl::peer_gone(None) else {
779 panic!()
780 };
781 assert!(msg.get_ref().is_none());
782 assert!(msg.get_mut().is_none());
783 assert!(msg.take().is_none());
784 assert!(matches!(
785 msg.ack::<crate::io::Base>().result,
786 ControlResult::Stop
787 ));
788 }
789}