fraiseql_server/realtime/
delivery.rs1use 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#[derive(Debug, Clone, Serialize, Deserialize)]
22pub struct EntityEvent {
23 pub entity: String,
25 pub event_kind: EventKindSerde,
27 pub new: Option<Value>,
29 pub old: Option<Value>,
31 pub timestamp: String,
33}
34
35#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
37#[serde(rename_all = "UPPERCASE")]
38pub enum EventKindSerde {
39 Insert,
41 Update,
43 Delete,
45}
46
47impl EventKindSerde {
48 #[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
69pub trait RlsEvaluator: Send + Sync + 'static {
85 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#[derive(Debug, Clone, Serialize)]
99pub struct ChangeMessage {
100 #[serde(rename = "type")]
102 pub msg_type: &'static str,
103 pub entity: String,
105 pub event: EventKindSerde,
107 pub new: Option<Value>,
109 pub old: Option<Value>,
111 pub timestamp: String,
113}
114
115impl ChangeMessage {
116 #[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
130pub struct EventDeliveryPipeline {
132 subscriptions: Arc<SubscriptionManager>,
134 connections: Arc<ConnectionManager>,
136 rls_evaluator: Arc<dyn RlsEvaluator>,
138 event_rx: mpsc::Receiver<EntityEvent>,
140}
141
142impl EventDeliveryPipeline {
143 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 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 async fn deliver_event(&self, event: &EntityEvent) {
168 let event_kind = event.event_kind.to_event_kind();
169
170 let Some(subscriber_details) = self.subscriptions.get_subscribers(&event.entity) else {
172 return;
173 };
174
175 let mut groups: HashMap<u64, Vec<(ConnectionId, Vec<FieldFilter>)>> = HashMap::new();
177 for (conn_id, details) in &subscriber_details {
178 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 let row = event.new.as_ref().or(event.old.as_ref());
192
193 let Ok(json) = serde_json::to_string(&ChangeMessage::from_event(event)) else {
195 return;
196 };
197
198 for (context_hash, connections) in &groups {
200 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 for (conn_id, field_filters) in connections {
214 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#[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 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
251fn 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 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
279fn compare_values(a: &Value, b: &Value) -> Option<std::cmp::Ordering> {
281 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 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
297fn 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}