1use 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
22pub const MAX_CLOCK_LEAD_MS: i64 = 30_000;
24
25#[derive(Debug, Clone, PartialEq, Eq)]
27pub struct VerifierKey {
28 pub key_id: String,
30 pub public_key: [u8; 32],
32 pub not_before_ms: Option<i64>,
35 pub not_after_ms: Option<i64>,
38}
39
40impl VerifierKey {
41 #[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 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#[derive(Debug, Clone, Copy, PartialEq, Eq, thiserror::Error)]
116#[error("invalid hook key list: {0}")]
117pub struct KeyListError(&'static str);
118
119#[derive(Debug, Clone, Copy, PartialEq, Eq, thiserror::Error)]
122#[non_exhaustive]
123pub enum VerifyError {
124 #[error("missing or repeated hook header {0}")]
126 Header(&'static str),
127 #[error("unsupported hook signature version")]
129 Version,
130 #[error("hook header {0} is not canonical")]
132 Malformed(&'static str),
133 #[error("unknown hook key id")]
135 UnknownKey,
136 #[error("hook key not valid at the request's creation time")]
138 KeyBounds,
139 #[error("audience is not this hook service's origin")]
141 Audience,
142 #[error("hook validity interval out of range")]
144 Interval,
145 #[error("hook request created in the future")]
147 ClockLead,
148 #[error("hook request expired")]
150 Expired,
151 #[error("hook digest does not match the body")]
153 Digest,
154 #[error("hook signature does not verify")]
156 Signature,
157 #[error("hook nonce replayed")]
159 Replay,
160}
161
162#[derive(Debug, Clone, PartialEq, Eq)]
164#[non_exhaustive]
165pub struct Verified {
166 pub key_id: String,
168 pub nonce: String,
170 pub created_ms: i64,
172 pub expires_ms: i64,
174}
175
176pub struct HookVerifier {
178 audience: String,
179 keys: Vec<VerifierKey>,
180 clock: Arc<dyn Fn() -> i64 + Send + Sync>,
181 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
220fn 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 #[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 #[must_use]
249 pub fn with_replay_protection(mut self) -> Self {
250 self.replay = Some(Mutex::new(HashMap::new()));
251 self
252 }
253
254 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 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}