use super::{Res, TryRes};
use crate::error::{
FromServerFnError, IntoAppError, SERVER_FN_ERROR_HEADER, ServerFnErrorErr,
ServerFnErrorResponseParts, ServerFnErrorWrapper,
};
use axum::body::Body;
use bytes::Bytes;
use futures::{Stream, TryStreamExt};
use http::{HeaderValue, Response, StatusCode, header};
impl<E> TryRes<E> for Response<Body>
where
E: Send + Sync + FromServerFnError,
{
fn try_from_string(content_type: &str, data: String) -> Result<Self, E> {
let builder = http::Response::builder();
builder
.status(200)
.header(header::CONTENT_TYPE, content_type)
.body(Body::from(data))
.map_err(|e| {
ServerFnErrorErr::Response(e.to_string()).into_app_error()
})
}
fn try_from_bytes(content_type: &str, data: Bytes) -> Result<Self, E> {
let builder = http::Response::builder();
builder
.status(200)
.header(header::CONTENT_TYPE, content_type)
.body(Body::from(data))
.map_err(|e| {
ServerFnErrorErr::Response(e.to_string()).into_app_error()
})
}
fn try_from_stream(
content_type: &str,
data: impl Stream<Item = Result<Bytes, Bytes>> + Send + 'static,
) -> Result<Self, E> {
let body =
Body::from_stream(data.map_err(|e| ServerFnErrorWrapper(E::de(e))));
let builder = http::Response::builder();
builder
.status(200)
.header(header::CONTENT_TYPE, content_type)
.body(body)
.map_err(|e| {
ServerFnErrorErr::Response(e.to_string()).into_app_error()
})
}
}
impl Res for Response<Body> {
fn error_response(path: &str, err: ServerFnErrorResponseParts) -> Self {
let status_code = err.status_code;
let content_type = err.content_type;
let body = err.body;
let mut builder = Response::builder()
.status(status_code)
.header(header::CONTENT_TYPE, content_type);
if let Ok(path_value) = HeaderValue::from_str(path) {
builder = builder.header(SERVER_FN_ERROR_HEADER, path_value);
}
builder.body(Body::from(body.clone())).unwrap_or_else(|_| {
let mut fallback = Response::new(Body::from(body));
*fallback.status_mut() = status_code;
fallback
})
}
fn redirect(&mut self, path: &str) {
if let Ok(path) = HeaderValue::from_str(path) {
self.headers_mut().insert(header::LOCATION, path);
*self.status_mut() = StatusCode::FOUND;
}
}
}
#[cfg(test)]
mod tests {
use super::*;
fn build_parts() -> ServerFnErrorResponseParts {
ServerFnErrorResponseParts::builder()
.body(Bytes::from_static(b"oops"))
.content_type("text/plain")
.status_code(StatusCode::INTERNAL_SERVER_ERROR)
.build()
}
#[test]
fn error_response_does_not_panic_on_invalid_path_header_bytes() {
let resp = <Response<Body> as Res>::error_response(
"/api/foo\r\nX-Evil: y",
build_parts(),
);
assert_eq!(resp.status(), StatusCode::INTERNAL_SERVER_ERROR);
assert!(resp.headers().get(SERVER_FN_ERROR_HEADER).is_none());
}
#[test]
fn error_response_includes_path_header_when_valid() {
let resp =
<Response<Body> as Res>::error_response("/api/foo", build_parts());
assert_eq!(resp.status(), StatusCode::INTERNAL_SERVER_ERROR);
assert_eq!(
resp.headers()
.get(SERVER_FN_ERROR_HEADER)
.map(|v| v.as_bytes()),
Some(&b"/api/foo"[..]),
);
assert_eq!(
resp.headers()
.get(header::CONTENT_TYPE)
.map(|v| v.as_bytes()),
Some(&b"text/plain"[..]),
);
}
}