aioduct 0.2.5

Async-native HTTP client built directly on hyper 1.x — no hyper-util, no legacy
Documentation
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,
};

/// Configuration for generating an RFC 9421 request signature base and headers.
#[derive(Clone, Debug, Eq, PartialEq)]
#[non_exhaustive]
pub struct MessageSignatureConfig {
    label: String,
    components: Vec<MessageSignatureComponent>,
    params: MessageSignatureParams,
}

impl MessageSignatureConfig {
    /// Create a signature configuration for the given signature label.
    ///
    /// The label is serialized as a Structured Fields dictionary key in both
    /// `Signature-Input` and `Signature`.
    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(),
        })
    }

    /// Return the signature label.
    pub fn label(&self) -> &str {
        &self.label
    }

    /// Return the ordered covered component list.
    pub fn components(&self) -> &[MessageSignatureComponent] {
        &self.components
    }

    /// Return the configured signature metadata parameters.
    pub fn params(&self) -> &MessageSignatureParams {
        &self.params
    }

    /// Add one covered component.
    pub fn component(mut self, component: MessageSignatureComponent) -> Self {
        self.components.push(component);
        self
    }

    /// Add covered components in order.
    pub fn components_iter(
        mut self,
        components: impl IntoIterator<Item = MessageSignatureComponent>,
    ) -> Self {
        self.components.extend(components);
        self
    }

    /// Configure the RFC 9651 top-level type for parsed `;sf` components covering this field.
    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
    }

    /// Set the `created` metadata parameter.
    pub fn created(mut self, created: u64) -> Self {
        self.params.created = Some(created);
        self
    }

    /// Set the `expires` metadata parameter.
    pub fn expires(mut self, expires: u64) -> Self {
        self.params.expires = Some(expires);
        self
    }

    /// Set the `nonce` metadata parameter.
    pub fn nonce(mut self, nonce: impl Into<String>) -> Self {
        self.params.nonce = Some(nonce.into());
        self
    }

    /// Set the `alg` metadata parameter.
    pub fn algorithm(mut self, algorithm: impl Into<String>) -> Self {
        self.params.algorithm = Some(algorithm.into());
        self
    }

    /// Set the `keyid` metadata parameter.
    pub fn key_id(mut self, key_id: impl Into<String>) -> Self {
        self.params.key_id = Some(key_id.into());
        self
    }

    /// Set the `tag` metadata parameter.
    pub fn tag(mut self, tag: impl Into<String>) -> Self {
        self.params.tag = Some(tag.into());
        self
    }

    /// Build the RFC 9421 signature base for a request.
    ///
    /// `target_uri` is the full request URI used for scheme, authority, path,
    /// query, and target-uri derived components. `request_target` is the final
    /// request URI form that will be sent on the wire, used for
    /// `@request-target`.
    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)
    }

    /// Build the RFC 9421 signature base for a request context.
    ///
    /// Use `MessageSignatureRequestContext::with_trailers(...)` when the
    /// signature covers caller-supplied trailer fields with `;tr`.
    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)
    }

    /// Build the RFC 9421 signature base for a response.
    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)
    }

    /// Build the RFC 9421 signature base for a response context.
    ///
    /// Use `MessageSignatureResponseContext::with_trailers(...)` when the
    /// signature covers caller-supplied trailer fields with `;tr`.
    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)
    }

    /// Build the RFC 9421 signature base for a response with its related request.
    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)
    }

    /// Build the RFC 9421 signature base for a response context with its related request.
    ///
    /// Use `with_trailers(...)` on either context when the signature covers
    /// caller-supplied trailer fields with `;tr`, including related request
    /// trailer fields with `;req`.
    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))
    }

    /// Format `Signature-Input` and `Signature` header values from signature bytes.
    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,
                }
            })?,
        })
    }

    /// Build the signature base, sign it, and format signature headers.
    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)
    }

    /// Build a request context signature base, sign it, and format signature headers.
    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)
    }

    /// Build a response signature base, sign it, and format signature headers.
    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)
    }

    /// Build a response context signature base, sign it, and format signature headers.
    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(())
}