Skip to main content

mcpmem_core/
subscriptions.rs

1//! Durable, network-free webhook subscription policy and SQLite persistence.
2use rusqlite::{Connection, OptionalExtension, params};
3use serde::{Deserialize, Serialize};
4use uuid::Uuid;
5
6use crate::errors::{MCSError, Result};
7use crate::events::{ChangeEvent, now_us, parse_uuid, sql_error};
8use crate::mutation::ChangeOperation;
9
10#[derive(Clone, Debug, Eq, PartialEq, Serialize, Deserialize)]
11#[serde(rename_all = "camelCase")]
12pub struct WebhookSubscription {
13    pub subscription_id: Uuid,
14    pub endpoint: String,
15    pub event_operations: Vec<ChangeOperation>,
16    pub entity_types: Vec<String>,
17    pub ignored_origins: Vec<String>,
18    pub consumer_origin: String,
19    pub secret_ref: String,
20    pub enabled: bool,
21}
22
23impl WebhookSubscription {
24    pub fn validate(self) -> Result<Self> {
25        if self.endpoint.trim().is_empty()
26            || self.consumer_origin.trim().is_empty()
27            || self.secret_ref.trim().is_empty()
28            || self.endpoint.len() > 2048
29            || self.consumer_origin.len() > 256
30            || self.secret_ref.len() > 512
31            || self.consumer_origin.chars().any(char::is_control)
32            || self.secret_ref.chars().any(char::is_control)
33        {
34            return Err(MCSError::InvalidParams(
35                "invalid webhook subscription".into(),
36            ));
37        }
38        Ok(self)
39    }
40}
41
42pub struct SubscriptionRepository<'a> {
43    conn: &'a Connection,
44}
45
46impl<'a> SubscriptionRepository<'a> {
47    pub const fn new(conn: &'a Connection) -> Self {
48        Self { conn }
49    }
50
51    pub fn upsert(&self, subscription: WebhookSubscription) -> Result<()> {
52        let subscription = subscription.validate()?;
53        let now = now_us();
54        self.conn.execute(
55            "INSERT INTO webhook_subscription(subscription_id,endpoint,event_operations,entity_types,ignored_origins,consumer_origin,secret_ref,enabled,created_at_us,updated_at_us) VALUES(?1,?2,?3,?4,?5,?6,?7,?8,?9,?9) ON CONFLICT(subscription_id) DO UPDATE SET endpoint=excluded.endpoint,event_operations=excluded.event_operations,entity_types=excluded.entity_types,ignored_origins=excluded.ignored_origins,consumer_origin=excluded.consumer_origin,secret_ref=excluded.secret_ref,enabled=excluded.enabled,updated_at_us=excluded.updated_at_us",
56            params![subscription.subscription_id.to_string(), subscription.endpoint, serde_json::to_string(&subscription.event_operations)?, serde_json::to_string(&subscription.entity_types)?, serde_json::to_string(&subscription.ignored_origins)?, subscription.consumer_origin, subscription.secret_ref, subscription.enabled, now],
57        ).map_err(sql_error)?;
58        Ok(())
59    }
60
61    pub fn get(&self, id: Uuid) -> Result<Option<WebhookSubscription>> {
62        self.conn.query_row("SELECT subscription_id,endpoint,event_operations,entity_types,ignored_origins,consumer_origin,secret_ref,enabled FROM webhook_subscription WHERE subscription_id=?1", [id.to_string()], row).optional().map_err(sql_error)?.map(decode).transpose()
63    }
64
65    pub fn delete(&self, id: Uuid) -> Result<bool> {
66        self.conn
67            .execute(
68                "DELETE FROM webhook_subscription WHERE subscription_id=?1",
69                [id.to_string()],
70            )
71            .map(|n| n == 1)
72            .map_err(sql_error)
73    }
74
75    /// Every subscription row, oldest first. The admin API lists through
76    /// this; the delivery worker reads [`matching`] instead, which applies
77    /// the enabled and filter predicates.
78    pub fn list(&self) -> Result<Vec<WebhookSubscription>> {
79        let mut statement = self
80            .conn
81            .prepare(
82                "SELECT subscription_id,endpoint,event_operations,entity_types,ignored_origins,consumer_origin,secret_ref,enabled \
83                 FROM webhook_subscription ORDER BY created_at_us, subscription_id",
84            )
85            .map_err(sql_error)?;
86        statement
87            .query_map([], row)
88            .map_err(sql_error)?
89            .collect::<std::result::Result<Vec<_>, _>>()
90            .map_err(sql_error)?
91            .into_iter()
92            .map(decode)
93            .collect::<Result<Vec<_>>>()
94    }
95
96    pub fn matching(&self, event: &ChangeEvent) -> Result<Vec<WebhookSubscription>> {
97        let entity_type = event
98            .change
99            .after
100            .as_ref()
101            .or(event.change.before.as_ref())
102            .map(|x| x.entity_type.as_str())
103            .ok_or_else(|| MCSError::MemoryError("event missing entity snapshot".into()))?;
104        let mut statement = self.conn.prepare("SELECT subscription_id,endpoint,event_operations,entity_types,ignored_origins,consumer_origin,secret_ref,enabled FROM webhook_subscription WHERE enabled=1").map_err(sql_error)?;
105        statement
106            .query_map([], row)
107            .map_err(sql_error)?
108            .collect::<std::result::Result<Vec<_>, _>>()
109            .map_err(sql_error)?
110            .into_iter()
111            .map(decode)
112            .collect::<Result<Vec<_>>>()
113            .map(|subscriptions| {
114                subscriptions
115                    .into_iter()
116                    .filter(|subscription| {
117                        subscription.consumer_origin != event.provenance.origin
118                            && !subscription
119                                .ignored_origins
120                                .contains(&event.provenance.origin)
121                            && (subscription.event_operations.is_empty()
122                                || subscription
123                                    .event_operations
124                                    .contains(&event.change.operation))
125                            && (subscription.entity_types.is_empty()
126                                || subscription
127                                    .entity_types
128                                    .iter()
129                                    .any(|kind| kind == entity_type))
130                    })
131                    .collect()
132            })
133    }
134
135    pub fn enqueue_matching(&self, event: &ChangeEvent) -> Result<usize> {
136        let matches = self.matching(event)?;
137        let mut inserted = 0;
138        for subscription in matches {
139            inserted += self.conn.execute("INSERT INTO event_outbox(delivery_id,event_id,subscription_id) VALUES(?1,?2,?3) ON CONFLICT(event_id,subscription_id) DO NOTHING", params![Uuid::new_v4().to_string(), event.event_id.to_string(), subscription.subscription_id.to_string()]).map_err(sql_error)?;
140        }
141        Ok(inserted)
142    }
143}
144
145type SubscriptionRow = (String, String, String, String, String, String, String, bool);
146fn row(row: &rusqlite::Row<'_>) -> rusqlite::Result<SubscriptionRow> {
147    Ok((
148        row.get(0)?,
149        row.get(1)?,
150        row.get(2)?,
151        row.get(3)?,
152        row.get(4)?,
153        row.get(5)?,
154        row.get(6)?,
155        row.get(7)?,
156    ))
157}
158fn decode(row: SubscriptionRow) -> Result<WebhookSubscription> {
159    Ok(WebhookSubscription {
160        subscription_id: parse_uuid(&row.0)?,
161        endpoint: row.1,
162        event_operations: serde_json::from_str(&row.2)?,
163        entity_types: serde_json::from_str(&row.3)?,
164        ignored_origins: serde_json::from_str(&row.4)?,
165        consumer_origin: row.5,
166        secret_ref: row.6,
167        enabled: row.7,
168    })
169}