soma_mcp_client/upstream/pool/
subject.rs1use 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;