contextvm_sdk/transport/server/
correlation_store.rs1use std::collections::{HashMap, HashSet};
4use std::num::NonZeroUsize;
5use std::sync::Arc;
6use std::time::{Duration, Instant};
7
8use lru::LruCache;
9use tokio::sync::RwLock;
10
11use crate::core::constants::DEFAULT_LRU_SIZE;
12
13#[derive(Debug, Clone)]
15pub struct RouteEntry {
16 pub client_pubkey: String,
18 pub original_request_id: serde_json::Value,
20 pub progress_token: Option<String>,
22 pub wrap_kind: Option<u16>,
25 pub registered_at: Instant,
27}
28
29struct Inner {
31 routes: LruCache<String, RouteEntry>,
33 progress_token_to_event: HashMap<String, String>,
35 client_event_ids: HashMap<String, HashSet<String>>,
37}
38
39impl Inner {
40 fn new(max_routes: usize) -> Self {
41 let routes =
42 LruCache::new(NonZeroUsize::new(max_routes).unwrap_or(NonZeroUsize::new(1).unwrap()));
43 Self {
44 routes,
45 progress_token_to_event: HashMap::new(),
46 client_event_ids: HashMap::new(),
47 }
48 }
49
50 fn cleanup_indexes(&mut self, event_id: &str, route: &RouteEntry) {
52 if let Some(ref token) = route.progress_token {
53 self.progress_token_to_event.remove(token);
54 }
55 if let Some(set) = self.client_event_ids.get_mut(&route.client_pubkey) {
56 set.remove(event_id);
57 if set.is_empty() {
58 self.client_event_ids.remove(&route.client_pubkey);
59 }
60 }
61 }
62
63 fn remove_route(&mut self, event_id: &str) -> Option<RouteEntry> {
65 let route = self.routes.pop(event_id)?;
66 self.cleanup_indexes(event_id, &route);
67 Some(route)
68 }
69}
70
71#[derive(Clone)]
76pub struct ServerEventRouteStore {
77 inner: Arc<RwLock<Inner>>,
78}
79
80impl Default for ServerEventRouteStore {
81 fn default() -> Self {
82 Self::new()
83 }
84}
85
86impl ServerEventRouteStore {
87 pub fn new() -> Self {
89 Self {
90 inner: Arc::new(RwLock::new(Inner::new(DEFAULT_LRU_SIZE))),
91 }
92 }
93
94 pub fn with_max_routes(max_routes: usize) -> Self {
97 Self {
98 inner: Arc::new(RwLock::new(Inner::new(max_routes))),
99 }
100 }
101
102 pub async fn register(
104 &self,
105 event_id: String,
106 client_pubkey: String,
107 original_request_id: serde_json::Value,
108 progress_token: Option<String>,
109 ) {
110 let mut inner = self.inner.write().await;
111
112 inner
114 .client_event_ids
115 .entry(client_pubkey.clone())
116 .or_default()
117 .insert(event_id.clone());
118
119 if let Some(ref token) = progress_token {
121 inner
122 .progress_token_to_event
123 .insert(token.clone(), event_id.clone());
124 }
125
126 let evicted = inner.routes.push(
128 event_id.clone(),
129 RouteEntry {
130 client_pubkey,
131 original_request_id,
132 progress_token,
133 wrap_kind: None,
134 registered_at: Instant::now(),
135 },
136 );
137
138 if let Some((evicted_key, evicted_route)) = evicted {
139 if evicted_key != event_id {
140 inner.cleanup_indexes(&evicted_key, &evicted_route);
142 }
143 }
144 }
145
146 pub async fn get(&self, event_id: &str) -> Option<String> {
148 self.inner
149 .read()
150 .await
151 .routes
152 .peek(event_id)
153 .map(|r| r.client_pubkey.clone())
154 }
155
156 pub async fn get_route(&self, event_id: &str) -> Option<RouteEntry> {
158 self.inner.read().await.routes.peek(event_id).cloned()
159 }
160
161 pub async fn pop(&self, event_id: &str) -> Option<RouteEntry> {
163 self.inner.write().await.remove_route(event_id)
164 }
165
166 pub async fn remove_for_client(&self, client_pubkey: &str) -> usize {
168 let mut inner = self.inner.write().await;
169
170 let event_ids = match inner.client_event_ids.remove(client_pubkey) {
171 Some(ids) => ids,
172 None => return 0,
173 };
174
175 let count = event_ids.len();
176 for event_id in &event_ids {
177 if let Some(route) = inner.routes.pop(event_id.as_str()) {
178 if let Some(ref token) = route.progress_token {
179 inner.progress_token_to_event.remove(token);
180 }
181 }
182 }
183 count
184 }
185
186 pub async fn has_event_route(&self, event_id: &str) -> bool {
188 self.inner.read().await.routes.contains(event_id)
189 }
190
191 pub async fn has_active_routes_for_client(&self, client_pubkey: &str) -> bool {
193 self.inner
194 .read()
195 .await
196 .client_event_ids
197 .get(client_pubkey)
198 .is_some_and(|set| !set.is_empty())
199 }
200
201 pub async fn get_event_id_by_progress_token(&self, token: &str) -> Option<String> {
203 self.inner
204 .read()
205 .await
206 .progress_token_to_event
207 .get(token)
208 .cloned()
209 }
210
211 pub async fn has_progress_token(&self, token: &str) -> bool {
213 self.inner
214 .read()
215 .await
216 .progress_token_to_event
217 .contains_key(token)
218 }
219
220 pub async fn event_route_count(&self) -> usize {
222 self.inner.read().await.routes.len()
223 }
224
225 pub async fn progress_token_count(&self) -> usize {
227 self.inner.read().await.progress_token_to_event.len()
228 }
229
230 pub async fn sweep_stale_routes(&self, timeout: Duration) -> Vec<String> {
234 let now = Instant::now();
235 let mut inner = self.inner.write().await;
236 let mut expired_keys = Vec::new();
237
238 for (key, entry) in inner.routes.iter() {
239 if now.duration_since(entry.registered_at) >= timeout {
240 expired_keys.push(key.clone());
241 }
242 }
243
244 for key in &expired_keys {
245 inner.remove_route(key);
246 }
247 expired_keys
248 }
249
250 pub async fn clear(&self) {
252 let mut inner = self.inner.write().await;
253 inner.routes.clear();
254 inner.progress_token_to_event.clear();
255 inner.client_event_ids.clear();
256 }
257}
258
259#[cfg(test)]
260mod tests {
261 use super::*;
262 use serde_json::json;
263
264 #[tokio::test]
265 async fn pop_on_empty_returns_none() {
266 let store = ServerEventRouteStore::new();
267 assert!(store.pop("nonexistent").await.is_none());
268 }
269
270 #[tokio::test]
271 async fn get_returns_without_removing() {
272 let store = ServerEventRouteStore::new();
273 store
274 .register("e1".into(), "pk1".into(), json!("r1"), None)
275 .await;
276 assert_eq!(store.get("e1").await.as_deref(), Some("pk1"));
277 assert_eq!(store.get("e1").await.as_deref(), Some("pk1"));
278 }
279
280 #[tokio::test]
281 async fn pop_removes_entry() {
282 let store = ServerEventRouteStore::new();
283 store
284 .register("e1".into(), "pk1".into(), json!("r1"), None)
285 .await;
286 let route = store.pop("e1").await.unwrap();
287 assert_eq!(route.client_pubkey, "pk1");
288 assert!(store.pop("e1").await.is_none());
289 }
290
291 #[tokio::test]
292 async fn remove_for_client_only_removes_matching() {
293 let store = ServerEventRouteStore::new();
294 store
295 .register("e1".into(), "pk1".into(), json!("r1"), None)
296 .await;
297 store
298 .register("e2".into(), "pk2".into(), json!("r2"), None)
299 .await;
300 store
301 .register("e3".into(), "pk1".into(), json!("r3"), None)
302 .await;
303
304 let removed = store.remove_for_client("pk1").await;
305 assert_eq!(removed, 2);
306
307 assert!(store.get("e1").await.is_none());
308 assert!(store.get("e3").await.is_none());
309 assert_eq!(store.get("e2").await.as_deref(), Some("pk2"));
310 }
311
312 #[tokio::test]
313 async fn remove_for_client_noop_when_no_match() {
314 let store = ServerEventRouteStore::new();
315 store
316 .register("e1".into(), "pk1".into(), json!("r1"), None)
317 .await;
318 let removed = store.remove_for_client("pk_other").await;
319 assert_eq!(removed, 0);
320 assert_eq!(store.get("e1").await.as_deref(), Some("pk1"));
321 }
322
323 #[tokio::test]
324 async fn clear_empties_store() {
325 let store = ServerEventRouteStore::new();
326 store
327 .register("e1".into(), "pk1".into(), json!("r1"), None)
328 .await;
329 store
330 .register("e2".into(), "pk2".into(), json!("r2"), None)
331 .await;
332 store.clear().await;
333 assert!(store.get("e1").await.is_none());
334 assert!(store.get("e2").await.is_none());
335 }
336
337 #[tokio::test]
338 async fn default_store_is_bounded() {
339 let store = ServerEventRouteStore::new();
340 for i in 0..=DEFAULT_LRU_SIZE {
341 store
342 .register(format!("e{i}"), "pk1".into(), json!(i), None)
343 .await;
344 }
345
346 assert_eq!(store.event_route_count().await, DEFAULT_LRU_SIZE);
347 assert!(!store.has_event_route("e0").await);
348 assert!(store.has_event_route(&format!("e{DEFAULT_LRU_SIZE}")).await);
349 }
350
351 #[tokio::test]
352 async fn sweep_stale_routes_removes_only_expired() {
353 let store = ServerEventRouteStore::new();
354
355 store
357 .register("old".into(), "pk1".into(), json!(1), Some("tok1".into()))
358 .await;
359
360 tokio::time::sleep(Duration::from_millis(20)).await;
361
362 store
364 .register("fresh".into(), "pk2".into(), json!(2), None)
365 .await;
366
367 let swept = store.sweep_stale_routes(Duration::from_millis(10)).await;
369 assert_eq!(swept.len(), 1);
370 assert_eq!(swept[0], "old");
371 assert!(!store.has_event_route("old").await);
372 assert!(store.has_event_route("fresh").await);
373 assert!(!store.has_progress_token("tok1").await);
375 assert!(!store.has_active_routes_for_client("pk1").await);
376 }
377
378 #[tokio::test]
379 async fn sweep_stale_routes_returns_zero_when_nothing_expired() {
380 let store = ServerEventRouteStore::new();
381 store
382 .register("e1".into(), "pk1".into(), json!(1), None)
383 .await;
384
385 let swept = store.sweep_stale_routes(Duration::from_secs(60)).await;
386 assert!(swept.is_empty());
387 assert!(store.has_event_route("e1").await);
388 }
389}