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)]
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>,
136 pub security_context_hash: u64,
138}
139
140pub struct SubscriptionManager {
145 entity_subscribers: DashMap<String, DashSet<ConnectionId>>,
147 connection_subscriptions: DashMap<ConnectionId, HashMap<String, SubscriptionDetails>>,
149 max_per_entity: usize,
151}
152
153impl SubscriptionManager {
154 #[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 pub fn subscribe(
173 &self,
174 connection_id: &str,
175 entity: &str,
176 details: SubscriptionDetails,
177 ) -> Result<bool, String> {
178 if let Some(subs) = self.connection_subscriptions.get(connection_id) {
180 if subs.contains_key(entity) {
181 return Ok(false);
182 }
183 }
184
185 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 self.entity_subscribers
196 .entry(entity.to_owned())
197 .or_default()
198 .insert(connection_id.to_owned());
199
200 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 #[must_use]
213 pub fn unsubscribe(&self, connection_id: &str, entity: &str) -> bool {
214 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 if let Some(set) = self.entity_subscribers.get(entity) {
223 set.remove(connection_id);
224 }
225 }
226
227 had_sub
228 }
229
230 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 #[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 #[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 #[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}