Skip to main content

soma_mcp_client/upstream/
pool.rs

1use 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;