Skip to main content

mkit_rpc/hooks/
verify.rs

1//! Verifying a signed hook request: the hook service's side of SPEC-SERVER
2//! §7.1, for a service that receives what [`HookSigner`](super::HookSigner)
3//! sends.
4//!
5//! [`HookVerifier`] depends only on the hash, the signature scheme and `std`:
6//! no server runtime, so a Rust hook implementer can use it. It performs
7//! every check §7.1 lists: the eight headers and their canonical forms, the
8//! key id against the key list (§7.2) with its validity bounds, the audience,
9//! the validity window, the digest of the exact body bytes, the strict
10//! Ed25519 signature over the BLAKE3 of the canonical string, and, when asked
11//! to, replay of a nonce inside its window.
12
13use std::collections::HashMap;
14use std::sync::{Arc, Mutex, PoisonError};
15
16use ed25519_dalek::{Signature, VerifyingKey};
17use mkit_core::hash::{hash, to_hex};
18use serde::Deserialize;
19
20use super::sign::{DOMAIN, MAX_VALIDITY};
21
22/// How far the sender's clock may lead the receiver's (SPEC-SERVER §7.1).
23pub const MAX_CLOCK_LEAD_MS: i64 = 30_000;
24
25/// One key of a §7.2 key list.
26#[derive(Debug, Clone, PartialEq, Eq)]
27pub struct VerifierKey {
28    /// The key id (`[A-Za-z0-9._-]`, 1-64 bytes).
29    pub key_id: String,
30    /// The Ed25519 public key.
31    pub public_key: [u8; 32],
32    /// The key is not valid for requests created before this epoch
33    /// millisecond, if set.
34    pub not_before_ms: Option<i64>,
35    /// The key is not valid for requests created after this epoch
36    /// millisecond, if set.
37    pub not_after_ms: Option<i64>,
38}
39
40impl VerifierKey {
41    /// A key with no validity bounds.
42    #[must_use]
43    pub fn new(key_id: impl Into<String>, public_key: [u8; 32]) -> Self {
44        Self {
45            key_id: key_id.into(),
46            public_key,
47            not_before_ms: None,
48            not_after_ms: None,
49        }
50    }
51
52    /// The keys of a §7.2 key-list JSON document.
53    ///
54    /// # Errors
55    /// [`KeyListError`] for anything the format refuses: another `version` or
56    /// `alg`, a public key that is not 64 lowercase hex digits, a bound that
57    /// is not a decimal `int64` string, a key id outside the §7.1 grammar, or
58    /// a repeated key id.
59    pub fn parse_list(json: &str) -> Result<Vec<Self>, KeyListError> {
60        #[derive(Deserialize)]
61        #[serde(rename_all = "camelCase")]
62        struct Entry {
63            key_id: String,
64            alg: String,
65            public_key: String,
66            not_before_ms: Option<String>,
67            not_after_ms: Option<String>,
68        }
69        #[derive(Deserialize)]
70        struct List {
71            version: u32,
72            keys: Vec<Entry>,
73        }
74        let list: List = serde_json::from_str(json).map_err(|_| KeyListError("malformed JSON"))?;
75        if list.version != 1 {
76            return Err(KeyListError("unsupported key-list version"));
77        }
78        let bound = |value: Option<String>| -> Result<Option<i64>, KeyListError> {
79            value
80                .map(|text| {
81                    text.parse::<i64>()
82                        .ok()
83                        .filter(|n| n.to_string() == text)
84                        .ok_or(KeyListError(
85                            "a validity bound is not a decimal int64 string",
86                        ))
87                })
88                .transpose()
89        };
90        let mut keys: Vec<Self> = Vec::with_capacity(list.keys.len());
91        for entry in list.keys {
92            if entry.alg != "ed25519" {
93                return Err(KeyListError("unsupported key algorithm"));
94            }
95            if !valid_key_id(&entry.key_id) {
96                return Err(KeyListError("key id outside [A-Za-z0-9._-], 1-64 bytes"));
97            }
98            if keys.iter().any(|known| known.key_id == entry.key_id) {
99                return Err(KeyListError("repeated key id"));
100            }
101            let public_key = lower_hex_32(&entry.public_key)
102                .ok_or(KeyListError("public key is not 64 lowercase hex digits"))?;
103            keys.push(Self {
104                key_id: entry.key_id,
105                public_key,
106                not_before_ms: bound(entry.not_before_ms)?,
107                not_after_ms: bound(entry.not_after_ms)?,
108            });
109        }
110        Ok(keys)
111    }
112}
113
114/// A refused key list.
115#[derive(Debug, Clone, Copy, PartialEq, Eq, thiserror::Error)]
116#[error("invalid hook key list: {0}")]
117pub struct KeyListError(&'static str);
118
119/// Why a request failed verification. The text is fixed and never quotes the
120/// request.
121#[derive(Debug, Clone, Copy, PartialEq, Eq, thiserror::Error)]
122#[non_exhaustive]
123pub enum VerifyError {
124    /// A required `X-Mkit-Hook-*` header is absent or sent twice.
125    #[error("missing or repeated hook header {0}")]
126    Header(&'static str),
127    /// `X-Mkit-Hook-Version` is not `1`.
128    #[error("unsupported hook signature version")]
129    Version,
130    /// A header value is not in its canonical form.
131    #[error("hook header {0} is not canonical")]
132    Malformed(&'static str),
133    /// The key id is not in the key list.
134    #[error("unknown hook key id")]
135    UnknownKey,
136    /// The request was created outside the key's validity bounds.
137    #[error("hook key not valid at the request's creation time")]
138    KeyBounds,
139    /// The audience is not this hook service's origin.
140    #[error("audience is not this hook service's origin")]
141    Audience,
142    /// The validity interval is not positive, or exceeds 300,000 ms.
143    #[error("hook validity interval out of range")]
144    Interval,
145    /// The request was created more than 30 s ahead of this clock.
146    #[error("hook request created in the future")]
147    ClockLead,
148    /// The request has expired.
149    #[error("hook request expired")]
150    Expired,
151    /// The digest does not match the body.
152    #[error("hook digest does not match the body")]
153    Digest,
154    /// The signature does not verify.
155    #[error("hook signature does not verify")]
156    Signature,
157    /// The nonce was already used inside its window.
158    #[error("hook nonce replayed")]
159    Replay,
160}
161
162/// A request that passed every check.
163#[derive(Debug, Clone, PartialEq, Eq)]
164#[non_exhaustive]
165pub struct Verified {
166    /// The key that signed it.
167    pub key_id: String,
168    /// Its nonce, 64 lowercase hex digits.
169    pub nonce: String,
170    /// Epoch milliseconds it was created.
171    pub created_ms: i64,
172    /// Epoch milliseconds it expires.
173    pub expires_ms: i64,
174}
175
176/// Verifies signed hook requests for one hook service.
177pub struct HookVerifier {
178    audience: String,
179    keys: Vec<VerifierKey>,
180    clock: Arc<dyn Fn() -> i64 + Send + Sync>,
181    /// Nonces seen, each until its request's expiry; `None`: no replay check.
182    replay: Option<Mutex<HashMap<String, i64>>>,
183}
184
185impl core::fmt::Debug for HookVerifier {
186    fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
187        f.debug_struct("HookVerifier")
188            .field("audience", &self.audience)
189            .field("keys", &self.keys.len())
190            .field("replay", &self.replay.is_some())
191            .finish_non_exhaustive()
192    }
193}
194
195fn valid_key_id(id: &str) -> bool {
196    (1..=64).contains(&id.len())
197        && id
198            .bytes()
199            .all(|b| b.is_ascii_alphanumeric() || b"._-".contains(&b))
200}
201
202fn is_lower_hex(value: &str, len: usize) -> bool {
203    value.len() == len
204        && value
205            .bytes()
206            .all(|b| matches!(b, b'0'..=b'9' | b'a'..=b'f'))
207}
208
209fn lower_hex_32(value: &str) -> Option<[u8; 32]> {
210    if !is_lower_hex(value, 64) {
211        return None;
212    }
213    let mut out = [0u8; 32];
214    for (byte, pair) in out.iter_mut().zip(value.as_bytes().chunks(2)) {
215        *byte = u8::from_str_radix(core::str::from_utf8(pair).ok()?, 16).ok()?;
216    }
217    Some(out)
218}
219
220/// Base-10 ASCII digits, no sign, no leading zero (a lone `0` is allowed).
221fn decimal(value: &str) -> Option<i64> {
222    let canonical = !value.is_empty()
223        && value.bytes().all(|b| b.is_ascii_digit())
224        && (value == "0" || !value.starts_with('0'));
225    if canonical { value.parse().ok() } else { None }
226}
227
228impl HookVerifier {
229    /// A verifier for the hook service whose canonical origin is `audience`,
230    /// trusting `keys`, reading time from `clock` (epoch milliseconds).
231    #[must_use]
232    pub fn new(
233        audience: impl Into<String>,
234        keys: Vec<VerifierKey>,
235        clock: impl Fn() -> i64 + Send + Sync + 'static,
236    ) -> Self {
237        Self {
238            audience: audience.into(),
239            keys,
240            clock: Arc::new(clock),
241            replay: None,
242        }
243    }
244
245    /// Also reject a nonce seen before, until its request's window closes
246    /// (SPEC-SERVER §7.1 step 5). The set is in memory: a hook that restarts
247    /// forgets it, and one that runs several instances must share one.
248    #[must_use]
249    pub fn with_replay_protection(mut self) -> Self {
250        self.replay = Some(Mutex::new(HashMap::new()));
251        self
252    }
253
254    /// Verify one request: `procedure` is the Connect path it arrived on,
255    /// `headers` every header it carried (names compare case-insensitively),
256    /// and `body` the exact bytes received.
257    ///
258    /// # Errors
259    /// The first [`VerifyError`] found. A request that fails is never
260    /// remembered for replay.
261    pub fn verify(
262        &self,
263        procedure: &str,
264        headers: &[(&str, &str)],
265        body: &[u8],
266    ) -> Result<Verified, VerifyError> {
267        let get = |name: &'static str| -> Result<&str, VerifyError> {
268            let mut found = headers
269                .iter()
270                .filter(|(n, _)| n.eq_ignore_ascii_case(name))
271                .map(|(_, v)| *v);
272            match (found.next(), found.next()) {
273                (Some(value), None) => Ok(value),
274                _ => Err(VerifyError::Header(name)),
275            }
276        };
277        // Every header first, before the body is looked at (§7.1).
278        let version = get("X-Mkit-Hook-Version")?;
279        let key_id = get("X-Mkit-Hook-Key-Id")?;
280        let audience = get("X-Mkit-Hook-Audience")?;
281        let created = get("X-Mkit-Hook-Created-At")?;
282        let expires = get("X-Mkit-Hook-Expires-At")?;
283        let nonce = get("X-Mkit-Hook-Nonce")?;
284        let digest = get("X-Mkit-Hook-Digest")?;
285        let signature = get("X-Mkit-Hook-Signature")?;
286        if version != "1" {
287            return Err(VerifyError::Version);
288        }
289        let created_ms =
290            decimal(created).ok_or(VerifyError::Malformed("X-Mkit-Hook-Created-At"))?;
291        let expires_ms =
292            decimal(expires).ok_or(VerifyError::Malformed("X-Mkit-Hook-Expires-At"))?;
293        if !is_lower_hex(nonce, 64) {
294            return Err(VerifyError::Malformed("X-Mkit-Hook-Nonce"));
295        }
296        if !valid_key_id(key_id) {
297            return Err(VerifyError::Malformed("X-Mkit-Hook-Key-Id"));
298        }
299        let signature_bytes = if is_lower_hex(signature, 128) {
300            let mut out = [0u8; 64];
301            for (byte, pair) in out.iter_mut().zip(signature.as_bytes().chunks(2)) {
302                *byte = u8::from_str_radix(core::str::from_utf8(pair).unwrap_or("zz"), 16)
303                    .map_err(|_| VerifyError::Malformed("X-Mkit-Hook-Signature"))?;
304            }
305            out
306        } else {
307            return Err(VerifyError::Malformed("X-Mkit-Hook-Signature"));
308        };
309        let key = self
310            .keys
311            .iter()
312            .find(|key| key.key_id == key_id)
313            .ok_or(VerifyError::UnknownKey)?;
314        if key.not_before_ms.is_some_and(|bound| created_ms < bound)
315            || key.not_after_ms.is_some_and(|bound| created_ms > bound)
316        {
317            return Err(VerifyError::KeyBounds);
318        }
319        if audience != self.audience {
320            return Err(VerifyError::Audience);
321        }
322        let now = (self.clock)();
323        let max_ms = i64::try_from(MAX_VALIDITY.as_millis()).unwrap_or(i64::MAX);
324        if expires_ms <= created_ms || expires_ms - created_ms > max_ms {
325            return Err(VerifyError::Interval);
326        }
327        if created_ms > now.saturating_add(MAX_CLOCK_LEAD_MS) {
328            return Err(VerifyError::ClockLead);
329        }
330        if expires_ms <= now {
331            return Err(VerifyError::Expired);
332        }
333        if digest != format!("body:{}", to_hex(&hash(body))) {
334            return Err(VerifyError::Digest);
335        }
336        let canonical = [
337            DOMAIN, key_id, audience, procedure, digest, created, expires, nonce,
338        ]
339        .join("\n");
340        VerifyingKey::from_bytes(&key.public_key)
341            .and_then(|public| {
342                public.verify_strict(
343                    &hash(canonical.as_bytes()),
344                    &Signature::from_bytes(&signature_bytes),
345                )
346            })
347            .map_err(|_| VerifyError::Signature)?;
348        if let Some(seen) = &self.replay {
349            let mut seen = seen.lock().unwrap_or_else(PoisonError::into_inner);
350            seen.retain(|_, until| *until > now);
351            match seen.entry(nonce.to_owned()) {
352                std::collections::hash_map::Entry::Occupied(_) => return Err(VerifyError::Replay),
353                std::collections::hash_map::Entry::Vacant(entry) => {
354                    entry.insert(expires_ms);
355                }
356            }
357        }
358        Ok(Verified {
359            key_id: key_id.to_owned(),
360            nonce: nonce.to_owned(),
361            created_ms,
362            expires_ms,
363        })
364    }
365}