use std::borrow::Cow;
use std::time::Duration;
use crate::{Error, Result};
use http::HeaderMap;
use http::HeaderValue;
use http::Method;
use http::header::HeaderName;
use http::uri::Authority;
use http::uri::PathAndQuery;
use http::uri::Scheme;
fn parse_query(query: &str) -> Vec<(String, String)> {
query
.split('&')
.filter(|pair| !pair.is_empty())
.map(|pair| {
let (key, value) = pair.split_once('=').unwrap_or((pair, ""));
(
percent_encoding::percent_decode_str(key)
.decode_utf8_lossy()
.into_owned(),
percent_encoding::percent_decode_str(value)
.decode_utf8_lossy()
.into_owned(),
)
})
.collect()
}
#[derive(Debug)]
pub struct SigningRequest {
pub method: Method,
pub scheme: Scheme,
pub authority: Authority,
pub path: String,
pub query: Vec<(String, String)>,
pub headers: HeaderMap,
}
impl SigningRequest {
pub fn build(parts: &mut http::request::Parts) -> Result<Self> {
let uri = parts.uri.clone().into_parts();
let paq = uri
.path_and_query
.unwrap_or_else(|| PathAndQuery::from_static("/"));
Ok(SigningRequest {
method: parts.method.clone(),
scheme: uri.scheme.unwrap_or(Scheme::HTTP),
authority: uri.authority.ok_or_else(|| {
Error::request_invalid("request without authority is invalid for signing")
})?,
path: paq.path().to_string(),
query: paq.query().map(parse_query).unwrap_or_default(),
headers: parts.headers.clone(),
})
}
pub fn apply(self, parts: &mut http::request::Parts) -> Result<()> {
#[cfg(debug_assertions)]
self.validate_request_view(parts)?;
parts.headers = self.headers;
Ok(())
}
#[cfg(debug_assertions)]
fn validate_request_view(&self, parts: &http::request::Parts) -> Result<()> {
let uri = parts.uri.clone().into_parts();
let paq = uri
.path_and_query
.unwrap_or_else(|| PathAndQuery::from_static("/"));
let scheme = uri.scheme.unwrap_or(Scheme::HTTP);
let authority = uri.authority.ok_or_else(|| {
Error::request_invalid("request without authority is invalid for signing")
})?;
let query = paq.query().map(parse_query).unwrap_or_default();
if self.method != parts.method
|| self.scheme != scheme
|| self.authority != authority
|| self.path != paq.path()
|| self.query != query
{
return Err(Error::request_invalid(
"signing request method or URI working view was modified",
));
}
Ok(())
}
pub fn path_percent_decoded(&self) -> Cow<'_, str> {
percent_encoding::percent_decode_str(&self.path).decode_utf8_lossy()
}
#[inline]
pub fn query_size(&self) -> usize {
self.query
.iter()
.map(|(k, v)| k.len() + v.len())
.sum::<usize>()
}
#[inline]
pub fn query_push(&mut self, key: impl Into<String>, value: impl Into<String>) {
self.query.push((key.into(), value.into()));
}
#[inline]
pub fn query_append(&mut self, query: &str) {
self.query.push((query.to_string(), "".to_string()));
}
pub fn query_to_vec_with_filter(&self, filter: impl Fn(&str) -> bool) -> Vec<(String, String)> {
self.query
.iter()
.filter(|(k, _)| filter(k))
.map(|(k, v)| (k.to_string(), v.to_string()))
.collect()
}
pub fn query_to_string(mut query: Vec<(String, String)>, sep: &str, join: &str) -> String {
let mut s = String::with_capacity(16);
query.sort();
for (idx, (k, v)) in query.into_iter().enumerate() {
if idx != 0 {
s.push_str(join);
}
s.push_str(&k);
if !v.is_empty() {
s.push_str(sep);
s.push_str(&v);
}
}
s
}
pub fn query_to_percent_decoded_string(
mut query: Vec<(String, String)>,
sep: &str,
join: &str,
) -> String {
let mut s = String::with_capacity(16);
query.sort();
for (idx, (k, v)) in query.into_iter().enumerate() {
if idx != 0 {
s.push_str(join);
}
s.push_str(&k);
if !v.is_empty() {
s.push_str(sep);
s.push_str(&percent_encoding::percent_decode_str(&v).decode_utf8_lossy());
}
}
s
}
#[inline]
pub fn header_get_or_default(&self, key: &HeaderName) -> Result<&str> {
match self.headers.get(key) {
Some(v) => v
.to_str()
.map_err(|e| Error::request_invalid("invalid header value").with_source(e)),
None => Ok(""),
}
}
pub fn header_value_normalize(v: &mut HeaderValue) {
let bs = v.as_bytes();
let starting_index = bs.iter().position(|b| *b != b' ').unwrap_or(0);
let ending_offset = bs.iter().rev().position(|b| *b != b' ').unwrap_or(0);
let ending_index = bs.len() - ending_offset;
*v = HeaderValue::from_bytes(&bs[starting_index..ending_index])
.expect("invalid header value")
}
pub fn header_name_to_vec_sorted(&self) -> Vec<&str> {
let mut h = self
.headers
.keys()
.map(|k| k.as_str())
.collect::<Vec<&str>>();
h.sort_unstable();
h
}
pub fn header_to_vec_with_prefix(&self, prefix: &str) -> Vec<(String, String)> {
self.headers
.iter()
.filter(|(k, _)| k.as_str().starts_with(prefix))
.map(|(k, v)| {
(
k.as_str().to_lowercase(),
v.to_str().expect("must be valid header").to_string(),
)
})
.collect()
}
pub fn header_to_string(mut headers: Vec<(String, String)>, sep: &str, join: &str) -> String {
let mut s = String::with_capacity(16);
headers.sort();
for (idx, (k, v)) in headers.into_iter().enumerate() {
if idx != 0 {
s.push_str(join);
}
s.push_str(&k);
s.push_str(sep);
s.push_str(&v);
}
s
}
}
#[derive(Copy, Clone, PartialEq, Eq)]
pub enum SigningMethod {
Header,
Query(Duration),
}
#[cfg(test)]
mod tests {
use super::*;
use http::{HeaderValue, Request};
const RAW_QUERY: &str = "slash=%2F&hash=%23&=%26&equals=%3D&space=%20&encoded-plus=%2B&literal-plus=+&double=%252F&dup=first&dup=second&=empty-key&empty=&flag&flag=&";
fn request_parts() -> http::request::Parts {
Request::get(format!("https://example.com/object%2Fname?{RAW_QUERY}"))
.header("x-original", " value ")
.body(())
.expect("request must build")
.into_parts()
.0
}
#[test]
fn build_is_read_only_and_parses_wire_query_once() {
let mut parts = request_parts();
let original = parts.clone();
let signing = SigningRequest::build(&mut parts).expect("signing request must build");
assert_eq!(parts.method, original.method);
assert_eq!(parts.uri, original.uri);
assert_eq!(parts.version, original.version);
assert_eq!(parts.headers, original.headers);
assert_eq!(signing.path, "/object%2Fname");
assert_eq!(
signing.query,
vec![
("slash".to_string(), "/".to_string()),
("hash".to_string(), "#".to_string()),
("amp".to_string(), "&".to_string()),
("equals".to_string(), "=".to_string()),
("space".to_string(), " ".to_string()),
("encoded-plus".to_string(), "+".to_string()),
("literal-plus".to_string(), "+".to_string()),
("double".to_string(), "%2F".to_string()),
("dup".to_string(), "first".to_string()),
("dup".to_string(), "second".to_string()),
(String::new(), "empty-key".to_string()),
("empty".to_string(), String::new()),
("flag".to_string(), String::new()),
("flag".to_string(), String::new()),
]
);
}
#[test]
fn build_error_leaves_request_unchanged() {
let mut parts = Request::get("/relative")
.header("x-original", "value")
.body(())
.expect("request must build")
.into_parts()
.0;
let original = parts.clone();
assert!(SigningRequest::build(&mut parts).is_err());
assert_eq!(parts.method, original.method);
assert_eq!(parts.uri, original.uri);
assert_eq!(parts.version, original.version);
assert_eq!(parts.headers, original.headers);
}
#[test]
fn apply_commits_only_headers() {
let mut parts = request_parts();
let original = parts.clone();
let mut signing =
SigningRequest::build(&mut parts).expect("signing request must build successfully");
signing
.headers
.insert("authorization", HeaderValue::from_static("signed"));
signing.apply(&mut parts).expect("apply must succeed");
assert_eq!(parts.method, original.method);
assert_eq!(parts.uri, original.uri);
assert_eq!(parts.version, original.version);
assert_eq!(
parts.headers.get("authorization"),
Some(&HeaderValue::from_static("signed"))
);
}
#[cfg(debug_assertions)]
#[test]
fn apply_rejects_modified_request_view_atomically() {
type ViewMutation = Box<dyn Fn(&mut SigningRequest)>;
let mutations: Vec<ViewMutation> = vec![
Box::new(|signing| signing.method = Method::POST),
Box::new(|signing| signing.scheme = Scheme::HTTP),
Box::new(|signing| signing.authority = "other.example.com".parse().unwrap()),
Box::new(|signing| signing.path.push_str("/changed")),
Box::new(|signing| signing.query_push("auth", "value")),
];
for mutate in mutations {
let mut parts = request_parts();
let original = parts.clone();
let mut signing =
SigningRequest::build(&mut parts).expect("signing request must build successfully");
signing
.headers
.insert("authorization", HeaderValue::from_static("signed"));
mutate(&mut signing);
assert!(signing.apply(&mut parts).is_err());
assert_eq!(parts.method, original.method);
assert_eq!(parts.uri, original.uri);
assert_eq!(parts.version, original.version);
assert_eq!(parts.headers, original.headers);
}
}
#[cfg(not(debug_assertions))]
#[test]
fn apply_omits_request_view_validation_in_release() {
let mut parts = request_parts();
let original = parts.clone();
let mut signing =
SigningRequest::build(&mut parts).expect("signing request must build successfully");
signing.method = Method::POST;
signing.scheme = Scheme::HTTP;
signing.authority = "other.example.com".parse().unwrap();
signing.path.push_str("/changed");
signing.query_push("auth", "value");
signing
.headers
.insert("authorization", HeaderValue::from_static("signed"));
signing.apply(&mut parts).expect("apply must succeed");
assert_eq!(parts.method, original.method);
assert_eq!(parts.uri, original.uri);
assert_eq!(parts.version, original.version);
assert_eq!(
parts.headers.get("authorization"),
Some(&HeaderValue::from_static("signed"))
);
}
}