use crate::error::{Error, Result};
use crate::protocol::browser_context::{BrowserContext, Cookie};
use crate::protocol::route::{FulfillOptions, Route};
use crate::protocol::route_params::merge_headers;
use bytes::Bytes;
use http::{HeaderMap, HeaderName, HeaderValue, Method, Request, Response, StatusCode, Uri};
use http_body_util::{BodyExt, Full};
use tower::{Service, ServiceExt};
pub use tower::BoxError;
pub type ServiceRequest = Request<Full<Bytes>>;
pub trait RouteService:
Service<ServiceRequest, Response = Response<Self::Body>, Future: Send, Error: Into<BoxError>>
+ Clone
+ Send
+ Sync
+ 'static
{
type Body: http_body::Body<Data: Send, Error: Into<BoxError>> + Send + 'static;
}
impl<S, B> RouteService for S
where
S: Service<ServiceRequest, Response = Response<B>> + Clone + Send + Sync + 'static,
S::Future: Send,
S::Error: Into<BoxError>,
B: http_body::Body + Send + 'static,
B::Data: Send,
B::Error: Into<BoxError>,
{
type Body = B;
}
pub(crate) async fn fulfill_from_service<S: RouteService>(
route: Route,
service: S,
context: Option<BrowserContext>,
) -> Result<()> {
let options = service_response(&route, service, context).await?;
route.fulfill(Some(options)).await
}
async fn service_response<S: RouteService>(
route: &Route,
service: S,
context: Option<BrowserContext>,
) -> Result<FulfillOptions> {
let request = route.request();
let context = context.or_else(|| request_context(route));
let is_webkit = context
.as_ref()
.and_then(BrowserContext::browser)
.is_some_and(|browser| browser.name() == "webkit");
let mut headers = request.header_pairs();
if is_webkit
&& !headers.iter().any(|(name, _)| name == "cookie")
&& let Some(context) = &context
{
let cookies = context.cookies(Some(&[request.url()])).await?;
if let Some(cookie) = cookie_header(&cookies) {
headers.push(("cookie".to_string(), cookie));
}
}
let http_request = request_from_parts(
request.method(),
request.url(),
headers
.iter()
.map(|(name, value)| (name.as_str(), value.as_str())),
request.post_data_buffer(),
)?;
let response = service.oneshot(http_request).await.map_err(|error| {
Error::ServerError(format!(
"route_service: the service failed for {}: {}",
request.url(),
error.into()
))
})?;
let (parts, body) = response.into_parts();
if is_webkit && parts.status.is_redirection() {
return Err(Error::ServerError(format!(
"route_service: WebKit cannot fulfill the {} redirect the service returned for {}",
parts.status,
request.url()
)));
}
let body = body.collect().await.map_err(|error| {
Error::ServerError(format!(
"route_service: the response body for {} failed: {}",
request.url(),
error.into()
))
})?;
Ok(fulfill_options(
parts.status,
&parts.headers,
body.to_bytes(),
))
}
fn request_context(route: &Route) -> Option<BrowserContext> {
route.request().frame()?.page()?.context().ok()
}
pub(crate) fn cookie_header(cookies: &[Cookie]) -> Option<String> {
if cookies.is_empty() {
return None;
}
Some(
cookies
.iter()
.map(|cookie| format!("{}={}", cookie.name, cookie.value))
.collect::<Vec<_>>()
.join("; "),
)
}
pub(crate) fn request_from_parts<'a>(
method: &str,
url: &str,
headers: impl IntoIterator<Item = (&'a str, &'a str)>,
body: Option<Vec<u8>>,
) -> Result<ServiceRequest> {
let method = Method::from_bytes(method.as_bytes()).map_err(|error| {
Error::ProtocolError(format!("route_service: invalid method {method:?}: {error}"))
})?;
let uri: Uri = url.parse().map_err(|error| {
Error::ProtocolError(format!("route_service: invalid URL {url:?}: {error}"))
})?;
let mut declared_length: Option<u64> = None;
let mut builder = Request::builder().method(method).uri(uri);
for (name, value) in headers {
if name == "content-length" {
declared_length = value.trim().parse().ok();
}
let (Ok(name), Ok(value)) = (
HeaderName::from_bytes(name.as_bytes()),
HeaderValue::from_bytes(value.as_bytes()),
) else {
continue;
};
builder = builder.header(name, value);
}
if body.is_none()
&& let Some(length) = declared_length
&& length > 0
{
return Err(Error::ProtocolError(format!(
"route_service: the driver did not expose the {length}-byte request body for {url} \
(multipart file uploads and very large bodies are withheld from route handlers); \
serve this request from a real listener"
)));
}
builder
.body(Full::new(Bytes::from(body.unwrap_or_default())))
.map_err(|error| Error::ProtocolError(format!("route_service: {error}")))
}
pub(crate) fn fulfill_options(
status: StatusCode,
headers: &HeaderMap,
body: Bytes,
) -> FulfillOptions {
let pairs = headers.iter().filter_map(|(name, value)| {
let name = name.as_str();
if name == "content-length" || name == "transfer-encoding" {
return None;
}
Some((name, String::from_utf8_lossy(value.as_bytes())))
});
FulfillOptions::builder()
.status(status.as_u16())
.headers(merge_headers(pairs, Some(", ")))
.body(Vec::from(body))
.build()
}
#[cfg(test)]
mod tests {
use super::*;
use crate::protocol::route_params::headers_of;
#[tokio::test]
async fn request_carries_method_absolute_uri_headers_and_body() {
let request = request_from_parts(
"POST",
"https://app.example/api/items?x=1",
[("content-type", "text/plain"), ("host", "app.example")],
Some(b"payload".to_vec()),
)
.unwrap();
assert_eq!(request.method(), Method::POST);
assert_eq!(request.uri().path(), "/api/items");
assert_eq!(request.uri().query(), Some("x=1"));
assert_eq!(request.uri().host(), Some("app.example"));
assert_eq!(request.headers()["content-type"], "text/plain");
assert_eq!(request.headers()["host"], "app.example");
let body = request.into_body().collect().await.unwrap().to_bytes();
assert_eq!(&body[..], b"payload");
}
#[tokio::test]
async fn request_without_a_body_has_an_empty_body() {
let request =
request_from_parts("GET", "http://app.example/", std::iter::empty(), None).unwrap();
let body = request.into_body().collect().await.unwrap().to_bytes();
assert!(body.is_empty());
}
#[test]
fn cookie_header_joins_the_jar_and_is_absent_when_empty() {
let cookie = |name: &str, value: &str| Cookie {
name: name.to_string(),
value: value.to_string(),
domain: "app.test".to_string(),
path: "/".to_string(),
expires: -1.0,
http_only: false,
secure: false,
same_site: None,
};
assert_eq!(cookie_header(&[]), None);
assert_eq!(
cookie_header(&[cookie("session", "abc"), cookie("theme", "dark")]).as_deref(),
Some("session=abc; theme=dark")
);
}
#[test]
fn request_carries_non_ascii_header_values() {
let request = request_from_parts(
"GET",
"https://app.example/",
[("x-name", "r\u{e9}sum\u{e9}")],
None,
)
.unwrap();
assert_eq!(
request.headers()["x-name"].as_bytes(),
"r\u{e9}sum\u{e9}".as_bytes()
);
}
#[test]
fn request_keeps_header_order_and_repeats() {
let request = request_from_parts(
"GET",
"https://app.example/",
[("cookie", "a=1"), ("accept", "*/*"), ("cookie", "b=2")],
None,
)
.unwrap();
let cookies: Vec<&str> = request
.headers()
.get_all("cookie")
.iter()
.map(|v| v.to_str().unwrap())
.collect();
assert_eq!(cookies, ["a=1", "b=2"]);
assert_eq!(request.headers().len(), 3);
}
#[test]
fn request_skips_pseudo_headers_and_keeps_the_rest() {
let request = request_from_parts(
"GET",
"https://app.example/",
[(":authority", "app.example"), ("accept", "*/*")],
None,
)
.unwrap();
assert_eq!(request.headers().len(), 1);
assert_eq!(request.headers()["accept"], "*/*");
}
#[test]
fn request_rejects_an_unparseable_url() {
let err = request_from_parts("GET", "not a url", std::iter::empty(), None).unwrap_err();
assert!(matches!(err, Error::ProtocolError(msg) if msg.contains("not a url")));
}
#[test]
fn request_rejects_an_invalid_method() {
let err = request_from_parts("G ET", "http://app.example/", std::iter::empty(), None)
.unwrap_err();
assert!(matches!(err, Error::ProtocolError(msg) if msg.contains("G ET")));
}
#[test]
fn fulfill_maps_status_headers_and_body() {
let mut map = HeaderMap::new();
map.insert("content-type", HeaderValue::from_static("application/json"));
map.insert("x-one", HeaderValue::from_static("1"));
let opts = fulfill_options(StatusCode::CREATED, &map, Bytes::from_static(b"{}"));
assert_eq!(opts.status, Some(201));
assert_eq!(opts.body.as_deref(), Some(&b"{}"[..]));
assert_eq!(opts.content_type, None);
assert_eq!(
opts.headers.unwrap(),
headers_of(&[("content-type", "application/json"), ("x-one", "1")])
);
}
#[test]
fn fulfill_joins_repeated_headers_and_newlines_set_cookie() {
let mut map = HeaderMap::new();
map.append("set-cookie", HeaderValue::from_static("a=1"));
map.append("set-cookie", HeaderValue::from_static("b=2"));
map.append("vary", HeaderValue::from_static("accept"));
map.append("vary", HeaderValue::from_static("origin"));
let merged = fulfill_options(StatusCode::OK, &map, Bytes::new())
.headers
.unwrap();
assert_eq!(merged["set-cookie"], "a=1\nb=2");
assert_eq!(merged["vary"], "accept, origin");
}
#[test]
fn request_refuses_a_body_the_driver_withheld() {
let with_body = request_from_parts(
"POST",
"https://app.example/",
[("content-length", "7")],
Some(b"payload".to_vec()),
)
.unwrap();
assert_eq!(with_body.headers()["content-length"], "7");
let bodiless_get = request_from_parts(
"GET",
"https://app.example/",
[("content-length", "0"), ("accept", "*/*")],
None,
)
.unwrap();
assert_eq!(bodiless_get.headers()["accept"], "*/*");
let withheld = request_from_parts(
"POST",
"https://app.example/upload",
[("content-length", "70000")],
None,
)
.unwrap_err();
assert!(
matches!(withheld, Error::ProtocolError(msg) if msg.contains("70000-byte") && msg.contains("withheld")),
);
}
#[test]
fn fulfill_keeps_non_ascii_header_values() {
let mut map = HeaderMap::new();
map.insert(
"content-disposition",
HeaderValue::from_bytes("attachment; filename=\"r\u{e9}sum\u{e9}.pdf\"".as_bytes())
.unwrap(),
);
let merged = fulfill_options(StatusCode::OK, &map, Bytes::new())
.headers
.unwrap();
assert_eq!(
merged["content-disposition"],
"attachment; filename=\"r\u{e9}sum\u{e9}.pdf\""
);
}
#[test]
fn fulfill_drops_framing_headers_the_body_makes_stale() {
let mut map = HeaderMap::new();
map.insert("content-length", HeaderValue::from_static("999"));
map.insert("transfer-encoding", HeaderValue::from_static("chunked"));
map.insert("etag", HeaderValue::from_static("\"v1\""));
let merged = fulfill_options(StatusCode::OK, &map, Bytes::from_static(b"abc"))
.headers
.unwrap();
assert_eq!(merged, headers_of(&[("etag", "\"v1\"")]));
}
}