use std::collections::HashSet;
use base64::Engine as _;
use http::header::{HeaderMap, HeaderName, HeaderValue};
use http::{Method, StatusCode, Uri};
use super::{
MessageSignatureBase, MessageSignatureComponent, MessageSignatureContext,
MessageSignatureError, MessageSignatureHeaders, MessageSignatureParams,
MessageSignatureRequestContext, MessageSignatureResponseContext, MessageSignatureSigner,
MessageSignatureStructuredFieldType,
};
#[derive(Clone, Debug, Eq, PartialEq)]
#[non_exhaustive]
pub struct MessageSignatureConfig {
label: String,
components: Vec<MessageSignatureComponent>,
params: MessageSignatureParams,
}
impl MessageSignatureConfig {
pub fn new(label: impl Into<String>) -> Result<Self, MessageSignatureError> {
let label = label.into();
validate_label(&label)?;
Ok(Self {
label,
components: Vec::new(),
params: MessageSignatureParams::default(),
})
}
pub fn label(&self) -> &str {
&self.label
}
pub fn components(&self) -> &[MessageSignatureComponent] {
&self.components
}
pub fn params(&self) -> &MessageSignatureParams {
&self.params
}
pub fn component(mut self, component: MessageSignatureComponent) -> Self {
self.components.push(component);
self
}
pub fn components_iter(
mut self,
components: impl IntoIterator<Item = MessageSignatureComponent>,
) -> Self {
self.components.extend(components);
self
}
pub fn structured_field_type(
mut self,
name: HeaderName,
field_type: MessageSignatureStructuredFieldType,
) -> Self {
for component in &mut self.components {
component.set_structured_field_type_for_header(&name, field_type);
}
self
}
pub fn created(mut self, created: u64) -> Self {
self.params.created = Some(created);
self
}
pub fn expires(mut self, expires: u64) -> Self {
self.params.expires = Some(expires);
self
}
pub fn nonce(mut self, nonce: impl Into<String>) -> Self {
self.params.nonce = Some(nonce.into());
self
}
pub fn algorithm(mut self, algorithm: impl Into<String>) -> Self {
self.params.algorithm = Some(algorithm.into());
self
}
pub fn key_id(mut self, key_id: impl Into<String>) -> Self {
self.params.key_id = Some(key_id.into());
self
}
pub fn tag(mut self, tag: impl Into<String>) -> Self {
self.params.tag = Some(tag.into());
self
}
pub fn signature_base(
&self,
method: &Method,
target_uri: &Uri,
request_target: &Uri,
headers: &HeaderMap,
) -> Result<MessageSignatureBase, MessageSignatureError> {
let context = MessageSignatureContext::request(method, target_uri, request_target, headers);
self.signature_base_for_context(&context)
}
pub fn signature_base_for_request_context(
&self,
request: MessageSignatureRequestContext<'_>,
) -> Result<MessageSignatureBase, MessageSignatureError> {
let context = MessageSignatureContext::request_with_trailers(
request.method(),
request.target_uri(),
request.request_target(),
request.headers(),
request.trailers(),
);
self.signature_base_for_context(&context)
}
pub fn response_signature_base(
&self,
status: StatusCode,
headers: &HeaderMap,
) -> Result<MessageSignatureBase, MessageSignatureError> {
let context = MessageSignatureContext::response(status, headers);
self.signature_base_for_context(&context)
}
pub fn response_signature_base_for_context(
&self,
response: MessageSignatureResponseContext<'_>,
) -> Result<MessageSignatureBase, MessageSignatureError> {
let context = MessageSignatureContext::response_with_trailers(
response.status(),
response.headers(),
response.trailers(),
);
self.signature_base_for_context(&context)
}
pub fn request_response_signature_base(
&self,
method: &Method,
target_uri: &Uri,
request_target: &Uri,
request_headers: &HeaderMap,
status: StatusCode,
response_headers: &HeaderMap,
) -> Result<MessageSignatureBase, MessageSignatureError> {
let context = MessageSignatureContext::request_response(
method,
target_uri,
request_target,
request_headers,
status,
response_headers,
);
self.signature_base_for_context(&context)
}
pub fn request_response_signature_base_for_context(
&self,
request: MessageSignatureRequestContext<'_>,
response: MessageSignatureResponseContext<'_>,
) -> Result<MessageSignatureBase, MessageSignatureError> {
let context = MessageSignatureContext::from_request_response_contexts(request, response);
self.signature_base_for_context(&context)
}
pub(crate) fn signature_base_for_context(
&self,
context: &MessageSignatureContext<'_>,
) -> Result<MessageSignatureBase, MessageSignatureError> {
self.validate_components()?;
let mut lines = Vec::with_capacity(self.components.len() + 1);
for component in &self.components {
let identifier = component.identifier()?;
let value = context.component_value(component)?;
ensure_component_value(component, &value)?;
lines.push(format!("{identifier}: {value}"));
}
let signature_params = self.signature_params_value()?;
lines.push(format!("\"@signature-params\": {signature_params}"));
let value = lines.join("\n");
if !value.is_ascii() {
return Err(MessageSignatureError::NonAsciiSignatureBase);
}
Ok(MessageSignatureBase::new(value))
}
pub fn headers_from_signature(
&self,
signature: impl AsRef<[u8]>,
) -> Result<MessageSignatureHeaders, MessageSignatureError> {
self.validate_components()?;
let signature_params = self.signature_params_value()?;
let signature_input = format!("{}={signature_params}", self.label);
let signature = format!(
"{}=:{}:",
self.label,
base64::engine::general_purpose::STANDARD.encode(signature.as_ref())
);
Ok(MessageSignatureHeaders {
label: self.label.clone(),
signature_input: HeaderValue::from_str(&signature_input).map_err(|source| {
MessageSignatureError::InvalidGeneratedHeader {
header: "Signature-Input",
source,
}
})?,
signature: HeaderValue::from_str(&signature).map_err(|source| {
MessageSignatureError::InvalidGeneratedHeader {
header: "Signature",
source,
}
})?,
})
}
pub fn sign_request(
&self,
method: &Method,
target_uri: &Uri,
request_target: &Uri,
headers: &HeaderMap,
signer: &(impl MessageSignatureSigner + ?Sized),
) -> Result<MessageSignatureHeaders, MessageSignatureError> {
let base = self.signature_base(method, target_uri, request_target, headers)?;
let signature = signer.sign(base.as_bytes())?;
self.headers_from_signature(signature)
}
pub fn sign_request_context(
&self,
request: MessageSignatureRequestContext<'_>,
signer: &(impl MessageSignatureSigner + ?Sized),
) -> Result<MessageSignatureHeaders, MessageSignatureError> {
let base = self.signature_base_for_request_context(request)?;
let signature = signer.sign(base.as_bytes())?;
self.headers_from_signature(signature)
}
pub fn sign_response(
&self,
status: StatusCode,
headers: &HeaderMap,
signer: &(impl MessageSignatureSigner + ?Sized),
) -> Result<MessageSignatureHeaders, MessageSignatureError> {
let base = self.response_signature_base(status, headers)?;
let signature = signer.sign(base.as_bytes())?;
self.headers_from_signature(signature)
}
pub fn sign_response_context(
&self,
response: MessageSignatureResponseContext<'_>,
signer: &(impl MessageSignatureSigner + ?Sized),
) -> Result<MessageSignatureHeaders, MessageSignatureError> {
let base = self.response_signature_base_for_context(response)?;
let signature = signer.sign(base.as_bytes())?;
self.headers_from_signature(signature)
}
pub(crate) fn validate_components(&self) -> Result<(), MessageSignatureError> {
validate_component_set(&self.components, true)
}
fn signature_params_value(&self) -> Result<String, MessageSignatureError> {
let mut out = String::new();
out.push('(');
for (index, component) in self.components.iter().enumerate() {
if index > 0 {
out.push(' ');
}
out.push_str(&component.identifier()?);
}
out.push(')');
out.push_str(&self.params.serialize()?);
Ok(out)
}
}
pub(crate) fn validate_component_set(
components: &[MessageSignatureComponent],
allow_empty: bool,
) -> Result<(), MessageSignatureError> {
if components.is_empty() && !allow_empty {
return Err(MessageSignatureError::EmptyComponents);
}
let mut seen = HashSet::new();
let mut seen_dictionary_keys = HashSet::new();
for component in components {
let identifier = component.identifier()?;
if !seen.insert(component.comparison_key()) {
return Err(MessageSignatureError::DuplicateComponent(identifier));
}
if let Some(identity) = component.dictionary_key_identity()
&& !seen_dictionary_keys.insert(identity)
{
return Err(MessageSignatureError::DuplicateComponent(identifier));
}
}
Ok(())
}
pub(crate) fn ensure_component_value(
component: &MessageSignatureComponent,
value: &str,
) -> Result<(), MessageSignatureError> {
if value.contains('\n') || value.contains('\r') {
return Err(MessageSignatureError::NewlineInComponentValue);
}
if value
.chars()
.any(|c| c.is_ascii_control() && (c != '\t' || !component.is_header_field()))
{
return Err(MessageSignatureError::ControlCharacterInComponentValue);
}
if !component.is_header_field()
&& (value.chars().next().is_some_and(char::is_whitespace)
|| value.chars().last().is_some_and(char::is_whitespace))
{
return Err(MessageSignatureError::InvalidDerivedComponentWhitespace);
}
Ok(())
}
pub(crate) fn validate_label(label: &str) -> Result<(), MessageSignatureError> {
let mut chars = label.chars();
let Some(first) = chars.next() else {
return Err(MessageSignatureError::InvalidLabel(label.to_owned()));
};
if !(first.is_ascii_lowercase() || first == '*') {
return Err(MessageSignatureError::InvalidLabel(label.to_owned()));
}
if chars.any(|c| {
!(c.is_ascii_lowercase() || c.is_ascii_digit() || matches!(c, '_' | '-' | '.' | '*'))
}) {
return Err(MessageSignatureError::InvalidLabel(label.to_owned()));
}
Ok(())
}