1use 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#[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}