Skip to main content

fraiseql_server/realtime/
delivery.rs

1//! Event delivery pipeline for the realtime broadcast system.
2//!
3//! Receives entity change events, groups subscriptions by security context hash,
4//! evaluates RLS once per group, applies field filters, and delivers events to
5//! authorized connections.
6
7use std::{collections::HashMap, sync::Arc};
8
9use futures::future::BoxFuture;
10use serde::{Deserialize, Serialize};
11use serde_json::Value;
12use tokio::sync::mpsc;
13use tracing::{debug, warn};
14
15use super::{
16    connections::{ConnectionId, ConnectionManager},
17    subscriptions::{EventKind, FieldFilter, FilterOperator, SubscriptionManager},
18};
19
20/// An entity change event to be broadcast to subscribers.
21#[derive(Debug, Clone, Serialize, Deserialize)]
22pub struct EntityEvent {
23    /// Entity name (e.g., `"Post"`).
24    pub entity:     String,
25    /// Type of change.
26    pub event_kind: EventKindSerde,
27    /// New row data (present for INSERT and UPDATE).
28    pub new:        Option<Value>,
29    /// Old row data (present for UPDATE and DELETE).
30    pub old:        Option<Value>,
31    /// Event timestamp (ISO 8601).
32    pub timestamp:  String,
33}
34
35/// Serializable event kind for JSON wire format.
36#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
37#[serde(rename_all = "UPPERCASE")]
38pub enum EventKindSerde {
39    /// Row inserted.
40    Insert,
41    /// Row updated.
42    Update,
43    /// Row deleted.
44    Delete,
45}
46
47impl EventKindSerde {
48    /// Convert to the internal `EventKind` type.
49    #[must_use]
50    pub const fn to_event_kind(self) -> EventKind {
51        match self {
52            Self::Insert => EventKind::Insert,
53            Self::Update => EventKind::Update,
54            Self::Delete => EventKind::Delete,
55        }
56    }
57}
58
59impl From<EventKind> for EventKindSerde {
60    fn from(kind: EventKind) -> Self {
61        match kind {
62            EventKind::Insert => Self::Insert,
63            EventKind::Update => Self::Update,
64            EventKind::Delete => Self::Delete,
65        }
66    }
67}
68
69/// Trait for evaluating row-level security on event delivery.
70///
71/// Implementations check whether a given security context hash is authorized
72/// to see a specific row of a given entity.
73///
74/// The trait is object-safe (returns `BoxFuture`) so it can be stored as
75/// `Arc<dyn RlsEvaluator>` without infecting the delivery pipeline with a
76/// type parameter.
77///
78/// # Production implementation
79///
80/// The concrete `SqlRlsEvaluator` runs
81/// `SELECT EXISTS(SELECT 1 FROM entity WHERE pk = $1 AND <rls_where_clauses>)`
82/// against the database — one query per distinct security-context group per
83/// event.
84pub trait RlsEvaluator: Send + Sync + 'static {
85    /// Check if the given security context can access the row.
86    ///
87    /// Returns `true` if the row should be delivered, `false` if it should
88    /// be silently dropped.
89    fn can_access<'a>(
90        &'a self,
91        context_hash: u64,
92        entity: &'a str,
93        row: &'a Value,
94    ) -> BoxFuture<'a, bool>;
95}
96
97/// A change event formatted for delivery to a client.
98#[derive(Debug, Clone, Serialize)]
99pub struct ChangeMessage {
100    /// Always `"change"`.
101    #[serde(rename = "type")]
102    pub msg_type:  &'static str,
103    /// Entity name.
104    pub entity:    String,
105    /// Event type (`INSERT`, `UPDATE`, `DELETE`).
106    pub event:     EventKindSerde,
107    /// New row data.
108    pub new:       Option<Value>,
109    /// Old row data.
110    pub old:       Option<Value>,
111    /// Timestamp (ISO 8601).
112    pub timestamp: String,
113}
114
115impl ChangeMessage {
116    /// Create from an `EntityEvent`.
117    #[must_use]
118    pub fn from_event(event: &EntityEvent) -> Self {
119        Self {
120            msg_type:  "change",
121            entity:    event.entity.clone(),
122            event:     event.event_kind,
123            new:       event.new.clone(),
124            old:       event.old.clone(),
125            timestamp: event.timestamp.clone(),
126        }
127    }
128}
129
130/// Event delivery pipeline that processes entity events and delivers to subscribers.
131pub struct EventDeliveryPipeline {
132    /// Subscription manager for looking up who receives what.
133    subscriptions: Arc<SubscriptionManager>,
134    /// Connection manager for sending events to connections.
135    connections:   Arc<ConnectionManager>,
136    /// RLS evaluator for access control.
137    rls_evaluator: Arc<dyn RlsEvaluator>,
138    /// Receiver for incoming entity events.
139    event_rx:      mpsc::Receiver<EntityEvent>,
140}
141
142impl EventDeliveryPipeline {
143    /// Create a new event delivery pipeline.
144    pub fn new(
145        subscriptions: Arc<SubscriptionManager>,
146        connections: Arc<ConnectionManager>,
147        rls_evaluator: Arc<dyn RlsEvaluator>,
148        event_rx: mpsc::Receiver<EntityEvent>,
149    ) -> Self {
150        Self {
151            subscriptions,
152            connections,
153            rls_evaluator,
154            event_rx,
155        }
156    }
157
158    /// Run the delivery loop. Processes events until the channel is closed.
159    pub async fn run(mut self) {
160        while let Some(event) = self.event_rx.recv().await {
161            self.deliver_event(&event).await;
162        }
163        debug!("Event delivery pipeline shutting down");
164    }
165
166    /// Deliver a single event to all authorized subscribers.
167    async fn deliver_event(&self, event: &EntityEvent) {
168        let event_kind = event.event_kind.to_event_kind();
169
170        // Get all subscribers for this entity
171        let Some(subscriber_details) = self.subscriptions.get_subscribers(&event.entity) else {
172            return;
173        };
174
175        // Group subscribers by security context hash for RLS coalescing
176        let mut groups: HashMap<u64, Vec<(ConnectionId, Vec<FieldFilter>)>> = HashMap::new();
177        for (conn_id, details) in &subscriber_details {
178            // Apply event type filter
179            if let Some(filter_kind) = details.event_filter {
180                if filter_kind != event_kind {
181                    continue;
182                }
183            }
184            groups
185                .entry(details.security_context_hash)
186                .or_default()
187                .push((conn_id.clone(), details.field_filters.clone()));
188        }
189
190        // Determine the row to check for RLS (prefer `new`, fall back to `old`)
191        let row = event.new.as_ref().or(event.old.as_ref());
192
193        // Serialize the change message once (shared across all connections)
194        let Ok(json) = serde_json::to_string(&ChangeMessage::from_event(event)) else {
195            return;
196        };
197
198        // Evaluate RLS once per group, then deliver
199        for (context_hash, connections) in &groups {
200            // RLS check: can this security context see this row?
201            if let Some(row) = row {
202                if !self.rls_evaluator.can_access(*context_hash, &event.entity, row).await {
203                    debug!(
204                        entity = %event.entity,
205                        context_hash = context_hash,
206                        "RLS denied event delivery"
207                    );
208                    continue;
209                }
210            }
211
212            // Deliver to each connection in this group
213            for (conn_id, field_filters) in connections {
214                // Apply field filters
215                if !evaluate_field_filters(field_filters, row) {
216                    continue;
217                }
218
219                if !self.connections.send_event(conn_id, json.clone()) {
220                    warn!(
221                        connection_id = %conn_id,
222                        "Failed to send event to connection (channel full or closed)"
223                    );
224                }
225            }
226        }
227    }
228}
229
230/// Evaluate field filters against a row.
231///
232/// Returns `true` if the row passes all filters (or if there are no filters).
233#[must_use]
234pub fn evaluate_field_filters(filters: &[FieldFilter], row: Option<&Value>) -> bool {
235    if filters.is_empty() {
236        return true;
237    }
238    let Some(row) = row else {
239        // No row data to filter against — pass through
240        return true;
241    };
242    for filter in filters {
243        let field_value = row.get(&filter.field);
244        if !evaluate_single_filter(field_value, &filter.operator, &filter.value) {
245            return false;
246        }
247    }
248    true
249}
250
251/// Evaluate a single filter comparison.
252fn evaluate_single_filter(
253    field_value: Option<&Value>,
254    operator: &FilterOperator,
255    filter_value: &Value,
256) -> bool {
257    let Some(field_value) = field_value else {
258        // Field not present in row — filter fails (except for Neq)
259        return matches!(operator, FilterOperator::Neq);
260    };
261
262    match operator {
263        FilterOperator::Eq => field_value == filter_value,
264        FilterOperator::Neq => field_value != filter_value,
265        FilterOperator::Gt => compare_values(field_value, filter_value).is_some_and(|o| o.is_gt()),
266        FilterOperator::Lt => compare_values(field_value, filter_value).is_some_and(|o| o.is_lt()),
267        FilterOperator::Gte => compare_values(field_value, filter_value).is_some_and(|o| o.is_ge()),
268        FilterOperator::Lte => compare_values(field_value, filter_value).is_some_and(|o| o.is_le()),
269        FilterOperator::In => {
270            if let Value::Array(arr) = filter_value {
271                arr.contains(field_value)
272            } else {
273                field_value == filter_value
274            }
275        },
276    }
277}
278
279/// Compare two JSON values numerically if possible, otherwise as strings.
280fn compare_values(a: &Value, b: &Value) -> Option<std::cmp::Ordering> {
281    // Try numeric comparison
282    let a_num = value_as_f64(a);
283    let b_num = value_as_f64(b);
284    if let (Some(a_f), Some(b_f)) = (a_num, b_num) {
285        return a_f.partial_cmp(&b_f);
286    }
287
288    // Fall back to string comparison
289    let a_str = a.as_str().or_else(|| if a.is_number() { None } else { Some("") });
290    let b_str = b.as_str().or_else(|| if b.is_number() { None } else { Some("") });
291    match (a_str, b_str) {
292        (Some(a_s), Some(b_s)) => Some(a_s.cmp(b_s)),
293        _ => None,
294    }
295}
296
297/// Try to extract a float from a JSON value.
298fn value_as_f64(v: &Value) -> Option<f64> {
299    match v {
300        Value::Number(n) => n.as_f64(),
301        Value::String(s) => s.parse::<f64>().ok(),
302        _ => None,
303    }
304}