fraiseql_server/realtime/
delivery.rs1use std::{
8 collections::{HashMap, HashSet},
9 sync::Arc,
10};
11
12use futures::future::BoxFuture;
13use serde::{Deserialize, Serialize};
14use serde_json::Value;
15use tokio::sync::mpsc;
16use tracing::{debug, warn};
17
18use super::{
19 connections::{ConnectionId, ConnectionManager},
20 subscription_policy::OwnerEnforcement,
21 subscriptions::{EventKind, FieldFilter, FilterOperator, SubscriptionManager},
22};
23
24#[derive(Debug, Clone, Serialize, Deserialize)]
26pub struct EntityEvent {
27 pub entity: String,
29 pub event_kind: EventKindSerde,
31 pub new: Option<Value>,
33 pub old: Option<Value>,
35 pub timestamp: String,
37}
38
39#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
41#[serde(rename_all = "UPPERCASE")]
42pub enum EventKindSerde {
43 Insert,
45 Update,
47 Delete,
49}
50
51impl EventKindSerde {
52 #[must_use]
54 pub const fn to_event_kind(self) -> EventKind {
55 match self {
56 Self::Insert => EventKind::Insert,
57 Self::Update => EventKind::Update,
58 Self::Delete => EventKind::Delete,
59 }
60 }
61}
62
63impl From<EventKind> for EventKindSerde {
64 fn from(kind: EventKind) -> Self {
65 match kind {
66 EventKind::Insert => Self::Insert,
67 EventKind::Update => Self::Update,
68 EventKind::Delete => Self::Delete,
69 }
70 }
71}
72
73pub trait RlsEvaluator: Send + Sync + 'static {
89 fn can_access<'a>(
94 &'a self,
95 context_hash: u64,
96 entity: &'a str,
97 row: &'a Value,
98 ) -> BoxFuture<'a, bool>;
99}
100
101#[derive(Debug, Clone, Serialize)]
103pub struct ChangeMessage {
104 #[serde(rename = "type")]
106 pub msg_type: &'static str,
107 pub entity: String,
109 pub event: EventKindSerde,
111 pub new: Option<Value>,
113 pub old: Option<Value>,
115 pub timestamp: String,
117}
118
119impl ChangeMessage {
120 #[must_use]
122 pub fn from_event(event: &EntityEvent) -> Self {
123 Self {
124 msg_type: "change",
125 entity: event.entity.clone(),
126 event: event.event_kind,
127 new: event.new.clone(),
128 old: event.old.clone(),
129 timestamp: event.timestamp.clone(),
130 }
131 }
132}
133
134pub struct EventDeliveryPipeline {
136 subscriptions: Arc<SubscriptionManager>,
138 connections: Arc<ConnectionManager>,
140 rls_evaluator: Arc<dyn RlsEvaluator>,
142 policy_entities: HashSet<String>,
148 event_rx: mpsc::Receiver<EntityEvent>,
150}
151
152impl EventDeliveryPipeline {
153 pub fn new(
155 subscriptions: Arc<SubscriptionManager>,
156 connections: Arc<ConnectionManager>,
157 rls_evaluator: Arc<dyn RlsEvaluator>,
158 event_rx: mpsc::Receiver<EntityEvent>,
159 ) -> Self {
160 Self::with_policy_entities(
161 subscriptions,
162 connections,
163 rls_evaluator,
164 HashSet::new(),
165 event_rx,
166 )
167 }
168
169 pub fn with_policy_entities(
172 subscriptions: Arc<SubscriptionManager>,
173 connections: Arc<ConnectionManager>,
174 rls_evaluator: Arc<dyn RlsEvaluator>,
175 policy_entities: HashSet<String>,
176 event_rx: mpsc::Receiver<EntityEvent>,
177 ) -> Self {
178 Self {
179 subscriptions,
180 connections,
181 rls_evaluator,
182 policy_entities,
183 event_rx,
184 }
185 }
186
187 pub async fn run(mut self) {
189 while let Some(event) = self.event_rx.recv().await {
190 self.deliver_event(&event).await;
191 }
192 debug!("Event delivery pipeline shutting down");
193 }
194
195 async fn deliver_event(&self, event: &EntityEvent) {
197 let event_kind = event.event_kind.to_event_kind();
198
199 let Some(subscriber_details) = self.subscriptions.get_subscribers(&event.entity) else {
201 return;
202 };
203
204 let policy_entity = self.policy_entities.contains(&event.entity);
207
208 let mut groups: HashMap<u64, Vec<(ConnectionId, Vec<FieldFilter>, OwnerEnforcement)>> =
210 HashMap::new();
211 for (conn_id, details) in &subscriber_details {
212 if let Some(filter_kind) = details.event_filter {
214 if filter_kind != event_kind {
215 continue;
216 }
217 }
218 groups.entry(details.security_context_hash).or_default().push((
219 conn_id.clone(),
220 details.field_filters.clone(),
221 details.owner_enforcement.clone(),
222 ));
223 }
224
225 let row = event.new.as_ref().or(event.old.as_ref());
227
228 let Ok(json) = serde_json::to_string(&ChangeMessage::from_event(event)) else {
230 return;
231 };
232
233 for (context_hash, connections) in &groups {
235 if let Some(row) = row {
237 if !self.rls_evaluator.can_access(*context_hash, &event.entity, row).await {
238 debug!(
239 entity = %event.entity,
240 context_hash = context_hash,
241 "RLS denied event delivery"
242 );
243 continue;
244 }
245 }
246
247 for (conn_id, field_filters, owner_enforcement) in connections {
249 if policy_entity && !owner_enforcement_admits(owner_enforcement, row) {
251 debug!(
252 entity = %event.entity,
253 connection_id = %conn_id,
254 "row-visibility policy denied event delivery (fail-closed)"
255 );
256 continue;
257 }
258
259 if !evaluate_field_filters(field_filters, row) {
261 continue;
262 }
263
264 if !self.connections.send_event(conn_id, json.clone()) {
265 warn!(
266 connection_id = %conn_id,
267 "Failed to send event to connection (channel full or closed)"
268 );
269 }
270 }
271 }
272 }
273}
274
275#[must_use]
284pub fn owner_enforcement_admits(enforcement: &OwnerEnforcement, row: Option<&Value>) -> bool {
285 match enforcement {
286 OwnerEnforcement::None => false,
287 OwnerEnforcement::Bypass => true,
288 OwnerEnforcement::Scoped(owner_filter) => {
289 row.is_some_and(|r| evaluate_field_filters(std::slice::from_ref(owner_filter), Some(r)))
290 },
291 }
292}
293
294#[must_use]
298pub fn evaluate_field_filters(filters: &[FieldFilter], row: Option<&Value>) -> bool {
299 if filters.is_empty() {
300 return true;
301 }
302 let Some(row) = row else {
303 return true;
305 };
306 for filter in filters {
307 let field_value = row.get(&filter.field);
308 if !evaluate_single_filter(field_value, &filter.operator, &filter.value) {
309 return false;
310 }
311 }
312 true
313}
314
315fn evaluate_single_filter(
317 field_value: Option<&Value>,
318 operator: &FilterOperator,
319 filter_value: &Value,
320) -> bool {
321 let Some(field_value) = field_value else {
322 return matches!(operator, FilterOperator::Neq);
324 };
325
326 match operator {
327 FilterOperator::Eq => field_value == filter_value,
328 FilterOperator::Neq => field_value != filter_value,
329 FilterOperator::Gt => compare_values(field_value, filter_value).is_some_and(|o| o.is_gt()),
330 FilterOperator::Lt => compare_values(field_value, filter_value).is_some_and(|o| o.is_lt()),
331 FilterOperator::Gte => compare_values(field_value, filter_value).is_some_and(|o| o.is_ge()),
332 FilterOperator::Lte => compare_values(field_value, filter_value).is_some_and(|o| o.is_le()),
333 FilterOperator::In => {
334 if let Value::Array(arr) = filter_value {
335 arr.contains(field_value)
336 } else {
337 field_value == filter_value
338 }
339 },
340 }
341}
342
343fn compare_values(a: &Value, b: &Value) -> Option<std::cmp::Ordering> {
345 let a_num = value_as_f64(a);
347 let b_num = value_as_f64(b);
348 if let (Some(a_f), Some(b_f)) = (a_num, b_num) {
349 return a_f.partial_cmp(&b_f);
350 }
351
352 let a_str = a.as_str().or_else(|| if a.is_number() { None } else { Some("") });
354 let b_str = b.as_str().or_else(|| if b.is_number() { None } else { Some("") });
355 match (a_str, b_str) {
356 (Some(a_s), Some(b_s)) => Some(a_s.cmp(b_s)),
357 _ => None,
358 }
359}
360
361fn value_as_f64(v: &Value) -> Option<f64> {
363 match v {
364 Value::Number(n) => n.as_f64(),
365 Value::String(s) => s.parse::<f64>().ok(),
366 _ => None,
367 }
368}