Skip to main content

fraiseql_server/realtime/
subscriptions.rs

1//! Subscription manager for tracking entity change subscriptions.
2//!
3//! Uses a two-level index for O(1) fan-out (entity → connections) and
4//! O(1) per-connection cleanup (connection → subscriptions).
5
6use std::collections::HashMap;
7
8use dashmap::{DashMap, DashSet};
9use serde_json::Value;
10
11use super::connections::ConnectionId;
12
13/// Event kind for filtering subscription events.
14#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
15#[non_exhaustive]
16pub enum EventKind {
17    /// `INSERT` — new row created.
18    Insert,
19    /// `UPDATE` — existing row modified.
20    Update,
21    /// `DELETE` — row removed.
22    Delete,
23}
24
25impl EventKind {
26    /// Parse an event kind from a string.
27    ///
28    /// # Errors
29    ///
30    /// Returns an error message if the string is not a recognized event kind.
31    pub fn parse(s: &str) -> Result<Self, String> {
32        match s.to_uppercase().as_str() {
33            "INSERT" => Ok(Self::Insert),
34            "UPDATE" => Ok(Self::Update),
35            "DELETE" => Ok(Self::Delete),
36            other => Err(format!("unknown event kind: {other}")),
37        }
38    }
39}
40
41/// Comparison operator for field filters.
42#[derive(Debug, Clone, PartialEq, Eq)]
43#[non_exhaustive]
44pub enum FilterOperator {
45    /// Equal (`eq`).
46    Eq,
47    /// Not equal (`neq`).
48    Neq,
49    /// Greater than (`gt`).
50    Gt,
51    /// Less than (`lt`).
52    Lt,
53    /// Greater than or equal (`gte`).
54    Gte,
55    /// Less than or equal (`lte`).
56    Lte,
57    /// In a set of values (`in`).
58    In,
59}
60
61/// A single field-level filter on subscription events.
62#[derive(Debug, Clone, PartialEq, Eq)]
63pub struct FieldFilter {
64    /// Field name to compare.
65    pub field:    String,
66    /// Comparison operator.
67    pub operator: FilterOperator,
68    /// Value to compare against.
69    pub value:    Value,
70}
71
72impl FilterOperator {
73    /// Parse an operator from its string representation.
74    ///
75    /// # Errors
76    ///
77    /// Returns an error message if the string is not a recognized operator.
78    pub fn parse(s: &str) -> Result<Self, String> {
79        match s {
80            "eq" => Ok(Self::Eq),
81            "neq" => Ok(Self::Neq),
82            "gt" => Ok(Self::Gt),
83            "lt" => Ok(Self::Lt),
84            "gte" => Ok(Self::Gte),
85            "lte" => Ok(Self::Lte),
86            "in" => Ok(Self::In),
87            other => Err(format!("unknown filter operator: {other}")),
88        }
89    }
90}
91
92/// Parse a filter value, coercing to number when possible.
93fn parse_filter_value(s: &str) -> Value {
94    if let Ok(n) = s.parse::<i64>() {
95        Value::Number(n.into())
96    } else if let Ok(f) = s.parse::<f64>() {
97        serde_json::Number::from_f64(f).map_or_else(|| Value::String(s.to_owned()), Value::Number)
98    } else {
99        Value::String(s.to_owned())
100    }
101}
102
103/// Parse a filter string in `field=op.value` format.
104///
105/// Multiple filters can be comma-separated: `"author_id=eq.123,status=neq.draft"`.
106///
107/// # Errors
108///
109/// Returns an error message if the filter string is malformed.
110pub fn parse_filter(filter_str: &str) -> Result<Vec<FieldFilter>, String> {
111    filter_str
112        .split(',')
113        .map(str::trim)
114        .filter(|p| !p.is_empty())
115        .map(|part| {
116            let (field, rest) =
117                part.split_once('=').ok_or_else(|| format!("invalid filter syntax: {part}"))?;
118            let (op_str, value_str) =
119                rest.split_once('.').ok_or_else(|| format!("invalid filter operator: {rest}"))?;
120            Ok(FieldFilter {
121                field:    field.to_owned(),
122                operator: FilterOperator::parse(op_str)?,
123                value:    parse_filter_value(value_str),
124            })
125        })
126        .collect()
127}
128
129/// Details for a single subscription held by a connection.
130#[derive(Debug, Clone)]
131pub struct SubscriptionDetails {
132    /// Optional event type filter (None = all events).
133    pub event_filter:          Option<EventKind>,
134    /// Client-supplied field-level filters applied to event payloads (cooperative — a
135    /// client can only narrow its own stream).
136    pub field_filters:         Vec<FieldFilter>,
137    /// Security context hash for RLS grouping.
138    pub security_context_hash: u64,
139    /// Server-owned row-visibility enforcement for a policy-declaring entity (#596),
140    /// resolved at subscribe time. The delivery pipeline **denies** a policy entity's
141    /// events to a subscription left at `OwnerEnforcement::None` — so the seam is
142    /// fail-closed by construction even if a future assembler skips the subscribe-time
143    /// wiring.
144    pub owner_enforcement:     super::subscription_policy::OwnerEnforcement,
145}
146
147/// Thread-safe subscription manager with two-level indexing.
148///
149/// Level 1: entity → set of connection IDs (for fan-out).
150/// Level 2: connection → map of entity → subscription details (for per-connection state).
151pub struct SubscriptionManager {
152    /// entity → set of connection IDs subscribed to it.
153    entity_subscribers:       DashMap<String, DashSet<ConnectionId>>,
154    /// `connection_id` → (entity → subscription details).
155    connection_subscriptions: DashMap<ConnectionId, HashMap<String, SubscriptionDetails>>,
156    /// Maximum subscriptions per entity (fan-out limit).
157    max_per_entity:           usize,
158}
159
160impl SubscriptionManager {
161    /// Create a new subscription manager with the given fan-out limit.
162    #[must_use]
163    pub fn new(max_per_entity: usize) -> Self {
164        Self {
165            entity_subscribers: DashMap::new(),
166            connection_subscriptions: DashMap::new(),
167            max_per_entity,
168        }
169    }
170
171    /// Subscribe a connection to an entity.
172    ///
173    /// Returns `Ok(true)` if this is a new subscription, `Ok(false)` if the
174    /// connection was already subscribed (idempotent).
175    ///
176    /// # Errors
177    ///
178    /// Returns an error if the fan-out limit for this entity is reached.
179    pub fn subscribe(
180        &self,
181        connection_id: &str,
182        entity: &str,
183        details: SubscriptionDetails,
184    ) -> Result<bool, String> {
185        // Check if already subscribed (idempotent)
186        if let Some(subs) = self.connection_subscriptions.get(connection_id) {
187            if subs.contains_key(entity) {
188                return Ok(false);
189            }
190        }
191
192        // Check fan-out limit
193        let current_count = self.entity_subscribers.get(entity).map_or(0, |set| set.len());
194        if current_count >= self.max_per_entity {
195            return Err(format!(
196                "subscription limit reached for entity {entity} ({} max)",
197                self.max_per_entity
198            ));
199        }
200
201        // Add to entity → connections index
202        self.entity_subscribers
203            .entry(entity.to_owned())
204            .or_default()
205            .insert(connection_id.to_owned());
206
207        // Add to connection → subscriptions index
208        self.connection_subscriptions
209            .entry(connection_id.to_owned())
210            .or_default()
211            .insert(entity.to_owned(), details);
212
213        Ok(true)
214    }
215
216    /// Unsubscribe a connection from an entity.
217    ///
218    /// Returns `true` if the subscription existed and was removed.
219    #[must_use]
220    pub fn unsubscribe(&self, connection_id: &str, entity: &str) -> bool {
221        // Remove from connection → subscriptions
222        let had_sub = self
223            .connection_subscriptions
224            .get_mut(connection_id)
225            .is_some_and(|mut subs| subs.remove(entity).is_some());
226
227        if had_sub {
228            // Remove from entity → connections
229            if let Some(set) = self.entity_subscribers.get(entity) {
230                set.remove(connection_id);
231            }
232        }
233
234        had_sub
235    }
236
237    /// Remove all subscriptions for a connection (called on disconnect).
238    pub fn unsubscribe_all(&self, connection_id: &str) {
239        if let Some((_, subs)) = self.connection_subscriptions.remove(connection_id) {
240            for entity in subs.keys() {
241                if let Some(set) = self.entity_subscribers.get(entity) {
242                    set.remove(connection_id);
243                }
244            }
245        }
246    }
247
248    /// Number of subscriptions for a given entity.
249    #[must_use]
250    pub fn count_for_entity(&self, entity: &str) -> usize {
251        self.entity_subscribers.get(entity).map_or(0, |set| set.len())
252    }
253
254    /// Number of entities a connection is subscribed to.
255    #[must_use]
256    pub fn count_for_connection(&self, connection_id: &str) -> usize {
257        self.connection_subscriptions.get(connection_id).map_or(0, |subs| subs.len())
258    }
259
260    /// Get all subscribers for an entity with their subscription details.
261    ///
262    /// Returns a vec of `(connection_id, details)` pairs, or `None` if no
263    /// connections are subscribed to this entity.
264    #[must_use]
265    pub fn get_subscribers(
266        &self,
267        entity: &str,
268    ) -> Option<Vec<(ConnectionId, SubscriptionDetails)>> {
269        let subscriber_set = self.entity_subscribers.get(entity)?;
270        if subscriber_set.is_empty() {
271            return None;
272        }
273
274        let mut result = Vec::with_capacity(subscriber_set.len());
275        for conn_id_ref in subscriber_set.iter() {
276            let conn_id = conn_id_ref.key().clone();
277            if let Some(subs) = self.connection_subscriptions.get(&conn_id) {
278                if let Some(details) = subs.get(entity) {
279                    result.push((conn_id, details.clone()));
280                }
281            }
282        }
283
284        if result.is_empty() {
285            None
286        } else {
287            Some(result)
288        }
289    }
290}