fraiseql_server/realtime/
subscriptions.rs1use std::collections::HashMap;
7
8use dashmap::{DashMap, DashSet};
9use serde_json::Value;
10
11use super::connections::ConnectionId;
12
13#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
15#[non_exhaustive]
16pub enum EventKind {
17 Insert,
19 Update,
21 Delete,
23}
24
25impl EventKind {
26 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#[derive(Debug, Clone, PartialEq, Eq)]
43#[non_exhaustive]
44pub enum FilterOperator {
45 Eq,
47 Neq,
49 Gt,
51 Lt,
53 Gte,
55 Lte,
57 In,
59}
60
61#[derive(Debug, Clone, PartialEq, Eq)]
63pub struct FieldFilter {
64 pub field: String,
66 pub operator: FilterOperator,
68 pub value: Value,
70}
71
72impl FilterOperator {
73 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
92fn 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
103pub 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#[derive(Debug, Clone)]
131pub struct SubscriptionDetails {
132 pub event_filter: Option<EventKind>,
134 pub field_filters: Vec<FieldFilter>,
137 pub security_context_hash: u64,
139 pub owner_enforcement: super::subscription_policy::OwnerEnforcement,
145}
146
147pub struct SubscriptionManager {
152 entity_subscribers: DashMap<String, DashSet<ConnectionId>>,
154 connection_subscriptions: DashMap<ConnectionId, HashMap<String, SubscriptionDetails>>,
156 max_per_entity: usize,
158}
159
160impl SubscriptionManager {
161 #[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 pub fn subscribe(
180 &self,
181 connection_id: &str,
182 entity: &str,
183 details: SubscriptionDetails,
184 ) -> Result<bool, String> {
185 if let Some(subs) = self.connection_subscriptions.get(connection_id) {
187 if subs.contains_key(entity) {
188 return Ok(false);
189 }
190 }
191
192 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 self.entity_subscribers
203 .entry(entity.to_owned())
204 .or_default()
205 .insert(connection_id.to_owned());
206
207 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 #[must_use]
220 pub fn unsubscribe(&self, connection_id: &str, entity: &str) -> bool {
221 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 if let Some(set) = self.entity_subscribers.get(entity) {
230 set.remove(connection_id);
231 }
232 }
233
234 had_sub
235 }
236
237 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 #[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 #[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 #[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}