Skip to content

Commit 9584ec3

Browse files
feat: add downstream subscriptions listen admission
Signed-off-by: Pratik Gandhi <gandhipratik203@gmail.com>
1 parent 3f13387 commit 9584ec3

12 files changed

Lines changed: 669 additions & 12 deletions

File tree

Lines changed: 106 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,106 @@
1+
use std::{
2+
collections::HashMap,
3+
sync::{Arc, Mutex},
4+
};
5+
6+
use rmcp::{
7+
model::{RequestId, SubscriptionFilter},
8+
service::SubscriptionSink,
9+
};
10+
11+
use crate::layers::request_context::GatewayRequestContext;
12+
13+
#[derive(Clone, Default)]
14+
pub(crate) struct DownstreamSubscriptionRegistry {
15+
inner: Arc<Mutex<HashMap<DownstreamSubscriptionKey, SubscriptionSink>>>,
16+
}
17+
18+
impl DownstreamSubscriptionRegistry {
19+
pub(crate) fn register(
20+
&self,
21+
context: &GatewayRequestContext,
22+
filter: &SubscriptionFilter,
23+
sink: &SubscriptionSink,
24+
) -> DownstreamSubscriptionGuard {
25+
let keys = subscription_keys(context, filter, sink.id());
26+
let mut subscriptions = self.inner.lock().expect("downstream subscription registry lock poisoned");
27+
for key in &keys {
28+
subscriptions.insert(key.clone(), sink.clone());
29+
}
30+
DownstreamSubscriptionGuard { registry: self.clone(), keys }
31+
}
32+
33+
fn remove_all(&self, keys: &[DownstreamSubscriptionKey]) {
34+
let mut subscriptions = self.inner.lock().expect("downstream subscription registry lock poisoned");
35+
for key in keys {
36+
subscriptions.remove(key);
37+
}
38+
}
39+
}
40+
41+
pub(crate) struct DownstreamSubscriptionGuard {
42+
registry: DownstreamSubscriptionRegistry,
43+
keys: Vec<DownstreamSubscriptionKey>,
44+
}
45+
46+
impl Drop for DownstreamSubscriptionGuard {
47+
fn drop(&mut self) {
48+
self.registry.remove_all(&self.keys);
49+
}
50+
}
51+
52+
#[derive(Clone, Debug, Eq, Hash, PartialEq)]
53+
pub(crate) struct DownstreamSubscriptionKey {
54+
principal: String,
55+
virtual_host_id: String,
56+
subscription_id: RequestId,
57+
notification: DownstreamSubscriptionNotification,
58+
}
59+
60+
#[derive(Clone, Debug, Eq, Hash, PartialEq)]
61+
pub(crate) enum DownstreamSubscriptionNotification {
62+
ToolsListChanged,
63+
PromptsListChanged,
64+
ResourcesListChanged,
65+
ResourceUpdated { uri: String },
66+
}
67+
68+
pub(super) fn subscription_keys(
69+
context: &GatewayRequestContext,
70+
filter: &SubscriptionFilter,
71+
subscription_id: &RequestId,
72+
) -> Vec<DownstreamSubscriptionKey> {
73+
let mut keys = Vec::new();
74+
if filter.tools_list_changed == Some(true) {
75+
keys.push(subscription_key(context, subscription_id, DownstreamSubscriptionNotification::ToolsListChanged));
76+
}
77+
if filter.prompts_list_changed == Some(true) {
78+
keys.push(subscription_key(context, subscription_id, DownstreamSubscriptionNotification::PromptsListChanged));
79+
}
80+
if filter.resources_list_changed == Some(true) {
81+
keys.push(subscription_key(context, subscription_id, DownstreamSubscriptionNotification::ResourcesListChanged));
82+
}
83+
if let Some(uris) = &filter.resource_subscriptions {
84+
keys.extend(uris.iter().map(|uri| {
85+
subscription_key(
86+
context,
87+
subscription_id,
88+
DownstreamSubscriptionNotification::ResourceUpdated { uri: uri.clone() },
89+
)
90+
}));
91+
}
92+
keys
93+
}
94+
95+
fn subscription_key(
96+
context: &GatewayRequestContext,
97+
subscription_id: &RequestId,
98+
notification: DownstreamSubscriptionNotification,
99+
) -> DownstreamSubscriptionKey {
100+
DownstreamSubscriptionKey {
101+
principal: context.principal().to_owned(),
102+
virtual_host_id: context.virtual_host_id().to_owned(),
103+
subscription_id: subscription_id.clone(),
104+
notification,
105+
}
106+
}

0 commit comments

Comments
 (0)