soma_mcp_client/upstream/
pool.rs1use std::collections::BTreeMap;
2use std::sync::{Arc, RwLock};
3
4use serde_json::Value;
5
6use crate::config::UpstreamConfig;
7use crate::process::guard::SpawnGuard;
8use crate::upstream::http_client::{decide_http_transport, transport_kind_for_decision};
9use crate::upstream::{
10 ResponseCaps, ToolDescriptor, TransportKind, UpstreamError, UpstreamHealth, UpstreamSnapshot,
11};
12
13pub mod connect_stdio;
14pub mod discovery;
15pub mod health;
16pub mod live;
17pub mod prompts;
18pub mod resources;
19#[cfg(feature = "oauth")]
20pub mod subject;
21pub mod tools;
22
23#[derive(Debug, Clone, PartialEq, Eq)]
24pub struct PoolOptions {
25 pub response_caps: ResponseCaps,
26 pub discovery_concurrency: usize,
27}
28
29impl Default for PoolOptions {
30 fn default() -> Self {
31 Self {
32 response_caps: ResponseCaps::default(),
33 discovery_concurrency: 8,
34 }
35 }
36}
37
38impl PoolOptions {
39 #[must_use]
40 pub fn normalized(mut self) -> Self {
41 self.discovery_concurrency = self.discovery_concurrency.max(1);
42 self
43 }
44}
45
46#[derive(Debug, Clone, PartialEq, Eq)]
47pub struct ToolCall {
48 pub upstream: String,
49 pub tool: String,
50 pub params: Value,
51}
52
53#[derive(Debug, Clone)]
54pub struct InProcessUpstream {
55 snapshot: UpstreamSnapshot,
56 tool_results: BTreeMap<String, Value>,
57}
58
59impl InProcessUpstream {
60 #[must_use]
61 pub fn new(name: impl Into<String>) -> Self {
62 Self {
63 snapshot: UpstreamSnapshot::empty(name, TransportKind::InProcess),
64 tool_results: BTreeMap::new(),
65 }
66 }
67
68 #[must_use]
69 pub fn with_tool(mut self, tool: ToolDescriptor, result: Value) -> Self {
70 self.tool_results.insert(tool.name.clone(), result);
71 self.snapshot.tools.push(tool);
72 self
73 }
74
75 #[must_use]
76 pub fn with_snapshot(mut self, snapshot: UpstreamSnapshot) -> Self {
77 self.snapshot = snapshot;
78 self
79 }
80
81 fn call_tool(&self, call: &ToolCall) -> Result<Value, UpstreamError> {
82 if !call.params.is_object() {
83 return Err(UpstreamError::ParamsMustBeObject);
84 }
85 self.tool_results
86 .get(&call.tool)
87 .cloned()
88 .ok_or_else(|| UpstreamError::NotExposed {
89 upstream: call.upstream.clone(),
90 item: call.tool.clone(),
91 })
92 }
93}
94
95struct PoolEntry {
96 config: UpstreamConfig,
97 snapshot: UpstreamSnapshot,
98 in_process: Option<InProcessUpstream>,
99 live: Option<Arc<live::LiveUpstream>>,
100}
101
102#[cfg(feature = "oauth")]
103struct SubjectPoolEntry {
104 snapshot: UpstreamSnapshot,
105 live: Arc<live::LiveUpstream>,
106}
107
108#[derive(Clone)]
109pub struct UpstreamPool {
110 entries: Arc<RwLock<BTreeMap<String, PoolEntry>>>,
111 #[cfg(feature = "oauth")]
112 subject_entries: Arc<RwLock<BTreeMap<(String, String), SubjectPoolEntry>>>,
113 #[cfg(feature = "oauth")]
114 oauth_provider: Arc<RwLock<Option<Arc<dyn crate::oauth::UpstreamOAuthProvider>>>>,
115 options: PoolOptions,
116}
117
118impl Default for UpstreamPool {
119 fn default() -> Self {
120 Self::new(PoolOptions::default())
121 }
122}
123
124impl UpstreamPool {
125 #[must_use]
126 pub fn new(options: PoolOptions) -> Self {
127 Self {
128 entries: Arc::new(RwLock::new(BTreeMap::new())),
129 #[cfg(feature = "oauth")]
130 subject_entries: Arc::new(RwLock::new(BTreeMap::new())),
131 #[cfg(feature = "oauth")]
132 oauth_provider: Arc::new(RwLock::new(None)),
133 options: options.normalized(),
134 }
135 }
136
137 #[must_use]
138 pub fn response_caps(&self) -> &ResponseCaps {
139 &self.options.response_caps
140 }
141
142 #[must_use]
143 pub fn discovery_concurrency(&self) -> usize {
144 self.options.discovery_concurrency
145 }
146
147 pub fn register_config(&self, config: UpstreamConfig) -> Result<(), UpstreamError> {
148 let transport = transport_for_config(&config);
149 let health = if config.enabled {
150 health_for_config(config.name.as_str(), transport)
151 } else {
152 UpstreamHealth::Disabled
153 };
154 let mut snapshot = UpstreamSnapshot::empty(config.name.clone(), transport);
155 snapshot.health = health;
156 let entry = PoolEntry {
157 config: config.clone(),
158 snapshot,
159 in_process: None,
160 live: None,
161 };
162 self.entries
163 .write()
164 .expect("upstream pool lock poisoned")
165 .insert(config.name.clone(), entry);
166 Ok(())
167 }
168
169 pub fn register_in_process(
170 &self,
171 config: UpstreamConfig,
172 upstream: InProcessUpstream,
173 ) -> Result<(), UpstreamError> {
174 let mut snapshot = upstream.snapshot.clone();
175 snapshot.name = config.name.clone();
176 snapshot.transport = TransportKind::InProcess;
177 snapshot.health = if config.enabled {
178 UpstreamHealth::Connected
179 } else {
180 UpstreamHealth::Disabled
181 };
182 let entry = PoolEntry {
183 config: config.clone(),
184 snapshot,
185 in_process: Some(upstream),
186 live: None,
187 };
188 self.entries
189 .write()
190 .expect("upstream pool lock poisoned")
191 .insert(config.name.clone(), entry);
192 Ok(())
193 }
194
195 pub async fn call_tool(&self, call: ToolCall) -> Result<Value, UpstreamError> {
196 self.ensure_connected(&call.upstream).await?;
197 let live_peer = {
198 let entries = self.entries.read().expect("upstream pool lock poisoned");
199 let entry =
200 entries
201 .get(&call.upstream)
202 .ok_or_else(|| UpstreamError::UnknownUpstream {
203 upstream: call.upstream.clone(),
204 })?;
205 ensure_routable(entry)?;
206 tools::ensure_tool_exposed(entry, &call.tool)?;
207 if let Some(in_process) = &entry.in_process {
208 let result = in_process.call_tool(&call)?;
209 let bytes = serde_json::to_vec(&result).map_or(usize::MAX, |bytes| bytes.len());
210 self.response_caps()
211 .enforce(crate::upstream::CapScope::ToolsCall, bytes)?;
212 return Ok(result);
213 }
214 entry.live.as_ref().map(|live| live.peer())
215 };
216 let Some(peer) = live_peer else {
217 return Err(UpstreamError::Unsupported {
218 upstream: call.upstream,
219 capability: "tools/call",
220 });
221 };
222 let upstream = call.upstream.clone();
223 let result = live::call_live_tool(&upstream, peer, call.tool, call.params).await?;
224 let bytes = serde_json::to_vec(&result).map_or(usize::MAX, |bytes| bytes.len());
225 self.response_caps()
226 .enforce(crate::upstream::CapScope::ToolsCall, bytes)?;
227 Ok(result)
228 }
229
230 pub async fn ensure_connected(&self, upstream: &str) -> Result<(), UpstreamError> {
231 let config = {
232 let entries = self.entries.read().expect("upstream pool lock poisoned");
233 let entry = entries
234 .get(upstream)
235 .ok_or_else(|| UpstreamError::UnknownUpstream {
236 upstream: upstream.to_owned(),
237 })?;
238 if entry.in_process.is_some() || entry.live.is_some() || !entry.config.enabled {
239 return Ok(());
240 }
241 entry.config.clone()
242 };
243 let context = live::LiveConnectContext::shared(self.response_caps());
244 let (live, snapshot) = live::connect_live(&config, &SpawnGuard::default(), context).await?;
245 let mut entries = self.entries.write().expect("upstream pool lock poisoned");
246 let entry = entries
247 .get_mut(upstream)
248 .ok_or_else(|| UpstreamError::UnknownUpstream {
249 upstream: upstream.to_owned(),
250 })?;
251 entry.snapshot = snapshot;
252 entry.live = Some(Arc::new(live));
253 Ok(())
254 }
255
256 pub async fn refresh_all(&self) {
257 let names = self
258 .entries
259 .read()
260 .expect("upstream pool lock poisoned")
261 .keys()
262 .cloned()
263 .collect::<Vec<_>>();
264 for name in names {
265 if let Err(error) = self.ensure_connected(&name).await {
266 let _ = self.record_discovery_error(&name, error);
267 }
268 }
269 }
270
271 pub(super) fn record_discovery_error(
272 &self,
273 upstream: &str,
274 error: UpstreamError,
275 ) -> Result<(), UpstreamError> {
276 let mut entries = self.entries.write().expect("upstream pool lock poisoned");
277 let entry = entries
278 .get_mut(upstream)
279 .ok_or_else(|| UpstreamError::UnknownUpstream {
280 upstream: upstream.to_owned(),
281 })?;
282 match error {
283 UpstreamError::Unsupported {
284 upstream,
285 capability,
286 } => {
287 entry.snapshot.health = UpstreamHealth::Unsupported {
288 reason: format!("upstream `{upstream}` does not support `{capability}`"),
289 };
290 }
291 other => {
292 entry.snapshot.health = UpstreamHealth::Degraded {
293 consecutive_failures: 1,
294 error: Some(other.to_string()),
295 };
296 entry.snapshot.stale = true;
297 }
298 }
299 Ok(())
300 }
301
302 fn snapshots(&self) -> Vec<UpstreamSnapshot> {
303 self.entries
304 .read()
305 .expect("upstream pool lock poisoned")
306 .values()
307 .map(|entry| entry.snapshot.clone())
308 .collect()
309 }
310
311 fn with_entry<T>(
312 &self,
313 upstream: &str,
314 f: impl FnOnce(&PoolEntry) -> Result<T, UpstreamError>,
315 ) -> Result<T, UpstreamError> {
316 let entries = self.entries.read().expect("upstream pool lock poisoned");
317 let entry = entries
318 .get(upstream)
319 .ok_or_else(|| UpstreamError::UnknownUpstream {
320 upstream: upstream.to_owned(),
321 })?;
322 f(entry)
323 }
324}
325
326fn ensure_routable(entry: &PoolEntry) -> Result<(), UpstreamError> {
327 if entry.snapshot.health.is_routable() {
328 return Ok(());
329 }
330 Err(UpstreamError::NotRoutable {
331 upstream: entry.snapshot.name.clone(),
332 reason: health_reason(&entry.snapshot.health),
333 })
334}
335
336fn health_reason(health: &UpstreamHealth) -> String {
337 match health {
338 UpstreamHealth::Connected => "connected".to_owned(),
339 UpstreamHealth::Disabled => "disabled".to_owned(),
340 UpstreamHealth::Degraded { error, .. } => error
341 .clone()
342 .unwrap_or_else(|| "capability degraded".to_owned()),
343 UpstreamHealth::Unsupported { reason } => reason.clone(),
344 }
345}
346
347fn transport_for_config(config: &UpstreamConfig) -> TransportKind {
348 if let Some(url) = config.url.as_deref() {
349 return transport_kind_for_decision(&decide_http_transport(url));
350 }
351 if config.command.is_some() {
352 return TransportKind::Stdio;
353 }
354 TransportKind::InProcess
355}
356
357fn health_for_config(name: &str, transport: TransportKind) -> UpstreamHealth {
358 match transport {
359 TransportKind::InProcess => UpstreamHealth::Unsupported {
360 reason: format!("configured upstream `{name}` has no live in-process connector"),
361 },
362 TransportKind::HttpJson | TransportKind::HttpSse | TransportKind::WebSocket => {
363 UpstreamHealth::Unsupported {
364 reason: format!("live upstream `{name}` is not connected yet"),
365 }
366 }
367 TransportKind::Stdio => UpstreamHealth::Unsupported {
368 reason: format!("live stdio upstream `{name}` is not connected yet"),
369 },
370 }
371}
372
373#[cfg(test)]
374#[path = "pool_tests.rs"]
375mod tests;