Skip to main content

mcpmem_webhook/
lib.rs

1//! Bounded, lease-fenced webhook delivery. Network implementation is injected.
2use hmac::{Hmac, Mac};
3use mcpmem_core::events::{EventDelivery, EventRepository, now_us};
4use mcpmem_core::subscriptions::{SubscriptionRepository, WebhookSubscription};
5use rusqlite::Connection;
6use serde::Serialize;
7use sha2::Sha256;
8use std::collections::{BTreeMap, BTreeSet};
9use std::net::{IpAddr, SocketAddr, ToSocketAddrs};
10use std::path::{Path, PathBuf};
11use thiserror::Error;
12use url::Url;
13
14const LEASE_US: i64 = 30_000_000;
15const MAX_ATTEMPTS: i64 = 8;
16const MAX_BODY: usize = 65_536;
17
18#[derive(Debug, Error)]
19pub enum WorkerError {
20    #[error("database: {0}")]
21    Database(#[from] rusqlite::Error),
22    #[error("core: {0}")]
23    Core(#[from] mcpmem_core::errors::MCSError),
24    #[error("policy: {0}")]
25    Policy(String),
26    #[error("secret: {0}")]
27    Secret(String),
28    #[error("delivery: {0}")]
29    Delivery(String),
30}
31#[derive(Clone, Debug)]
32pub struct SigningKey(Vec<u8>);
33impl SigningKey {
34    pub fn new(bytes: Vec<u8>) -> Result<Self, WorkerError> {
35        if bytes.is_empty() {
36            Err(WorkerError::Secret("empty signing key".into()))
37        } else {
38            Ok(Self(bytes))
39        }
40    }
41}
42pub trait SecretProvider: Send + Sync {
43    fn signing_key(&self, reference: &str) -> Result<SigningKey, WorkerError>;
44}
45pub trait Resolver: Send + Sync {
46    fn resolve(&self, hostname: &str) -> Result<Vec<IpAddr>, WorkerError>;
47}
48pub struct SystemResolver;
49impl Resolver for SystemResolver {
50    fn resolve(&self, hostname: &str) -> Result<Vec<IpAddr>, WorkerError> {
51        (hostname, 443)
52            .to_socket_addrs()
53            .map_err(|e| WorkerError::Policy(format!("DNS resolution failed: {e}")))
54            .map(|addresses| addresses.map(|address| address.ip()).collect())
55    }
56}
57pub struct StaticSecretProvider(pub BTreeMap<String, SigningKey>);
58impl SecretProvider for StaticSecretProvider {
59    fn signing_key(&self, reference: &str) -> Result<SigningKey, WorkerError> {
60        self.0
61            .get(reference)
62            .cloned()
63            .ok_or_else(|| WorkerError::Secret("secret reference is not configured".into()))
64    }
65}
66
67/// The delivery policy the binary reads from its configuration file. The
68/// allowlist names the HTTPS hostnames the worker may deliver to; a `secret`
69/// reference maps to an already-loaded signing key. Empty is the fail-closed
70/// default: no allowed host and no key means the worker refuses everything.
71#[derive(Clone, Debug, Default)]
72pub struct WebhookConfigFile {
73    pub allowlist: BTreeSet<String>,
74    pub secrets: BTreeMap<String, SigningKey>,
75}
76#[derive(Clone, Debug)]
77pub struct ValidatedEndpoint {
78    pub url: Url,
79    pub address: SocketAddr,
80}
81#[derive(Clone, Debug)]
82pub struct SignedRequest {
83    pub body: Vec<u8>,
84    pub event_id: String,
85    pub timestamp_us: i64,
86    pub signature: String,
87}
88#[derive(Clone, Copy, Debug)]
89pub struct DeliveryResponse {
90    pub status: u16,
91    pub retry_after_us: Option<i64>,
92}
93pub trait DeliveryConnector: Send + Sync {
94    fn send(
95        &self,
96        endpoint: &ValidatedEndpoint,
97        request: SignedRequest,
98    ) -> Result<DeliveryResponse, WorkerError>;
99}
100pub trait HttpsTransport: Send + Sync {
101    fn post(
102        &self,
103        endpoint: &ValidatedEndpoint,
104        request: SignedRequest,
105    ) -> Result<DeliveryResponse, WorkerError>;
106}
107pub struct ReqwestTransport;
108#[derive(Clone, Debug)]
109pub struct RequestPlan {
110    pub host: String,
111    pub address: SocketAddr,
112    pub follow_redirects: bool,
113    pub idempotency_key: String,
114    pub timestamp: String,
115    pub signature: String,
116}
117pub fn request_plan(
118    endpoint: &ValidatedEndpoint,
119    request: &SignedRequest,
120) -> Result<RequestPlan, WorkerError> {
121    Ok(RequestPlan {
122        host: endpoint
123            .url
124            .host_str()
125            .ok_or_else(|| WorkerError::Policy("validated endpoint lacks hostname".into()))?
126            .to_owned(),
127        address: endpoint.address,
128        follow_redirects: false,
129        idempotency_key: request.event_id.clone(),
130        timestamp: request.timestamp_us.to_string(),
131        signature: request.signature.clone(),
132    })
133}
134impl HttpsTransport for ReqwestTransport {
135    fn post(
136        &self,
137        endpoint: &ValidatedEndpoint,
138        request: SignedRequest,
139    ) -> Result<DeliveryResponse, WorkerError> {
140        let plan = request_plan(endpoint, &request)?;
141        let client = reqwest::blocking::Client::builder()
142            .redirect(if plan.follow_redirects {
143                reqwest::redirect::Policy::limited(10)
144            } else {
145                reqwest::redirect::Policy::none()
146            })
147            .resolve(&plan.host, plan.address)
148            .build()
149            .map_err(|e| WorkerError::Delivery(e.to_string()))?;
150        let response = client
151            .post(endpoint.url.clone())
152            .header("Idempotency-Key", plan.idempotency_key)
153            .header("X-Memory-Timestamp", plan.timestamp)
154            .header("X-Memory-Signature", plan.signature)
155            .body(request.body)
156            .send()
157            .map_err(|e| WorkerError::Delivery(e.to_string()))?;
158        let retry_after_us = response
159            .headers()
160            .get(reqwest::header::RETRY_AFTER)
161            .and_then(|value| value.to_str().ok())
162            .and_then(|value| value.parse::<i64>().ok())
163            .map(|seconds| seconds.saturating_mul(1_000_000));
164        Ok(DeliveryResponse {
165            status: response.status().as_u16(),
166            retry_after_us,
167        })
168    }
169}
170pub struct HttpsConnector<T = ReqwestTransport>(pub T);
171impl HttpsConnector {
172    pub const fn production() -> Self {
173        Self(ReqwestTransport)
174    }
175}
176impl<T: HttpsTransport> DeliveryConnector for HttpsConnector<T> {
177    fn send(
178        &self,
179        endpoint: &ValidatedEndpoint,
180        request: SignedRequest,
181    ) -> Result<DeliveryResponse, WorkerError> {
182        self.0.post(endpoint, request)
183    }
184}
185
186pub fn validate_endpoint(
187    endpoint: &str,
188    allowlist: &BTreeSet<String>,
189    resolver: &dyn Resolver,
190) -> Result<ValidatedEndpoint, WorkerError> {
191    let url = Url::parse(endpoint).map_err(|e| WorkerError::Policy(e.to_string()))?;
192    if url.scheme() != "https"
193        || url.port_or_known_default() != Some(443)
194        || url.username() != ""
195        || url.password().is_some()
196        || url.fragment().is_some()
197        || url.host_str().is_none()
198        || url
199            .host_str()
200            .is_some_and(|host| host.parse::<IpAddr>().is_ok())
201    {
202        return Err(WorkerError::Policy(
203            "endpoint must be https, port 443, hostname-only, without fragment".into(),
204        ));
205    }
206    let hostname = url.host_str().expect("checked hostname");
207    if !allowlist.contains(hostname) {
208        return Err(WorkerError::Policy(
209            "endpoint hostname is not allowlisted".into(),
210        ));
211    }
212    let addresses = resolver.resolve(hostname)?;
213    let ip = addresses
214        .first()
215        .copied()
216        .ok_or_else(|| WorkerError::Policy("endpoint resolution returned no address".into()))?;
217    if !is_public(ip) {
218        return Err(WorkerError::Policy(
219            "endpoint resolved to non-public address".into(),
220        ));
221    }
222    if !addresses.iter().copied().all(is_public) {
223        return Err(WorkerError::Policy(
224            "endpoint resolution contains non-public address".into(),
225        ));
226    }
227    Ok(ValidatedEndpoint {
228        url,
229        address: SocketAddr::new(ip, 443),
230    })
231}
232const fn is_public(ip: IpAddr) -> bool {
233    match ip {
234        IpAddr::V4(v) => {
235            !(v.is_private()
236                || v.is_loopback()
237                || v.is_link_local()
238                || v.is_broadcast()
239                || v.is_unspecified()
240                || v.is_multicast()
241                || v.octets()[0] == 0
242                || v.octets()[0] >= 224)
243        }
244        IpAddr::V6(v) => {
245            !(v.is_loopback()
246                || v.is_unspecified()
247                || v.is_multicast()
248                || v.is_unique_local()
249                || v.is_unicast_link_local())
250        }
251    }
252}
253
254#[derive(Debug, Default, Eq, PartialEq)]
255pub struct DeliveryReport {
256    pub claimed: usize,
257    pub completed: usize,
258    pub retried: usize,
259    pub dead: usize,
260}
261pub struct WebhookWorker<C, S, R> {
262    database: PathBuf,
263    connector: C,
264    secrets: S,
265    allowlist: BTreeSet<String>,
266    resolver: R,
267    lease_us: i64,
268}
269pub trait WorkerPoll: Send + Sync {
270    fn poll(&self, now_us: i64) -> Result<DeliveryReport, WorkerError>;
271}
272impl<C: DeliveryConnector, S: SecretProvider, R: Resolver> WebhookWorker<C, S, R> {
273    pub fn new(
274        database: impl AsRef<Path>,
275        connector: C,
276        secrets: S,
277        allowlist: BTreeSet<String>,
278        resolver: R,
279    ) -> Self {
280        Self {
281            database: database.as_ref().to_path_buf(),
282            connector,
283            secrets,
284            allowlist,
285            resolver,
286            lease_us: LEASE_US,
287        }
288    }
289    pub const fn with_lease_us(mut self, lease_us: i64) -> Self {
290        self.lease_us = lease_us;
291        self
292    }
293    pub fn run_once(&self, now: i64) -> Result<DeliveryReport, WorkerError> {
294        let conn = Connection::open(&self.database)?;
295        mcpmem_core::schema::initialize_database(&conn)?;
296        let events = EventRepository::new(&conn);
297        let Some(delivery) = events.claim_due(now, self.lease_us)? else {
298            return Ok(DeliveryReport::default());
299        };
300        let mut report = DeliveryReport {
301            claimed: 1,
302            ..DeliveryReport::default()
303        };
304        let outcome = SubscriptionRepository::new(&conn)
305            .get(delivery.subscription_id)?
306            .filter(|s| s.enabled)
307            .ok_or_else(|| WorkerError::Policy("subscription missing or disabled".into()))
308            .and_then(|subscription| self.deliver(&subscription, &delivery, now));
309        match outcome {
310            Ok(response) if (200..300).contains(&response.status) => {
311                if events.complete(&delivery, now_us())? {
312                    report.completed = 1;
313                }
314            }
315            Ok(response) => {
316                let dead =
317                    !(response.status == 408 || response.status == 429 || response.status >= 500)
318                        || delivery.attempts >= MAX_ATTEMPTS;
319                let delay = response
320                    .retry_after_us
321                    .unwrap_or(1_000_000)
322                    .clamp(1_000_000, 3_600_000_000);
323                if events.retry(
324                    &delivery,
325                    now_us(),
326                    now.saturating_add(delay),
327                    &format!("http {}", response.status),
328                    dead,
329                )? {
330                    if dead {
331                        report.dead = 1
332                    } else {
333                        report.retried = 1
334                    }
335                }
336            }
337            Err(error) => {
338                let dead = delivery.attempts >= MAX_ATTEMPTS
339                    || matches!(error, WorkerError::Policy(_) | WorkerError::Secret(_));
340                if events.retry(
341                    &delivery,
342                    now_us(),
343                    now.saturating_add(1_000_000),
344                    &error.to_string(),
345                    dead,
346                )? {
347                    if dead {
348                        report.dead = 1
349                    } else {
350                        report.retried = 1
351                    }
352                }
353            }
354        }
355        Ok(report)
356    }
357    fn deliver(
358        &self,
359        subscription: &WebhookSubscription,
360        delivery: &EventDelivery,
361        now: i64,
362    ) -> Result<DeliveryResponse, WorkerError> {
363        let endpoint = validate_endpoint(&subscription.endpoint, &self.allowlist, &self.resolver)?;
364        let body = envelope(delivery)?;
365        let key = self.secrets.signing_key(&subscription.secret_ref)?;
366        let signature = signature(&key, now, &body)?;
367        self.connector.send(
368            &endpoint,
369            SignedRequest {
370                body,
371                event_id: delivery.event.event_id.to_string(),
372                timestamp_us: now,
373                signature,
374            },
375        )
376    }
377}
378impl<C: DeliveryConnector, S: SecretProvider, R: Resolver> WorkerPoll for WebhookWorker<C, S, R> {
379    fn poll(&self, now_us: i64) -> Result<DeliveryReport, WorkerError> {
380        self.run_once(now_us)
381    }
382}
383#[derive(Serialize)]
384#[serde(rename_all = "camelCase")]
385struct Envelope<'a> {
386    version: u8,
387    event_id: String,
388    transaction_id: String,
389    entity_id: i64,
390    entity_revision: i64,
391    operation: mcpmem_core::mutation::ChangeOperation,
392    occurred_at_us: i64,
393    origin: &'a str,
394    correlation_id: String,
395    causation_id: Option<String>,
396    hop_count: u8,
397    old_name: Option<&'a str>,
398    new_name: Option<&'a str>,
399}
400fn envelope(delivery: &EventDelivery) -> Result<Vec<u8>, WorkerError> {
401    let event = &delivery.event;
402    let body = serde_json::to_vec(&Envelope {
403        version: 2,
404        event_id: event.event_id.to_string(),
405        transaction_id: event.transaction_id.to_string(),
406        entity_id: event.entity_id,
407        entity_revision: event.entity_revision,
408        operation: event.change.operation,
409        occurred_at_us: event.occurred_at_us,
410        origin: &event.provenance.origin,
411        correlation_id: event.provenance.correlation_id.to_string(),
412        causation_id: event.provenance.causation_id.map(|id| id.to_string()),
413        hop_count: event.provenance.hop_count,
414        old_name: (event.change.operation == mcpmem_core::mutation::ChangeOperation::Rename)
415            .then_some(event.change.old_name.as_deref())
416            .flatten(),
417        new_name: (event.change.operation == mcpmem_core::mutation::ChangeOperation::Rename)
418            .then_some(event.change.new_name.as_deref())
419            .flatten(),
420    })
421    .map_err(mcpmem_core::errors::MCSError::from)?;
422    if body.len() > MAX_BODY {
423        return Err(WorkerError::Policy("envelope exceeds 64KiB".into()));
424    }
425    Ok(body)
426}
427fn signature(key: &SigningKey, timestamp: i64, body: &[u8]) -> Result<String, WorkerError> {
428    let mut mac =
429        Hmac::<Sha256>::new_from_slice(&key.0).map_err(|e| WorkerError::Secret(e.to_string()))?;
430    mac.update(timestamp.to_string().as_bytes());
431    mac.update(b".");
432    mac.update(body);
433    Ok(hex(mac.finalize().into_bytes().as_slice()))
434}
435fn hex(bytes: &[u8]) -> String {
436    bytes.iter().map(|b| format!("{b:02x}")).collect()
437}