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)]
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    /// Field-level filters applied to event payloads.
135    pub field_filters:         Vec<FieldFilter>,
136    /// Security context hash for RLS grouping.
137    pub security_context_hash: u64,
138}
139
140/// Thread-safe subscription manager with two-level indexing.
141///
142/// Level 1: entity → set of connection IDs (for fan-out).
143/// Level 2: connection → map of entity → subscription details (for per-connection state).
144pub struct SubscriptionManager {
145    /// entity → set of connection IDs subscribed to it.
146    entity_subscribers:       DashMap<String, DashSet<ConnectionId>>,
147    /// `connection_id` → (entity → subscription details).
148    connection_subscriptions: DashMap<ConnectionId, HashMap<String, SubscriptionDetails>>,
149    /// Maximum subscriptions per entity (fan-out limit).
150    max_per_entity:           usize,
151}
152
153impl SubscriptionManager {
154    /// Create a new subscription manager with the given fan-out limit.
155    #[must_use]
156    pub fn new(max_per_entity: usize) -> Self {
157        Self {
158            entity_subscribers: DashMap::new(),
159            connection_subscriptions: DashMap::new(),
160            max_per_entity,
161        }
162    }
163
164    /// Subscribe a connection to an entity.
165    ///
166    /// Returns `Ok(true)` if this is a new subscription, `Ok(false)` if the
167    /// connection was already subscribed (idempotent).
168    ///
169    /// # Errors
170    ///
171    /// Returns an error if the fan-out limit for this entity is reached.
172    pub fn subscribe(
173        &self,
174        connection_id: &str,
175        entity: &str,
176        details: SubscriptionDetails,
177    ) -> Result<bool, String> {
178        // Check if already subscribed (idempotent)
179        if let Some(subs) = self.connection_subscriptions.get(connection_id) {
180            if subs.contains_key(entity) {
181                return Ok(false);
182            }
183        }
184
185        // Check fan-out limit
186        let current_count = self.entity_subscribers.get(entity).map_or(0, |set| set.len());
187        if current_count >= self.max_per_entity {
188            return Err(format!(
189                "subscription limit reached for entity {entity} ({} max)",
190                self.max_per_entity
191            ));
192        }
193
194        // Add to entity → connections index
195        self.entity_subscribers
196            .entry(entity.to_owned())
197            .or_default()
198            .insert(connection_id.to_owned());
199
200        // Add to connection → subscriptions index
201        self.connection_subscriptions
202            .entry(connection_id.to_owned())
203            .or_default()
204            .insert(entity.to_owned(), details);
205
206        Ok(true)
207    }
208
209    /// Unsubscribe a connection from an entity.
210    ///
211    /// Returns `true` if the subscription existed and was removed.
212    #[must_use]
213    pub fn unsubscribe(&self, connection_id: &str, entity: &str) -> bool {
214        // Remove from connection → subscriptions
215        let had_sub = self
216            .connection_subscriptions
217            .get_mut(connection_id)
218            .is_some_and(|mut subs| subs.remove(entity).is_some());
219
220        if had_sub {
221            // Remove from entity → connections
222            if let Some(set) = self.entity_subscribers.get(entity) {
223                set.remove(connection_id);
224            }
225        }
226
227        had_sub
228    }
229
230    /// Remove all subscriptions for a connection (called on disconnect).
231    pub fn unsubscribe_all(&self, connection_id: &str) {
232        if let Some((_, subs)) = self.connection_subscriptions.remove(connection_id) {
233            for entity in subs.keys() {
234                if let Some(set) = self.entity_subscribers.get(entity) {
235                    set.remove(connection_id);
236                }
237            }
238        }
239    }
240
241    /// Number of subscriptions for a given entity.
242    #[must_use]
243    pub fn count_for_entity(&self, entity: &str) -> usize {
244        self.entity_subscribers.get(entity).map_or(0, |set| set.len())
245    }
246
247    /// Number of entities a connection is subscribed to.
248    #[must_use]
249    pub fn count_for_connection(&self, connection_id: &str) -> usize {
250        self.connection_subscriptions.get(connection_id).map_or(0, |subs| subs.len())
251    }
252
253    /// Get all subscribers for an entity with their subscription details.
254    ///
255    /// Returns a vec of `(connection_id, details)` pairs, or `None` if no
256    /// connections are subscribed to this entity.
257    #[must_use]
258    pub fn get_subscribers(
259        &self,
260        entity: &str,
261    ) -> Option<Vec<(ConnectionId, SubscriptionDetails)>> {
262        let subscriber_set = self.entity_subscribers.get(entity)?;
263        if subscriber_set.is_empty() {
264            return None;
265        }
266
267        let mut result = Vec::with_capacity(subscriber_set.len());
268        for conn_id_ref in subscriber_set.iter() {
269            let conn_id = conn_id_ref.key().clone();
270            if let Some(subs) = self.connection_subscriptions.get(&conn_id) {
271                if let Some(details) = subs.get(entity) {
272                    result.push((conn_id, details.clone()));
273                }
274            }
275        }
276
277        if result.is_empty() {
278            None
279        } else {
280            Some(result)
281        }
282    }
283}