Skip to main content

soma_mcp_client/upstream/pool/
subject.rs

1use serde_json::{Map, Value};
2
3use crate::oauth::UpstreamOAuthProvider;
4use crate::process::guard::SpawnGuard;
5use crate::upstream::{
6    CapScope, PromptDescriptor, ResourceDescriptor, ToolDescriptor, UpstreamError, UpstreamHealth,
7    UpstreamSnapshot,
8};
9
10use super::tools::matches_filter;
11use super::{live, SubjectPoolEntry, ToolCall, UpstreamPool};
12
13impl UpstreamPool {
14    pub fn install_oauth_provider(&self, provider: std::sync::Arc<dyn UpstreamOAuthProvider>) {
15        *self
16            .oauth_provider
17            .write()
18            .expect("oauth provider lock poisoned") = Some(provider);
19        self.subject_entries
20            .write()
21            .expect("subject pool lock poisoned")
22            .clear();
23    }
24
25    pub fn evict_oauth_subject(&self, upstream: &str, subject: &str) {
26        self.subject_entries
27            .write()
28            .expect("subject pool lock poisoned")
29            .remove(&(upstream.to_owned(), subject.to_owned()));
30        if let Some(provider) = self
31            .oauth_provider
32            .read()
33            .expect("oauth provider lock poisoned")
34            .as_ref()
35        {
36            provider.evict_subject(upstream, subject);
37        }
38    }
39
40    pub async fn discover_for_subject(
41        &self,
42        subject: Option<&str>,
43    ) -> Result<Vec<UpstreamSnapshot>, UpstreamError> {
44        let Some(subject) = subject else {
45            return self.discover().await;
46        };
47        let configs = self.configured_upstreams();
48        for (name, oauth_enabled) in configs {
49            let result = if oauth_enabled {
50                self.ensure_subject_connected(&name, subject).await
51            } else {
52                self.ensure_connected(&name).await
53            };
54            if let Err(error) = result {
55                tracing::warn!(upstream = %name, subject, error = %error, "subject discovery failed");
56            }
57        }
58        Ok(self.snapshots_for_subject(subject))
59    }
60
61    pub async fn exposed_tools_for_subject(
62        &self,
63        upstream: &str,
64        subject: Option<&str>,
65    ) -> Result<Vec<ToolDescriptor>, UpstreamError> {
66        let Some(subject) = subject else {
67            return self.exposed_tools(upstream);
68        };
69        let Some((snapshot, config)) = self.subject_snapshot_and_config(upstream, subject)? else {
70            return self.exposed_tools(upstream);
71        };
72        let tools: Vec<ToolDescriptor> = snapshot
73            .tools
74            .into_iter()
75            .filter(|tool| matches_filter(config.expose_tools.as_deref(), &tool.name))
76            .collect();
77        let bytes = serde_json::to_vec(&tools).map_or(usize::MAX, |bytes| bytes.len());
78        self.response_caps().enforce(CapScope::ToolsList, bytes)?;
79        Ok(tools)
80    }
81
82    pub async fn call_tool_for_subject(
83        &self,
84        call: ToolCall,
85        subject: Option<&str>,
86    ) -> Result<Value, UpstreamError> {
87        let Some(subject) = subject else {
88            return self.call_tool(call).await;
89        };
90        if !self.config_is_oauth(&call.upstream)? {
91            return self.call_tool(call).await;
92        }
93        self.ensure_subject_connected(&call.upstream, subject)
94            .await?;
95        let (peer, upstream) = self.with_subject_entry(&call.upstream, subject, |entry| {
96            ensure_subject_routable(&entry.snapshot)?;
97            if !entry
98                .snapshot
99                .tools
100                .iter()
101                .any(|candidate| candidate.name == call.tool)
102            {
103                return Err(UpstreamError::NotExposed {
104                    upstream: call.upstream.clone(),
105                    item: call.tool.clone(),
106                });
107            }
108            Ok((entry.live.peer(), entry.snapshot.name.clone()))
109        })?;
110        let result = live::call_live_tool(&upstream, peer, call.tool, call.params).await?;
111        let bytes = serde_json::to_vec(&result).map_or(usize::MAX, |bytes| bytes.len());
112        self.response_caps().enforce(CapScope::ToolsCall, bytes)?;
113        Ok(result)
114    }
115
116    pub async fn list_resources_for_subject(
117        &self,
118        upstream: &str,
119        subject: Option<&str>,
120    ) -> Result<Vec<ResourceDescriptor>, UpstreamError> {
121        let Some(subject) = subject else {
122            return self.list_resources(upstream).await;
123        };
124        let Some((snapshot, config)) = self.subject_snapshot_and_config(upstream, subject)? else {
125            return self.list_resources(upstream).await;
126        };
127        if !config.proxy_resources {
128            return Ok(Vec::new());
129        }
130        let resources: Vec<ResourceDescriptor> = snapshot
131            .resources
132            .into_iter()
133            .filter(|resource| matches_filter(config.expose_resources.as_deref(), &resource.uri))
134            .collect();
135        let bytes = serde_json::to_vec(&resources).map_or(usize::MAX, |bytes| bytes.len());
136        self.response_caps()
137            .enforce(CapScope::ResourcesList, bytes)?;
138        Ok(resources)
139    }
140
141    pub async fn read_resource_for_subject(
142        &self,
143        upstream: &str,
144        uri: &str,
145        subject: Option<&str>,
146    ) -> Result<Value, UpstreamError> {
147        let Some(subject) = subject else {
148            return self.read_resource(upstream, uri).await;
149        };
150        if !self.config_is_oauth(upstream)? {
151            return self.read_resource(upstream, uri).await;
152        }
153        self.ensure_subject_connected(upstream, subject).await?;
154        let peer = self.with_subject_entry(upstream, subject, |entry| {
155            ensure_subject_routable(&entry.snapshot)?;
156            Ok(entry.live.peer())
157        })?;
158        let value = live::read_live_resource(upstream, peer, uri.to_owned()).await?;
159        let bytes = serde_json::to_vec(&value).map_or(usize::MAX, |bytes| bytes.len());
160        self.response_caps()
161            .enforce(CapScope::ResourcesRead, bytes)?;
162        Ok(value)
163    }
164
165    pub async fn list_prompts_for_subject(
166        &self,
167        upstream: &str,
168        subject: Option<&str>,
169    ) -> Result<Vec<PromptDescriptor>, UpstreamError> {
170        let Some(subject) = subject else {
171            return self.list_prompts(upstream).await;
172        };
173        let Some((snapshot, config)) = self.subject_snapshot_and_config(upstream, subject)? else {
174            return self.list_prompts(upstream).await;
175        };
176        if !config.proxy_prompts {
177            return Ok(Vec::new());
178        }
179        let prompts: Vec<PromptDescriptor> = snapshot
180            .prompts
181            .into_iter()
182            .filter(|prompt| matches_filter(config.expose_prompts.as_deref(), &prompt.name))
183            .collect();
184        let bytes = serde_json::to_vec(&prompts).map_or(usize::MAX, |bytes| bytes.len());
185        self.response_caps().enforce(CapScope::PromptsList, bytes)?;
186        Ok(prompts)
187    }
188
189    pub async fn get_prompt_for_subject(
190        &self,
191        upstream: &str,
192        name: &str,
193        arguments: Option<Map<String, Value>>,
194        subject: Option<&str>,
195    ) -> Result<Value, UpstreamError> {
196        let Some(subject) = subject else {
197            return self.get_prompt(upstream, name, arguments).await;
198        };
199        if !self.config_is_oauth(upstream)? {
200            return self.get_prompt(upstream, name, arguments).await;
201        }
202        self.ensure_subject_connected(upstream, subject).await?;
203        let peer = self.with_subject_entry(upstream, subject, |entry| {
204            ensure_subject_routable(&entry.snapshot)?;
205            Ok(entry.live.peer())
206        })?;
207        let value = live::get_live_prompt(upstream, peer, name.to_owned(), arguments).await?;
208        let bytes = serde_json::to_vec(&value).map_or(usize::MAX, |bytes| bytes.len());
209        self.response_caps().enforce(CapScope::PromptsGet, bytes)?;
210        Ok(value)
211    }
212
213    async fn ensure_subject_connected(
214        &self,
215        upstream: &str,
216        subject: &str,
217    ) -> Result<(), UpstreamError> {
218        let key = (upstream.to_owned(), subject.to_owned());
219        if self
220            .subject_entries
221            .read()
222            .expect("subject pool lock poisoned")
223            .contains_key(&key)
224        {
225            return Ok(());
226        }
227        let config = self.config_for_subject(upstream)?;
228        if !config.enabled {
229            return Ok(());
230        }
231        if config.oauth.is_none() {
232            return self.ensure_connected(upstream).await;
233        }
234        let provider = self.oauth_provider()?;
235        let context = live::LiveConnectContext::oauth(self.response_caps(), subject, provider);
236        let (live, snapshot) = live::connect_live(&config, &SpawnGuard::default(), context).await?;
237        self.subject_entries
238            .write()
239            .expect("subject pool lock poisoned")
240            .insert(
241                key,
242                SubjectPoolEntry {
243                    snapshot,
244                    live: std::sync::Arc::new(live),
245                },
246            );
247        Ok(())
248    }
249
250    fn configured_upstreams(&self) -> Vec<(String, bool)> {
251        self.entries
252            .read()
253            .expect("upstream pool lock poisoned")
254            .iter()
255            .map(|(name, entry)| (name.clone(), entry.config.oauth.is_some()))
256            .collect()
257    }
258
259    fn snapshots_for_subject(&self, subject: &str) -> Vec<UpstreamSnapshot> {
260        let entries = self.entries.read().expect("upstream pool lock poisoned");
261        let subject_entries = self
262            .subject_entries
263            .read()
264            .expect("subject pool lock poisoned");
265        entries
266            .iter()
267            .filter_map(|(name, entry)| {
268                if entry.config.oauth.is_some() {
269                    return subject_entries
270                        .get(&(name.clone(), subject.to_owned()))
271                        .map(|entry| entry.snapshot.clone());
272                }
273                Some(entry.snapshot.clone())
274            })
275            .collect()
276    }
277
278    fn subject_snapshot_and_config(
279        &self,
280        upstream: &str,
281        subject: &str,
282    ) -> Result<Option<(UpstreamSnapshot, crate::config::UpstreamConfig)>, UpstreamError> {
283        if !self.config_is_oauth(upstream)? {
284            return Ok(None);
285        }
286        let snapshot =
287            self.with_subject_entry(upstream, subject, |entry| Ok(entry.snapshot.clone()))?;
288        let config = self.config_for_subject(upstream)?;
289        Ok(Some((snapshot, config)))
290    }
291
292    fn with_subject_entry<T>(
293        &self,
294        upstream: &str,
295        subject: &str,
296        f: impl FnOnce(&SubjectPoolEntry) -> Result<T, UpstreamError>,
297    ) -> Result<T, UpstreamError> {
298        let entries = self
299            .subject_entries
300            .read()
301            .expect("subject pool lock poisoned");
302        let entry = entries
303            .get(&(upstream.to_owned(), subject.to_owned()))
304            .ok_or_else(|| UpstreamError::NotRoutable {
305                upstream: upstream.to_owned(),
306                reason: "subject connection is not established".to_owned(),
307            })?;
308        f(entry)
309    }
310
311    fn oauth_provider(&self) -> Result<std::sync::Arc<dyn UpstreamOAuthProvider>, UpstreamError> {
312        self.oauth_provider
313            .read()
314            .expect("oauth provider lock poisoned")
315            .clone()
316            .ok_or_else(|| UpstreamError::LiveConnect {
317                upstream: "oauth".to_owned(),
318                message: "upstream OAuth runtime is not configured".to_owned(),
319            })
320    }
321
322    fn config_for_subject(
323        &self,
324        upstream: &str,
325    ) -> Result<crate::config::UpstreamConfig, UpstreamError> {
326        self.with_entry(upstream, |entry| Ok(entry.config.clone()))
327    }
328
329    fn config_is_oauth(&self, upstream: &str) -> Result<bool, UpstreamError> {
330        self.with_entry(upstream, |entry| Ok(entry.config.oauth.is_some()))
331    }
332}
333
334fn ensure_subject_routable(snapshot: &UpstreamSnapshot) -> Result<(), UpstreamError> {
335    if snapshot.health == UpstreamHealth::Connected {
336        return Ok(());
337    }
338    Err(UpstreamError::NotRoutable {
339        upstream: snapshot.name.clone(),
340        reason: "subject-scoped upstream is not connected".to_owned(),
341    })
342}
343
344#[cfg(test)]
345#[path = "subject_tests.rs"]
346mod tests;