Skip to main content

renox_core/
webhook.rs

1//! Webhooks: calls from payment gateways and other services, received
2//! safely. For each call Renox:
3//!
4//! 1. checks it came from the provider ([`Webhook::verify`], usually a
5//!    signature over the raw body), answering 401 otherwise;
6//! 2. stores it in `webhook_calls`, once per provider event id, so the
7//!    provider's retries of the same event are answered 200 and not
8//!    processed twice;
9//! 3. answers 200 at once and runs [`Webhook::handle`] in a queue worker,
10//!    retrying on errors; `webhook:failed` lists what failed and
11//!    `webhook:retry <id>` runs a call again.
12//!
13//! ```
14//! # use renox::prelude::*;
15//! # #[derive(serde::Deserialize)] struct Invoice { id: String, status: String, external_id: String }
16//! # async fn mark_paid(_: &Db, _: &str) -> Result { Ok(()) }
17//! use renox::webhook;
18//!
19//! struct Xendit;
20//!
21//! impl Webhook for Xendit {
22//!     const PROVIDER: &'static str = "xendit";
23//!
24//!     fn verify(req: &WebhookRequest, state: &AppState) -> Result {
25//!         let token = webhook::secret(state, "XENDIT_CALLBACK_TOKEN")?;
26//!         webhook::ensure(req.header("x-callback-token").is_some_and(|t| webhook::same(t, &token)))
27//!     }
28//!
29//!     fn event_id(req: &WebhookRequest) -> Result<String> {
30//!         let invoice: Invoice = req.json()?;
31//!         Ok(format!("{}:{}", invoice.id, invoice.status))
32//!     }
33//!
34//!     async fn handle(call: WebhookCall, ctx: JobContext) -> Result {
35//!         let invoice: Invoice = call.json()?;
36//!         mark_paid(&ctx.state.db, &invoice.external_id).await
37//!     }
38//! }
39//!
40//! // Module::routes:   Routes::new().webhook::<Xendit>("/webhooks/xendit")
41//! // Module::register: app.webhook::<Xendit>();
42//! ```
43//!
44//! Webhook routes skip CSRF (callers have no session) and keep working in
45//! maintenance mode (calls are stored and processed as usual).
46
47use std::collections::HashMap;
48use std::future::Future;
49use std::pin::Pin;
50use std::sync::Arc;
51use std::time::Duration;
52
53use anyhow::anyhow;
54use axum::body::Bytes;
55use axum::extract::State;
56use axum::http::{HeaderMap, StatusCode};
57use axum::response::{IntoResponse, Response};
58use hmac::{Hmac, KeyInit, Mac};
59use serde::de::DeserializeOwned;
60use serde::{Deserialize, Serialize};
61use sha2::{Digest, Sha256, Sha512};
62
63use crate::db::{Db, Migration};
64use crate::queue::{Job, JobContext, unix_now};
65use crate::{AppState, Error, Result};
66
67pub(crate) const MIGRATIONS: [Migration; 2] = [
68    crate::db::framework_migration!("webhook", "00010101000300_create_webhook_calls_table"),
69    Migration::new(
70        "00010101000301_store_webhook_payloads_as_bytes",
71        include_str!("../migrations/webhook/00010101000301_store_webhook_payloads_as_bytes.up.sql"),
72        Some(include_str!(
73            "../migrations/webhook/00010101000301_store_webhook_payloads_as_bytes.down.sql"
74        )),
75    )
76    .postgres(
77        include_str!(
78            "../migrations/webhook/00010101000301_store_webhook_payloads_as_bytes.postgres.up.sql"
79        ),
80        Some(include_str!(
81            "../migrations/webhook/00010101000301_store_webhook_payloads_as_bytes.postgres.down.sql"
82        )),
83    ),
84];
85
86/// Tells Renox how to receive one provider's webhooks.
87pub trait Webhook: Send + Sync + 'static {
88    /// A short, stable name such as `"midtrans"`; stored with each call.
89    const PROVIDER: &'static str;
90
91    /// Accepts the call only if it really comes from the provider, usually
92    /// by checking a signature over [`WebhookRequest::body`]. An error
93    /// answers 401 and nothing is stored.
94    fn verify(request: &WebhookRequest, state: &AppState) -> Result;
95
96    /// The provider's id for this event. A call whose id was seen before is
97    /// answered 200 without being processed again. Include the status when
98    /// a provider sends several calls for one object (e.g. `"{id}:{status}"`).
99    fn event_id(request: &WebhookRequest) -> Result<String>;
100
101    /// Processes a stored call in a queue worker. An error marks the call
102    /// failed and retries it (five attempts in all).
103    fn handle(call: WebhookCall, ctx: JobContext) -> impl Future<Output = Result> + Send;
104}
105
106/// The incoming call: headers and the raw body, exactly as signed.
107#[derive(Debug, Clone)]
108#[non_exhaustive]
109pub struct WebhookRequest {
110    /// The request's headers.
111    pub headers: HeaderMap,
112    /// The raw body, unparsed.
113    pub body: Bytes,
114}
115
116impl WebhookRequest {
117    /// A header's value; `None` if it is missing or not visible ASCII.
118    pub fn header(&self, name: &str) -> Option<&str> {
119        self.headers.get(name).and_then(|v| v.to_str().ok())
120    }
121
122    /// The body as JSON.
123    pub fn json<T: DeserializeOwned>(&self) -> Result<T> {
124        serde_json::from_slice(&self.body).map_err(|err| Error::BadRequest(err.to_string()))
125    }
126
127    /// The body as a urlencoded form.
128    pub fn form<T: DeserializeOwned>(&self) -> Result<T> {
129        serde_urlencoded::from_bytes(&self.body).map_err(|err| Error::BadRequest(err.to_string()))
130    }
131}
132
133setting_enum! {
134    /// Where a stored webhook call is, in `webhook_calls.status`.
135    pub enum WebhookStatus ("webhook_calls.status") {
136        /// Stored, waiting for its job (or retried).
137        Received = "received",
138        /// Handled without an error.
139        Processed = "processed",
140        /// Its handler failed; `webhook:retry <id>` runs it again.
141        Failed = "failed",
142    }
143}
144
145/// A stored call, as `handle` gets it.
146#[derive(Debug, Clone, Serialize)]
147#[non_exhaustive]
148pub struct WebhookCall {
149    /// The `webhook_calls` row id, for `webhook:retry`.
150    pub id: i64,
151    /// The [`Webhook::PROVIDER`] that received it.
152    pub provider: String,
153    /// The provider's id, or `sha256:` and its hash when longer than 200 bytes.
154    pub event_id: String,
155    /// The body exactly as received.
156    pub payload: Vec<u8>,
157    /// Received, processed or failed.
158    pub status: WebhookStatus,
159    /// The last processing error; `None` unless `failed`.
160    pub error: Option<String>,
161    /// When it arrived.
162    pub received_at: crate::db::DateTime,
163    /// When processing succeeded; `None` until then.
164    pub processed_at: Option<crate::db::DateTime>,
165}
166
167impl WebhookCall {
168    /// The payload as JSON.
169    pub fn json<T: DeserializeOwned>(&self) -> Result<T> {
170        Ok(serde_json::from_slice(&self.payload)?)
171    }
172
173    /// The payload as a urlencoded form.
174    pub fn form<T: DeserializeOwned>(&self) -> Result<T> {
175        Ok(serde_urlencoded::from_bytes(&self.payload)?)
176    }
177
178    /// The payload as text; an error if it isn't UTF-8.
179    pub fn text(&self) -> Result<&str> {
180        Ok(std::str::from_utf8(&self.payload)?)
181    }
182
183    /// The call with this id, if any.
184    pub async fn find(db: &Db, id: i64) -> Result<Option<Self>> {
185        let row = crate::db::sql(format!("SELECT {COLUMNS} FROM webhook_calls WHERE id = ?"))
186            .bind(id)
187            .fetch_optional(db)
188            .await?;
189        row.map(|row| from_row(&row).map_err(Into::into))
190            .transpose()
191    }
192
193    /// Calls whose processing failed, oldest first.
194    pub async fn failed(db: &Db) -> Result<Vec<Self>> {
195        let rows = crate::db::sql(format!(
196            "SELECT {COLUMNS} FROM webhook_calls WHERE status = 'failed' ORDER BY id"
197        ))
198        .fetch_all(db)
199        .await?;
200        Ok(rows
201            .iter()
202            .map(from_row)
203            .collect::<std::result::Result<_, _>>()?)
204    }
205}
206
207const COLUMNS: &str = "id, provider, event_id, payload, status, error, received_at, processed_at";
208
209fn from_row(row: &crate::db::Row) -> std::result::Result<WebhookCall, crate::db::DbError> {
210    Ok(WebhookCall {
211        id: row.try_get("id")?,
212        provider: row.try_get("provider")?,
213        event_id: row.try_get("event_id")?,
214        payload: row.try_get("payload")?,
215        status: WebhookStatus::parse(&row.try_get::<String>("status")?)
216            .map_err(|err| crate::db::DbError::from(sqlx::Error::Decode(err.into())))?,
217        error: row.try_get("error")?,
218        received_at: crate::db::from_unix(row.try_get("received_at")?),
219        processed_at: row
220            .try_get::<Option<i64>>("processed_at")?
221            .map(crate::db::from_unix),
222    })
223}
224
225pub(crate) type HandleFn = Arc<
226    dyn Fn(WebhookCall, JobContext) -> Pin<Box<dyn Future<Output = Result> + Send>> + Send + Sync,
227>;
228
229/// Registered providers and their `handle`.
230pub(crate) type Handlers = Arc<HashMap<&'static str, HandleFn>>;
231
232pub(crate) fn handler<W: Webhook>() -> HandleFn {
233    Arc::new(|call, ctx| Box::pin(W::handle(call, ctx)))
234}
235
236/// The route handler behind `Routes::webhook::<W>(path)`.
237pub(crate) async fn receive<W: Webhook>(
238    State(state): State<AppState>,
239    headers: HeaderMap,
240    body: Bytes,
241) -> Response {
242    let request = WebhookRequest { headers, body };
243    if let Err(err) = W::verify(&request, &state) {
244        tracing::warn!(provider = W::PROVIDER, error = ?err, "webhook refused");
245        return (StatusCode::UNAUTHORIZED, "invalid webhook").into_response();
246    }
247    let event_id = match W::event_id(&request) {
248        Ok(id) if !id.trim().is_empty() => stored_event_id(id),
249        Ok(_) => return Error::BadRequest("the webhook has no event id".into()).into_response(),
250        Err(err) => return err.into_response(),
251    };
252    match store(&state, W::PROVIDER, &event_id, &request.body).await {
253        Ok(Some(id)) => {
254            tracing::info!(provider = W::PROVIDER, event_id, id, "webhook received");
255            (StatusCode::OK, "ok").into_response()
256        }
257        Ok(None) => {
258            tracing::info!(provider = W::PROVIDER, event_id, "webhook already received");
259            (StatusCode::OK, "already received").into_response()
260        }
261        // The provider will send it again.
262        Err(err) => err.into_response(),
263    }
264}
265
266/// Event ids longer than this are stored as their hash, so the unique index
267/// stays small (PostgreSQL refuses index rows over about 2.7 KB).
268const EVENT_ID_MAX: usize = 200;
269
270fn stored_event_id(id: String) -> String {
271    if id.len() <= EVENT_ID_MAX {
272        return id;
273    }
274    format!("sha256:{}", sha256_hex(&id))
275}
276
277/// Stores the call and queues its processing in one transaction; `None`
278/// when this event was stored before.
279async fn store(
280    state: &AppState,
281    provider: &str,
282    event_id: &str,
283    body: &[u8],
284) -> Result<Option<i64>> {
285    let mut tx = state.db.begin().await?;
286    let id: Option<i64> = crate::db::sql(
287        "INSERT INTO webhook_calls (provider, event_id, payload, status, received_at) \
288         VALUES (?, ?, ?, 'received', ?) ON CONFLICT (provider, event_id) DO NOTHING RETURNING id",
289    )
290    .bind(provider)
291    .bind(event_id)
292    .bind(body.to_vec())
293    .bind(unix_now())
294    .scalar_optional(&mut tx)
295    .await?;
296    if let Some(id) = id {
297        state
298            .queue
299            .dispatch_in(&mut tx, ProcessWebhook { call_id: id })
300            .await?;
301    }
302    tx.commit().await?;
303    if id.is_some() {
304        state.queue.wake_workers();
305    }
306    Ok(id)
307}
308
309/// Runs a stored call again (e.g. after fixing a bug); `false` if there's no such call.
310pub async fn retry(state: &AppState, id: i64) -> Result<bool> {
311    let mut tx = state.db.begin().await?;
312    let changed = crate::db::sql(
313        "UPDATE webhook_calls SET status = 'received', error = NULL, processed_at = NULL WHERE id = ?",
314    )
315    .bind(id)
316    .execute(&mut tx)
317    .await?;
318    if changed == 0 {
319        return Ok(false);
320    }
321    state
322        .queue
323        .dispatch_in(&mut tx, ProcessWebhook { call_id: id })
324        .await?;
325    tx.commit().await?;
326    state.queue.wake_workers();
327    Ok(true)
328}
329
330/// The queue job that runs `Webhook::handle` for a stored call.
331#[derive(Serialize, Deserialize)]
332pub(crate) struct ProcessWebhook {
333    call_id: i64,
334}
335
336impl Job for ProcessWebhook {
337    const NAME: &'static str = "renox:webhook";
338    const MAX_ATTEMPTS: u32 = 5;
339
340    fn backoff(attempt: u32) -> Duration {
341        Duration::from_secs(30 * u64::from(attempt))
342    }
343
344    async fn handle(self, ctx: JobContext) -> Result {
345        let db = ctx.state.db.clone();
346        let Some(call) = WebhookCall::find(&db, self.call_id).await? else {
347            return Ok(());
348        };
349        if call.status == WebhookStatus::Processed {
350            return Ok(());
351        }
352        let Some(handle) = ctx.state.webhooks.get(call.provider.as_str()).cloned() else {
353            let error = format!("no webhook `{}` is registered", call.provider);
354            mark(&db, call.id, WebhookStatus::Failed, Some(&error)).await?;
355            return Err(anyhow!(error).into());
356        };
357        let id = call.id;
358        // Its own task, so a panicking handler marks the call failed too.
359        let outcome = match tokio::spawn(crate::clock::carry(handle(call, ctx))).await {
360            Ok(outcome) => outcome,
361            Err(join) => Err(anyhow!(
362                "the webhook handler panicked: {}",
363                join.try_into_panic()
364                    .map(|panic| crate::error::panic_message(&*panic))
365                    .unwrap_or_else(|join| join.to_string())
366            )
367            .into()),
368        };
369        match outcome {
370            Ok(()) => mark(&db, id, WebhookStatus::Processed, None).await,
371            Err(err) => {
372                mark(&db, id, WebhookStatus::Failed, Some(&format!("{err:?}"))).await?;
373                Err(err)
374            }
375        }
376    }
377}
378
379async fn mark(db: &Db, id: i64, status: WebhookStatus, error: Option<&str>) -> Result {
380    let processed_at = (status == WebhookStatus::Processed).then(unix_now);
381    crate::db::sql("UPDATE webhook_calls SET status = ?, error = ?, processed_at = ? WHERE id = ?")
382        .bind(status.as_str())
383        .bind(error)
384        .bind(processed_at)
385        .bind(id)
386        .execute(db)
387        .await?;
388    Ok(())
389}
390
391// Signature helpers.
392
393/// A secret from `.env` (or `Config::vars`), e.g.
394/// `secret(state, "MIDTRANS_SERVER_KEY")`; missing is an error.
395pub fn secret(state: &AppState, name: &str) -> Result<String> {
396    state
397        .config
398        .var(name)
399        .ok_or_else(|| anyhow!("set {name} in .env to receive these webhooks").into())
400}
401
402/// `Ok` if `valid`, otherwise the error `verify` returns for a forged call.
403pub fn ensure(valid: bool) -> Result {
404    if valid {
405        Ok(())
406    } else {
407        Err(Error::Unauthorized)
408    }
409}
410
411/// Compares two strings in time independent of where they differ.
412pub fn same(a: &str, b: &str) -> bool {
413    crate::crypto::constant_time_eq(a, b)
414}
415
416fn hex(bytes: &[u8]) -> String {
417    bytes.iter().map(|b| format!("{b:02x}")).collect()
418}
419
420/// Lowercase hex SHA-256 of `data`.
421pub fn sha256_hex(data: impl AsRef<[u8]>) -> String {
422    hex(&Sha256::digest(data.as_ref()))
423}
424
425/// Lowercase hex SHA-512 of `data` (Midtrans' `signature_key`).
426pub fn sha512_hex(data: impl AsRef<[u8]>) -> String {
427    hex(&Sha512::digest(data.as_ref()))
428}
429
430/// Lowercase hex HMAC-SHA256 of `data` with `key`.
431pub fn hmac_sha256_hex(key: impl AsRef<[u8]>, data: impl AsRef<[u8]>) -> String {
432    let mut mac = <Hmac<Sha256> as KeyInit>::new_from_slice(key.as_ref())
433        .expect("HMAC accepts keys of any length");
434    mac.update(data.as_ref());
435    hex(&mac.finalize().into_bytes())
436}
437
438/// Lowercase hex HMAC-SHA512 of `data` with `key`.
439pub fn hmac_sha512_hex(key: impl AsRef<[u8]>, data: impl AsRef<[u8]>) -> String {
440    let mut mac = <Hmac<Sha512> as KeyInit>::new_from_slice(key.as_ref())
441        .expect("HMAC accepts keys of any length");
442    mac.update(data.as_ref());
443    hex(&mac.finalize().into_bytes())
444}
445
446/// Checks a hex HMAC-SHA256 `signature` of `body`, with or without a
447/// `sha256=` prefix (GitHub, Shopify-style hex, …), ignoring hex case.
448pub fn verify_hmac_sha256(key: impl AsRef<[u8]>, body: &[u8], signature: &str) -> bool {
449    let signature = signature.trim();
450    let signature = signature.strip_prefix("sha256=").unwrap_or(signature);
451    same(&signature.to_ascii_lowercase(), &hmac_sha256_hex(key, body))
452}
453
454/// Checks a Stripe-style signature header, `t=<unix time>,v1=<hex>[,v1=…]`,
455/// where each `v1` is HMAC-SHA256 of `"{t}.{body}"`, and refuses one older or
456/// newer than `tolerance` (Stripe uses five minutes) so it can't be replayed.
457pub fn verify_timestamped(
458    key: impl AsRef<[u8]>,
459    body: &[u8],
460    header: &str,
461    tolerance: Duration,
462) -> bool {
463    let mut timestamp = None;
464    let mut signatures = Vec::new();
465    for part in header.split(',') {
466        match part.trim().split_once('=') {
467            Some(("t", value)) => timestamp = value.parse::<i64>().ok(),
468            Some(("v1", value)) => signatures.push(value.to_ascii_lowercase()),
469            _ => {}
470        }
471    }
472    let Some(timestamp) = timestamp else {
473        return false;
474    };
475    if (unix_now() - timestamp).unsigned_abs() > tolerance.as_secs() {
476        return false;
477    }
478    let mut signed = format!("{timestamp}.").into_bytes();
479    signed.extend_from_slice(body);
480    let expected = hmac_sha256_hex(key, &signed);
481    signatures.iter().any(|s| same(s, &expected))
482}
483
484#[cfg(test)]
485mod tests {
486    use super::*;
487
488    #[test]
489    fn digests_match_known_values() {
490        assert_eq!(
491            sha256_hex("abc"),
492            "ba7816bf8f01cfea414140de5dae2223b00361a396177a9cb410ff61f20015ad"
493        );
494        assert!(sha512_hex("abc").starts_with("ddaf35a193617aba"));
495        // RFC 4231 test case 2.
496        assert_eq!(
497            hmac_sha256_hex("Jefe", "what do ya want for nothing?"),
498            "5bdcc146bf60754e6a042426089575c75a003f089d2739839dec58b964ec3843"
499        );
500        assert!(
501            hmac_sha512_hex("Jefe", "what do ya want for nothing?").starts_with("164b7a7bfcf819e2")
502        );
503    }
504
505    #[test]
506    fn checks_hmac_and_timestamped_signatures() {
507        let body = br#"{"id":"evt_1"}"#;
508        let signature = hmac_sha256_hex("secret", body);
509        assert!(verify_hmac_sha256("secret", body, &signature));
510        assert!(verify_hmac_sha256(
511            "secret",
512            body,
513            &format!("sha256={}", signature.to_uppercase())
514        ));
515        assert!(!verify_hmac_sha256("other", body, &signature));
516
517        let now = unix_now();
518        let signed = |t: i64, key: &str| {
519            let mut data = format!("{t}.").into_bytes();
520            data.extend_from_slice(body);
521            format!("t={t},v1=bad,v1={}", hmac_sha256_hex(key, data))
522        };
523        let five_minutes = Duration::from_secs(300);
524        assert!(verify_timestamped(
525            "secret",
526            body,
527            &signed(now, "secret"),
528            five_minutes
529        ));
530        assert!(!verify_timestamped(
531            "secret",
532            body,
533            &signed(now, "other"),
534            five_minutes
535        ));
536        assert!(!verify_timestamped(
537            "secret",
538            body,
539            &signed(now - 600, "secret"),
540            five_minutes
541        ));
542        assert!(!verify_timestamped(
543            "secret",
544            b"{}",
545            &signed(now, "secret"),
546            five_minutes
547        ));
548        assert!(!verify_timestamped("secret", body, "v1=abc", five_minutes));
549    }
550}