Skip to main content

contextvm_sdk/transport/server/
correlation_store.rs

1//! Server-side event route store for mapping event IDs to client routes.
2
3use 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/// A route entry for an in-flight request.
14#[derive(Debug, Clone)]
15pub struct RouteEntry {
16    /// The client's public key that originated this request.
17    pub client_pubkey: String,
18    /// The original JSON-RPC request ID (before replacement with event ID).
19    pub original_request_id: serde_json::Value,
20    /// Optional progress token for this request.
21    pub progress_token: Option<String>,
22    /// The outer gift-wrap event kind that carried this request (e.g. 1059 or 21059).
23    /// Populated from the inbound event once wrap-kind mirroring is wired; `None` until then.
24    pub wrap_kind: Option<u16>,
25    /// When the route was registered.
26    pub registered_at: Instant,
27}
28
29/// Internal state behind the lock.
30struct Inner {
31    /// Primary index: event_id → route entry (LRU-ordered).
32    routes: LruCache<String, RouteEntry>,
33    /// Secondary index: progress_token → event_id.
34    progress_token_to_event: HashMap<String, String>,
35    /// Secondary index: client_pubkey → set of event_ids.
36    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    /// Clean up secondary indexes for a removed route.
51    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    /// Remove a single route and clean up all secondary indexes.
64    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/// Maps event IDs to full route entries for response routing on the server side.
72///
73/// An optional capacity limit enables LRU eviction; when the limit is reached
74/// the oldest entry is evicted and its secondary indexes are cleaned up.
75#[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    /// Create a new store with the default capacity
88    pub fn new() -> Self {
89        Self {
90            inner: Arc::new(RwLock::new(Inner::new(DEFAULT_LRU_SIZE))),
91        }
92    }
93
94    /// Create a store with an upper bound on event routes.
95    /// When the limit is reached the oldest entry is evicted.
96    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    /// Register a route for an incoming request.
103    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        // Update client index.
113        inner
114            .client_event_ids
115            .entry(client_pubkey.clone())
116            .or_default()
117            .insert(event_id.clone());
118
119        // Update progress token index.
120        if let Some(ref token) = progress_token {
121            inner
122                .progress_token_to_event
123                .insert(token.clone(), event_id.clone());
124        }
125
126        // Insert into LRU; handle possible eviction.
127        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                // A different entry was evicted due to capacity — clean up its indexes.
141                inner.cleanup_indexes(&evicted_key, &evicted_route);
142            }
143        }
144    }
145
146    /// Returns the client public key for the given event ID without removing it.
147    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    /// Returns the full route entry for the given event ID without removing it.
157    pub async fn get_route(&self, event_id: &str) -> Option<RouteEntry> {
158        self.inner.read().await.routes.peek(event_id).cloned()
159    }
160
161    /// Removes and returns the full route entry for the given event ID.
162    pub async fn pop(&self, event_id: &str) -> Option<RouteEntry> {
163        self.inner.write().await.remove_route(event_id)
164    }
165
166    /// Removes all routes for a given client public key. Returns the count removed.
167    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    /// Check whether a route exists for the given event ID.
187    pub async fn has_event_route(&self, event_id: &str) -> bool {
188        self.inner.read().await.routes.contains(event_id)
189    }
190
191    /// Check whether the given client has any active routes.
192    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    /// Look up the event ID associated with a progress token.
202    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    /// Check whether a progress token mapping exists.
212    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    /// Number of event routes currently tracked.
221    pub async fn event_route_count(&self) -> usize {
222        self.inner.read().await.routes.len()
223    }
224
225    /// Number of progress token mappings currently tracked.
226    pub async fn progress_token_count(&self) -> usize {
227        self.inner.read().await.progress_token_to_event.len()
228    }
229
230    /// Remove all route entries older than `timeout`.
231    /// (Routes for expired sessions are already cleaned by `cleanup_sessions`.)
232    /// Returns the event IDs of the removed entries.
233    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    /// Remove all route entries and secondary indexes
251    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        // Insert a route that will age past the threshold.
356        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        // Insert a fresh route.
363        store
364            .register("fresh".into(), "pk2".into(), json!(2), None)
365            .await;
366
367        // Sweep with 10ms timeout — "old" should be removed, "fresh" should remain.
368        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        // Secondary indexes should also be cleaned.
374        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}