1use std::marker::PhantomData;
2
3use crate::http::error::HttpError;
4use crate::http::header::{HeaderMap, HeaderName, HeaderValue};
5use crate::http::{Response, ResponseBuilder, StatusCode};
6use crate::util::{Bytes, BytesMut, Either};
7
8use super::error::{InternalError, WebResponseError};
9use super::{HttpRequest, State};
10
11pub trait Responder<St: State = ()> {
53 async fn respond_to(self, st: &St, req: &HttpRequest) -> Response;
55
56 fn with_status(self, status: StatusCode) -> CustomResponder<Self, St>
70 where
71 Self: Sized,
72 {
73 CustomResponder::new(self).with_status(status)
74 }
75
76 fn with_header<K, V>(self, key: K, value: V) -> CustomResponder<Self, St>
98 where
99 Self: Sized,
100 HeaderName: TryFrom<K>,
101 HeaderValue: TryFrom<V>,
102 <HeaderName as TryFrom<K>>::Error: Into<HttpError>,
103 <HeaderValue as TryFrom<V>>::Error: Into<HttpError>,
104 {
105 CustomResponder::new(self).with_header(key, value)
106 }
107}
108
109impl<St: State> Responder<St> for Response {
110 #[inline]
111 async fn respond_to(self, _: &St, _: &HttpRequest) -> Response {
112 self
113 }
114}
115
116impl<St: State> Responder<St> for ResponseBuilder {
117 #[inline]
118 async fn respond_to(mut self, _: &St, _: &HttpRequest) -> Response {
119 self.build()
120 }
121}
122
123impl<T, St> Responder<St> for Option<T>
124where
125 T: Responder<St>,
126 St: State,
127{
128 async fn respond_to(self, st: &St, req: &HttpRequest) -> Response {
129 match self {
130 Some(t) => t.respond_to(st, req).await,
131 None => Response::builder(StatusCode::NOT_FOUND).build(),
132 }
133 }
134}
135
136impl<St, T, E> Responder<St> for Result<T, E>
137where
138 St: State,
139 T: Responder<St>,
140 E: WebResponseError<St, St::Error>,
141{
142 async fn respond_to(self, st: &St, req: &HttpRequest) -> Response {
143 match self {
144 Ok(val) => val.respond_to(st, req).await,
145 Err(e) => e.error_response(st),
146 }
147 }
148}
149
150impl<St, T> Responder<St> for (T, StatusCode)
151where
152 St: State,
153 T: Responder<St>,
154{
155 async fn respond_to(self, st: &St, req: &HttpRequest) -> Response {
156 let mut res = self.0.respond_to(st, req).await;
157 *res.status_mut() = self.1;
158 res
159 }
160}
161
162impl<St: State> Responder<St> for &'static str {
163 async fn respond_to(self, _: &St, _: &HttpRequest) -> Response {
164 Response::builder(StatusCode::OK)
165 .content_type("text/plain; charset=utf-8")
166 .body(self)
167 }
168}
169
170impl<St: State> Responder<St> for &'static [u8] {
171 async fn respond_to(self, _: &St, _: &HttpRequest) -> Response {
172 Response::builder(StatusCode::OK)
173 .content_type("application/octet-stream")
174 .body(self)
175 }
176}
177
178impl<St: State> Responder<St> for String {
179 async fn respond_to(self, _: &St, _: &HttpRequest) -> Response {
180 Response::builder(StatusCode::OK)
181 .content_type("text/plain; charset=utf-8")
182 .body(self)
183 }
184}
185
186impl<St: State> Responder<St> for &String {
187 async fn respond_to(self, _: &St, _: &HttpRequest) -> Response {
188 Response::builder(StatusCode::OK)
189 .content_type("text/plain; charset=utf-8")
190 .body(self)
191 }
192}
193
194impl<St: State> Responder<St> for Bytes {
195 async fn respond_to(self, _: &St, _: &HttpRequest) -> Response {
196 Response::builder(StatusCode::OK)
197 .content_type("application/octet-stream")
198 .body(self)
199 }
200}
201
202impl<St: State> Responder<St> for BytesMut {
203 async fn respond_to(self, _: &St, _: &HttpRequest) -> Response {
204 Response::builder(StatusCode::OK)
205 .content_type("application/octet-stream")
206 .body(self)
207 }
208}
209
210impl Responder<()> for () {
211 async fn respond_to(self, (): &(), _: &HttpRequest) -> Response {
212 Response::builder(StatusCode::OK).build()
213 }
214}
215
216#[derive(derive_more::Debug)]
218#[debug("CustomResponder")]
219pub struct CustomResponder<T: Responder<St>, St: State> {
220 responder: T,
221 status: Option<StatusCode>,
222 headers: Option<HeaderMap>,
223 error: Option<HttpError>,
224 _t: PhantomData<St>,
225}
226
227impl<T: Responder<St>, St: State> CustomResponder<T, St> {
228 fn new(responder: T) -> Self {
229 CustomResponder {
230 responder,
231 status: None,
232 headers: None,
233 error: None,
234 _t: PhantomData,
235 }
236 }
237
238 pub fn with_status(mut self, status: StatusCode) -> Self {
250 self.status = Some(status);
251 self
252 }
253
254 pub fn with_header<K, V>(mut self, key: K, value: V) -> Self
274 where
275 HeaderName: TryFrom<K>,
276 HeaderValue: TryFrom<V>,
277 <HeaderName as TryFrom<K>>::Error: Into<HttpError>,
278 <HeaderValue as TryFrom<V>>::Error: Into<HttpError>,
279 {
280 if self.headers.is_none() {
281 self.headers = Some(HeaderMap::new());
282 }
283
284 match HeaderName::try_from(key) {
285 Ok(key) => match HeaderValue::try_from(value) {
286 Ok(value) => {
287 self.headers.as_mut().unwrap().append(key, value);
288 }
289 Err(e) => self.error = Some(e.into()),
290 },
291 Err(e) => self.error = Some(e.into()),
292 }
293 self
294 }
295}
296
297impl<T: Responder<St>, St: State> Responder<St> for CustomResponder<T, St> {
298 async fn respond_to(self, st: &St, req: &HttpRequest) -> Response {
299 if let Some(err) = self.error {
300 return Response::from(err);
301 }
302 let mut res = self.responder.respond_to(st, req).await;
303
304 if let Some(status) = self.status {
305 *res.status_mut() = status;
306 }
307 if let Some(headers) = self.headers {
308 for key in headers.keys() {
309 res.headers_mut().remove(key);
310 }
311 for (k, v) in &headers {
312 res.headers_mut().append(k.clone(), v.clone());
313 }
314 }
315 res
316 }
317}
318
319impl<St, A, B> Responder<St> for Either<A, B>
337where
338 St: State,
339 A: Responder<St>,
340 B: Responder<St>,
341{
342 async fn respond_to(self, st: &St, req: &HttpRequest) -> Response {
343 match self {
344 Either::Left(a) => a.respond_to(st, req).await,
345 Either::Right(b) => b.respond_to(st, req).await,
346 }
347 }
348}
349
350impl<St, T> Responder<St> for InternalError<T>
351where
352 St: State,
353 T: std::fmt::Debug + std::fmt::Display + 'static,
354{
355 async fn respond_to(self, st: &St, _: &HttpRequest) -> Response {
356 WebResponseError::<St, St::Error>::error_response(&self, st)
357 }
358}
359
360#[cfg(test)]
361pub(crate) mod tests {
362 use super::*;
363 use crate::http::Response as HttpResponse;
364 use crate::http::body::{Body, ResponseBody};
365 use crate::http::header::CONTENT_TYPE;
366 use crate::web;
367 use crate::web::test::{TestRequest, init_service};
368
369 fn responder<T: Responder>(responder: T) -> impl Responder {
370 responder
371 }
372
373 #[crate::rt_test]
374 async fn test_either_responder() {
375 let srv = init_service(web::App::new().service(web::resource("/index.html").to(
376 async move |req: HttpRequest| {
377 if req.query_string().is_empty() {
378 Either::Left(HttpResponse::BadRequest())
379 } else {
380 Either::Right("hello")
381 }
382 },
383 )))
384 .await;
385
386 let req = TestRequest::with_uri("/index.html").to_request();
387 let resp = srv.call(req).await.unwrap();
388 assert_eq!(resp.status(), StatusCode::BAD_REQUEST);
389
390 let req = TestRequest::with_uri("/index.html?query=test").to_request();
391 let resp = srv.call(req).await.unwrap();
392 assert_eq!(resp.status(), StatusCode::OK);
393 }
394
395 #[crate::rt_test]
396 async fn test_option_responder() {
397 let srv = init_service(
398 web::App::new()
399 .service(web::resource("/none").to(async || Option::<&'static str>::None))
400 .service(web::resource("/some").to(async || Some("some"))),
401 )
402 .await;
403
404 let req = TestRequest::with_uri("/none").to_request();
405 let resp = srv.call(req).await.unwrap();
406 assert_eq!(resp.status(), StatusCode::NOT_FOUND);
407
408 let req = TestRequest::with_uri("/some").to_request();
409 let resp = srv.call(req).await.unwrap();
410 assert_eq!(resp.status(), StatusCode::OK);
411 if let ResponseBody::Body(Body::Bytes(b)) = resp.body() {
412 let bytes: Bytes = b.clone();
413 assert_eq!(bytes, Bytes::from_static(b"some"));
414 } else {
415 panic!()
416 }
417 }
418
419 #[crate::rt_test]
420 async fn test_responder() {
421 let req = TestRequest::default().to_http_request();
422
423 let resp: HttpResponse = responder("test").respond_to(&(), &req).await;
424 assert_eq!(resp.status(), StatusCode::OK);
425 assert_eq!(resp.get_body_ref(), b"test");
426 assert_eq!(
427 resp.headers().get(CONTENT_TYPE).unwrap(),
428 HeaderValue::from_static("text/plain; charset=utf-8")
429 );
430
431 let resp: HttpResponse = responder(&b"test"[..]).respond_to(&(), &req).await;
432 assert_eq!(resp.status(), StatusCode::OK);
433 assert_eq!(resp.get_body_ref(), b"test");
434 assert_eq!(
435 resp.headers().get(CONTENT_TYPE).unwrap(),
436 HeaderValue::from_static("application/octet-stream")
437 );
438
439 let resp: HttpResponse = responder("test".to_string()).respond_to(&(), &req).await;
440 assert_eq!(resp.status(), StatusCode::OK);
441 assert_eq!(resp.get_body_ref(), b"test");
442 assert_eq!(
443 resp.headers().get(CONTENT_TYPE).unwrap(),
444 HeaderValue::from_static("text/plain; charset=utf-8")
445 );
446
447 let resp: HttpResponse = responder(&"test".to_string()).respond_to(&(), &req).await;
448 assert_eq!(resp.status(), StatusCode::OK);
449 assert_eq!(resp.get_body_ref(), b"test");
450 assert_eq!(
451 resp.headers().get(CONTENT_TYPE).unwrap(),
452 HeaderValue::from_static("text/plain; charset=utf-8")
453 );
454
455 let resp: HttpResponse = responder(Bytes::from_static(b"test"))
456 .respond_to(&(), &req)
457 .await;
458 assert_eq!(resp.status(), StatusCode::OK);
459 assert_eq!(resp.get_body_ref(), b"test");
460 assert_eq!(
461 resp.headers().get(CONTENT_TYPE).unwrap(),
462 HeaderValue::from_static("application/octet-stream")
463 );
464
465 let resp: HttpResponse = responder(BytesMut::from(b"test".as_ref()))
466 .respond_to(&(), &req)
467 .await;
468 assert_eq!(resp.status(), StatusCode::OK);
469 assert_eq!(resp.get_body_ref(), b"test");
470 assert_eq!(
471 resp.headers().get(CONTENT_TYPE).unwrap(),
472 HeaderValue::from_static("application/octet-stream")
473 );
474
475 let resp: HttpResponse = responder(InternalError::new("err", StatusCode::BAD_REQUEST))
477 .respond_to(&(), &req)
478 .await;
479 assert_eq!(resp.status(), StatusCode::BAD_REQUEST);
480 }
481
482 #[crate::rt_test]
483 async fn test_result_responder() {
484 let req = TestRequest::default().to_http_request();
485
486 let resp: HttpResponse = Responder::<()>::respond_to(
488 Ok::<String, std::convert::Infallible>("test".to_string()),
489 &(),
490 &req,
491 )
492 .await;
493 assert_eq!(resp.status(), StatusCode::OK);
494 assert_eq!(resp.get_body_ref(), b"test");
495 assert_eq!(
496 resp.headers().get(CONTENT_TYPE).unwrap(),
497 HeaderValue::from_static("text/plain; charset=utf-8")
498 );
499
500 let res = responder(Err::<String, _>(InternalError::new(
501 "err",
502 StatusCode::BAD_REQUEST,
503 )))
504 .respond_to(&(), &req)
505 .await;
506 assert_eq!(res.status(), StatusCode::BAD_REQUEST);
507 }
508
509 #[crate::rt_test]
510 async fn test_custom_responder() {
511 let req = TestRequest::default().to_http_request();
512 let res = responder("test".to_string())
513 .with_status(StatusCode::BAD_REQUEST)
514 .respond_to(&(), &req)
515 .await;
516 assert_eq!(res.status(), StatusCode::BAD_REQUEST);
517 assert_eq!(res.get_body_ref(), b"test");
518
519 let res = responder("test".to_string())
520 .with_header("content-type", "json")
521 .respond_to(&(), &req)
522 .await;
523
524 assert_eq!(res.status(), StatusCode::OK);
525 assert_eq!(res.get_body_ref(), b"test");
526 assert_eq!(
527 res.headers().get(CONTENT_TYPE).unwrap(),
528 HeaderValue::from_static("json")
529 );
530 }
531
532 #[crate::rt_test]
533 async fn test_custom_responder_headers() {
534 let req = TestRequest::default().to_http_request();
535 let res = responder("test".to_string())
536 .with_header("x-test", "1")
537 .with_header("x-test", "2")
538 .respond_to(&(), &req)
539 .await;
540 assert_eq!(res.status(), StatusCode::OK);
541 let values: Vec<_> = res.headers().get_all("x-test").collect();
542 assert_eq!(
543 values,
544 [HeaderValue::from_static("1"), HeaderValue::from_static("2")]
545 );
546
547 let res = responder(
549 HttpResponse::Ok()
550 .header(CONTENT_TYPE, "text/plain")
551 .header(CONTENT_TYPE, "text/html")
552 .build(),
553 )
554 .with_header(CONTENT_TYPE, "json")
555 .respond_to(&(), &req)
556 .await;
557 let values: Vec<_> = res.headers().get_all(CONTENT_TYPE).collect();
558 assert_eq!(values, [HeaderValue::from_static("json")]);
559
560 let res = responder("test".to_string())
561 .with_header("bad header", "1")
562 .respond_to(&(), &req)
563 .await;
564 assert_eq!(res.status(), StatusCode::INTERNAL_SERVER_ERROR);
565
566 let res = responder("test".to_string())
567 .with_header("x-test", "bad\nvalue")
568 .with_status(StatusCode::CREATED)
569 .respond_to(&(), &req)
570 .await;
571 assert_eq!(res.status(), StatusCode::INTERNAL_SERVER_ERROR);
572 assert!(res.headers().get("x-test").is_none());
573 }
574
575 #[crate::rt_test]
576 async fn test_tuple_responder_with_status_code() {
577 let req = TestRequest::default().to_http_request();
578 let res =
579 Responder::<()>::respond_to(("test".to_string(), StatusCode::BAD_REQUEST), &(), &req)
580 .await;
581 assert_eq!(res.status(), StatusCode::BAD_REQUEST);
582 assert_eq!(res.get_body_ref(), b"test");
583
584 let req = TestRequest::default().to_http_request();
585 let res = CustomResponder::<_, ()>::new(("test".to_string(), StatusCode::OK))
586 .with_header("content-type", "json")
587 .respond_to(&(), &req)
588 .await;
589 assert_eq!(res.status(), StatusCode::OK);
590 assert_eq!(res.get_body_ref(), b"test");
591 assert_eq!(
592 res.headers().get(CONTENT_TYPE).unwrap(),
593 HeaderValue::from_static("json")
594 );
595 }
596}