Skip to main content

mcp_utils/client/
config.rs

1use futures::future::BoxFuture;
2use rmcp::{RoleServer, service::DynService, transport::streamable_http_client::StreamableHttpClientTransportConfig};
3use schemars::JsonSchema;
4use serde::{Deserialize, Serialize};
5use serde_json::Value;
6use std::collections::{BTreeMap, HashMap};
7use std::fmt::{Debug, Formatter};
8use std::num::NonZeroU16;
9use std::path::Path;
10use utils::is_false;
11use utils::variables::{VarError, Vars};
12
13#[derive(Debug, Clone, Default, Deserialize, Serialize, JsonSchema)]
14pub struct McpConfig {
15    #[serde(alias = "mcpServers")]
16    pub servers: BTreeMap<String, McpServerConfig>,
17}
18
19#[doc = include_str!("../docs/mcp_server_config.md")]
20#[derive(Debug, Clone, Deserialize, Serialize, JsonSchema, PartialEq)]
21#[serde(untagged)]
22pub enum McpServerConfig {
23    Stdio(StdioServerConfig),
24    Remote(RemoteServerConfig),
25    InMemory(InMemoryServerConfig),
26}
27
28#[derive(Debug, Clone, Deserialize, Serialize, JsonSchema, PartialEq)]
29#[serde(deny_unknown_fields)]
30pub struct StdioServerConfig {
31    /// Transport discriminant; always `stdio`.
32    #[serde(rename = "type", default)]
33    pub type_: StdioType,
34
35    /// Executable launched to run the MCP server over stdio.
36    pub command: String,
37
38    /// Command-line arguments passed to the executable.
39    #[serde(default)]
40    pub args: Vec<String>,
41
42    /// Environment variables set for the server process.
43    #[serde(default)]
44    pub env: HashMap<String, String>,
45
46    /// Expose this server's tools through Aether's tool proxy.
47    #[serde(default, skip_serializing_if = "is_false")]
48    pub proxy: bool,
49}
50
51#[derive(Debug, Clone, Deserialize, Serialize, JsonSchema, PartialEq)]
52#[serde(rename_all = "camelCase", deny_unknown_fields)]
53pub struct McpOAuthConfig {
54    pub client_id: String,
55    pub callback_port: NonZeroU16,
56}
57
58#[derive(Debug, Clone, Deserialize, Serialize, JsonSchema, PartialEq)]
59#[serde(deny_unknown_fields)]
60pub struct RemoteServerConfig {
61    /// Transport discriminant; `http` (streamable HTTP) or `sse` (Server-Sent Events).
62    #[serde(rename = "type")]
63    pub type_: RemoteType,
64
65    /// Base URL of the remote MCP server.
66    pub url: String,
67
68    /// Extra HTTP headers sent with every request.
69    #[serde(default)]
70    pub headers: HashMap<String, String>,
71
72    /// OAuth settings for a pre-registered public client.
73    #[serde(default, skip_serializing_if = "Option::is_none")]
74    pub oauth: Option<McpOAuthConfig>,
75
76    /// Expose this server's tools through Aether's tool proxy.
77    #[serde(default, skip_serializing_if = "is_false")]
78    pub proxy: bool,
79}
80
81#[derive(Debug, Clone, Deserialize, Serialize, JsonSchema, PartialEq)]
82#[serde(deny_unknown_fields)]
83pub struct InMemoryServerConfig {
84    /// Transport discriminant; always `in-memory`.
85    #[serde(rename = "type")]
86    pub type_: InMemoryType,
87
88    /// Arguments passed to the built-in (in-process) server.
89    #[serde(default)]
90    pub args: Vec<String>,
91
92    /// Optional JSON input passed to the built-in server at startup.
93    #[serde(default)]
94    pub input: Option<Value>,
95
96    /// Expose this server's tools through Aether's tool proxy.
97    #[serde(default, skip_serializing_if = "is_false")]
98    pub proxy: bool,
99}
100
101#[derive(Debug, Clone, Copy, Default, Deserialize, Serialize, JsonSchema, PartialEq)]
102pub enum StdioType {
103    #[default]
104    #[serde(rename = "stdio")]
105    Stdio,
106}
107
108#[derive(Debug, Clone, Copy, Deserialize, Serialize, JsonSchema, PartialEq)]
109pub enum RemoteType {
110    #[serde(rename = "http")]
111    Http,
112    #[serde(rename = "sse")]
113    Sse,
114}
115
116#[derive(Debug, Clone, Copy, Deserialize, Serialize, JsonSchema, PartialEq)]
117pub enum InMemoryType {
118    #[serde(rename = "in-memory")]
119    InMemory,
120}
121
122pub struct McpServer {
123    pub name: String,
124    pub transport: McpTransport,
125    pub proxy: bool,
126}
127
128pub enum McpTransport {
129    Stdio { command: String, args: Vec<String>, env: HashMap<String, String> },
130    Http(McpHttpConfig),
131    InMemory { server: Box<dyn DynService<RoleServer>> },
132}
133
134#[derive(Debug, Clone)]
135pub struct McpHttpConfig {
136    pub transport: StreamableHttpClientTransportConfig,
137    pub oauth: Option<McpOAuthConfig>,
138}
139
140impl McpHttpConfig {
141    pub fn oauth_client_id(&self) -> Option<&str> {
142        self.oauth.as_ref().map(|oauth| oauth.client_id.as_str())
143    }
144
145    pub fn callback_port(&self) -> Option<NonZeroU16> {
146        self.oauth.as_ref().map(|oauth| oauth.callback_port)
147    }
148}
149
150impl From<StreamableHttpClientTransportConfig> for McpHttpConfig {
151    fn from(transport: StreamableHttpClientTransportConfig) -> Self {
152        Self { transport, oauth: None }
153    }
154}
155
156impl McpServer {
157    pub fn new(name: impl Into<String>, transport: McpTransport, proxy: bool) -> Self {
158        Self { name: name.into(), transport, proxy }
159    }
160
161    /// Clone this server config. Fails for [`McpTransport::InMemory`], whose
162    /// boxed service cannot be duplicated and so cannot be shared across
163    /// independently-spawned MCP managers.
164    pub fn try_clone(&self) -> Result<Self, McpServerCloneError> {
165        let transport = match &self.transport {
166            McpTransport::Stdio { command, args, env } => {
167                McpTransport::Stdio { command: command.clone(), args: args.clone(), env: env.clone() }
168            }
169            McpTransport::Http(config) => McpTransport::Http(config.clone()),
170            McpTransport::InMemory { .. } => return Err(McpServerCloneError(self.name.clone())),
171        };
172        Ok(Self { name: self.name.clone(), transport, proxy: self.proxy })
173    }
174}
175
176#[derive(Debug, thiserror::Error)]
177#[error("in-memory MCP server `{0}` cannot be cloned across runtimes")]
178pub struct McpServerCloneError(pub String);
179
180impl Debug for McpServer {
181    fn fmt(&self, f: &mut Formatter<'_>) -> std::fmt::Result {
182        f.debug_struct("McpServer")
183            .field("name", &self.name)
184            .field("transport", &self.transport)
185            .field("proxy", &self.proxy)
186            .finish()
187    }
188}
189
190impl Debug for McpTransport {
191    fn fmt(&self, f: &mut Formatter<'_>) -> std::fmt::Result {
192        match self {
193            McpTransport::Stdio { command, args, env } => {
194                f.debug_struct("Stdio").field("command", command).field("args", args).field("env", env).finish()
195            }
196            McpTransport::Http(config) => f.debug_tuple("Http").field(config).finish(),
197            McpTransport::InMemory { .. } => f.debug_struct("InMemory").field("server", &"<DynService>").finish(),
198        }
199    }
200}
201
202pub type ServerFactory =
203    Box<dyn Fn(Vec<String>, Option<Value>) -> BoxFuture<'static, Box<dyn DynService<RoleServer>>> + Send + Sync>;
204
205#[derive(Debug, thiserror::Error)]
206pub enum ParseError {
207    #[error("Failed to read config file: {0}")]
208    IoError(#[from] std::io::Error),
209
210    #[error("Invalid JSON: {0}")]
211    JsonError(#[from] serde_json::Error),
212
213    #[error("Variable expansion failed: {0}")]
214    VarError(#[from] VarError),
215
216    #[error("InMemory server factory '{0}' not registered")]
217    FactoryNotFound(String),
218
219    #[error("Invalid nested config in tool-proxy: {0}")]
220    InvalidNestedConfig(String),
221}
222
223impl McpConfig {
224    pub fn new(servers: BTreeMap<String, McpServerConfig>) -> Self {
225        Self { servers }
226    }
227
228    pub fn from_json_file(path: impl AsRef<Path>) -> Result<Self, ParseError> {
229        let content = std::fs::read_to_string(path)?;
230        Self::from_json(&content)
231    }
232
233    pub fn from_json_files<T: AsRef<Path>>(paths: &[T]) -> Result<Self, ParseError> {
234        let mut merged = BTreeMap::new();
235        for path in paths {
236            let raw = Self::from_json_file(path)?;
237            merged.extend(raw.servers);
238        }
239        Ok(Self::new(merged))
240    }
241
242    pub fn from_json(json: &str) -> Result<Self, ParseError> {
243        Ok(serde_json::from_str(json)?)
244    }
245
246    pub async fn into_servers(
247        self,
248        factories: &HashMap<String, ServerFactory>,
249        vars: &Vars,
250    ) -> Result<Vec<McpServer>, ParseError> {
251        self.into_servers_with_proxy(factories, vars, false).await
252    }
253
254    pub async fn into_servers_with_proxy(
255        self,
256        factories: &HashMap<String, ServerFactory>,
257        vars: &Vars,
258        force_proxy: bool,
259    ) -> Result<Vec<McpServer>, ParseError> {
260        let mut servers = Vec::with_capacity(self.servers.len());
261        for (name, config) in self.servers {
262            servers.push(config.into_server(name, factories, vars, force_proxy).await?);
263        }
264        Ok(servers)
265    }
266
267    pub fn mark_all_proxy(&mut self) {
268        for server in self.servers.values_mut() {
269            server.set_proxy(true);
270        }
271    }
272}
273
274impl McpServerConfig {
275    pub fn proxy(&self) -> bool {
276        match self {
277            McpServerConfig::Stdio(config) => config.proxy,
278            McpServerConfig::Remote(config) => config.proxy,
279            McpServerConfig::InMemory(config) => config.proxy,
280        }
281    }
282
283    pub fn set_proxy(&mut self, value: bool) {
284        match self {
285            McpServerConfig::Stdio(config) => config.proxy = value,
286            McpServerConfig::Remote(config) => config.proxy = value,
287            McpServerConfig::InMemory(config) => config.proxy = value,
288        }
289    }
290
291    pub async fn into_server(
292        self,
293        name: String,
294        factories: &HashMap<String, ServerFactory>,
295        vars: &Vars,
296        force_proxy: bool,
297    ) -> Result<McpServer, ParseError> {
298        let proxy = force_proxy || self.proxy();
299        let transport = self.into_transport(name.clone(), factories, vars).await?;
300        Ok(McpServer::new(name, transport, proxy))
301    }
302
303    async fn into_transport(
304        self,
305        name: String,
306        factories: &HashMap<String, ServerFactory>,
307        vars: &Vars,
308    ) -> Result<McpTransport, ParseError> {
309        match self {
310            McpServerConfig::Stdio(StdioServerConfig { command, args, env, .. }) => Ok(McpTransport::Stdio {
311                command: vars.expand(&command)?,
312                args: args.into_iter().map(|a| vars.expand(&a)).collect::<Result<Vec<_>, _>>()?,
313                env: env
314                    .into_iter()
315                    .map(|(k, v)| Ok((k, vars.expand(&v)?)))
316                    .collect::<Result<HashMap<_, _>, VarError>>()?,
317            }),
318
319            McpServerConfig::Remote(RemoteServerConfig { url, headers, oauth, .. }) => {
320                let auth_header = headers.get("Authorization").map(|v| vars.expand(v)).transpose()?.map(|auth| {
321                    // rmcp adds `Bearer`  to the auth header.
322                    auth.split_once(' ')
323                        .filter(|(scheme, _)| scheme.eq_ignore_ascii_case("Bearer"))
324                        .map_or(auth.as_str(), |(_, rest)| rest)
325                        .to_string()
326                });
327
328                let mut transport = StreamableHttpClientTransportConfig::with_uri(vars.expand(&url)?);
329                if let Some(auth) = auth_header {
330                    transport = transport.auth_header(auth);
331                }
332
333                let oauth = oauth
334                    .map(|oauth| -> Result<McpOAuthConfig, VarError> {
335                        Ok(McpOAuthConfig { client_id: vars.expand(&oauth.client_id)?, ..oauth })
336                    })
337                    .transpose()?;
338
339                Ok(McpTransport::Http(McpHttpConfig { transport, oauth }))
340            }
341
342            McpServerConfig::InMemory(InMemoryServerConfig { args, input, .. }) => {
343                let server_factory = factories.get(&name).ok_or_else(|| ParseError::FactoryNotFound(name.clone()))?;
344                let expanded_args = args.into_iter().map(|a| vars.expand(&a)).collect::<Result<Vec<_>, VarError>>()?;
345                let server = server_factory(expanded_args, input).await;
346                Ok(McpTransport::InMemory { server })
347            }
348        }
349    }
350}
351
352#[cfg(test)]
353mod tests {
354    use super::*;
355    use std::fs;
356    use tempfile::tempdir;
357
358    fn write_config(dir: &Path, name: &str, json: &str) -> std::path::PathBuf {
359        let path = dir.join(name);
360        fs::write(&path, json).unwrap();
361        path
362    }
363
364    fn stdio_config(command: &str) -> String {
365        format!(r#"{{"servers": {{"coding": {{"type": "stdio", "command": "{command}"}}}}}}"#)
366    }
367
368    #[test]
369    fn from_json_accepts_mcp_servers_key() {
370        let config = McpConfig::from_json(r#"{"mcpServers": {"alpha": {"type": "stdio", "command": "a"}}}"#).unwrap();
371        assert_eq!(config.servers.len(), 1);
372        assert!(config.servers.contains_key("alpha"));
373    }
374
375    #[test]
376    fn from_json_defaults_missing_type_to_stdio() {
377        let config = McpConfig::from_json(
378            r#"{"mcpServers": {"devtools": {"command": "npx", "args": ["-y", "chrome-devtools-mcp"]}}}"#,
379        )
380        .unwrap();
381        match config.servers.get("devtools").unwrap() {
382            McpServerConfig::Stdio(StdioServerConfig { command, args, proxy, .. }) => {
383                assert_eq!(command, "npx");
384                assert_eq!(args, &["-y", "chrome-devtools-mcp"]);
385                assert!(!proxy);
386            }
387            other => panic!("expected Stdio server, got {other:?}"),
388        }
389    }
390
391    #[test]
392    fn from_json_accepts_server_proxy_true() {
393        let config =
394            McpConfig::from_json(r#"{"servers": {"playwright": {"type": "stdio", "command": "npx", "proxy": true}}}"#)
395                .unwrap();
396        assert!(config.servers.get("playwright").unwrap().proxy());
397    }
398
399    #[test]
400    fn from_json_rejects_proxy_server_type() {
401        let result = McpConfig::from_json(r#"{"servers":{"tools":{"type":"proxy","servers":{}}}}"#);
402        assert!(result.is_err());
403    }
404
405    #[test]
406    fn false_proxy_omits_during_serialization() {
407        let config =
408            McpConfig::from_json(r#"{"servers": {"coding": {"type": "stdio", "command": "a", "proxy": false}}}"#)
409                .unwrap();
410        let serialized = serde_json::to_string(&config).unwrap();
411        assert!(!serialized.contains("proxy"));
412    }
413
414    #[test]
415    fn true_proxy_serializes() {
416        let config =
417            McpConfig::from_json(r#"{"servers": {"coding": {"type": "stdio", "command": "a", "proxy": true}}}"#)
418                .unwrap();
419        let serialized = serde_json::to_string(&config).unwrap();
420        assert!(serialized.contains("proxy"));
421    }
422
423    #[test]
424    fn from_json_rejects_unknown_type() {
425        let result = McpConfig::from_json(r#"{"servers": {"bad": {"type": "htp", "url": "https://example.com"}}}"#);
426        assert!(result.is_err());
427    }
428
429    #[test]
430    fn from_json_files_empty_returns_empty_servers() {
431        let result = McpConfig::from_json_files::<&str>(&[]).unwrap();
432        assert!(result.servers.is_empty());
433    }
434
435    #[test]
436    fn from_json_files_single_file_matches_from_json_file() {
437        let dir = tempdir().unwrap();
438        let path = write_config(dir.path(), "a.json", &stdio_config("ls"));
439
440        let single = McpConfig::from_json_file(&path).unwrap();
441        let multi = McpConfig::from_json_files(&[&path]).unwrap();
442
443        assert_eq!(single.servers.len(), multi.servers.len());
444        assert!(multi.servers.contains_key("coding"));
445    }
446
447    #[test]
448    fn from_json_files_merges_disjoint_servers() {
449        let dir = tempdir().unwrap();
450        let a = write_config(dir.path(), "a.json", r#"{"servers": {"alpha": {"type": "stdio", "command": "a"}}}"#);
451        let b = write_config(dir.path(), "b.json", r#"{"servers": {"beta": {"type": "stdio", "command": "b"}}}"#);
452
453        let merged = McpConfig::from_json_files(&[a, b]).unwrap();
454        assert_eq!(merged.servers.len(), 2);
455        assert!(merged.servers.contains_key("alpha"));
456        assert!(merged.servers.contains_key("beta"));
457    }
458
459    #[test]
460    fn from_json_files_last_file_wins_on_collision_including_proxy() {
461        let dir = tempdir().unwrap();
462        let a = write_config(
463            dir.path(),
464            "a.json",
465            r#"{"servers":{"coding":{"type":"stdio","command":"from_a","proxy":true}}}"#,
466        );
467        let b = write_config(dir.path(), "b.json", r#"{"servers":{"coding":{"type":"stdio","command":"from_b"}}}"#);
468
469        let merged_ab = McpConfig::from_json_files(&[&a, &b]).unwrap();
470        match merged_ab.servers.get("coding").unwrap() {
471            McpServerConfig::Stdio(StdioServerConfig { command, proxy, .. }) => {
472                assert_eq!(command, "from_b");
473                assert!(!proxy);
474            }
475            other => panic!("expected Stdio, got {other:?}"),
476        }
477
478        let merged_ba = McpConfig::from_json_files(&[&b, &a]).unwrap();
479        match merged_ba.servers.get("coding").unwrap() {
480            McpServerConfig::Stdio(StdioServerConfig { command, proxy, .. }) => {
481                assert_eq!(command, "from_a");
482                assert!(*proxy);
483            }
484            other => panic!("expected Stdio, got {other:?}"),
485        }
486    }
487
488    #[test]
489    fn mark_all_proxy_sets_every_server() {
490        let mut config = McpConfig::from_json(
491            r#"{"servers":{"a":{"type":"stdio","command":"a"},"b":{"type":"http","url":"https://example.com"}}}"#,
492        )
493        .unwrap();
494        config.mark_all_proxy();
495        assert!(config.servers.values().all(McpServerConfig::proxy));
496    }
497
498    #[test]
499    fn from_json_files_propagates_io_error_on_missing_file() {
500        let dir = tempdir().unwrap();
501        let missing = dir.path().join("does-not-exist.json");
502        let result = McpConfig::from_json_files(&[missing]);
503        assert!(matches!(result, Err(ParseError::IoError(_))));
504    }
505
506    #[test]
507    fn from_json_files_propagates_json_error_on_invalid_file() {
508        let dir = tempdir().unwrap();
509        let bad = write_config(dir.path(), "bad.json", "not valid json");
510        let result = McpConfig::from_json_files(&[bad]);
511        assert!(matches!(result, Err(ParseError::JsonError(_))));
512    }
513
514    #[tokio::test]
515    async fn into_servers_preserves_proxy_flags() {
516        let json = r#"{
517            "servers": {
518                "github": {"type": "stdio", "command": "g"},
519                "playwright": {"type": "stdio", "command": "p", "proxy": true}
520            }
521        }"#;
522        let config = McpConfig::from_json(json).unwrap();
523        let servers = config.into_servers(&HashMap::new(), &Vars::new()).await.unwrap();
524
525        assert_eq!(servers.len(), 2);
526        assert!(!servers.iter().find(|s| s.name == "github").unwrap().proxy);
527        assert!(servers.iter().find(|s| s.name == "playwright").unwrap().proxy);
528    }
529
530    #[tokio::test]
531    async fn into_servers_with_proxy_forces_proxy_flags() {
532        let config =
533            McpConfig::from_json(r#"{"servers":{"github":{"type":"stdio","command":"g","proxy":false}}}"#).unwrap();
534        let servers = config.into_servers_with_proxy(&HashMap::new(), &Vars::new(), true).await.unwrap();
535        assert!(servers[0].proxy);
536    }
537
538    #[tokio::test]
539    async fn into_transport_expands_workspace_var_in_stdio_args() {
540        let config = McpConfig::from_json(
541            r#"{"servers":{"coding":{"type":"stdio","command":"server","args":["--root","${WORKSPACE}/src"]}}}"#,
542        )
543        .unwrap();
544        let vars = Vars::new().with("WORKSPACE", "/workspace");
545        let servers = config.into_servers(&HashMap::new(), &vars).await.unwrap();
546
547        match &servers[0].transport {
548            McpTransport::Stdio { args, .. } => {
549                assert_eq!(args, &["--root", "/workspace/src"]);
550            }
551            other => panic!("expected Stdio transport, got {other:?}"),
552        }
553    }
554
555    #[tokio::test]
556    async fn into_transport_strips_bearer_prefix_from_auth_header() -> Result<(), String> {
557        let config = McpConfig::from_json(
558            r#"{"servers":{"weather":{"type":"http","url":"http://127.0.0.1:9000/mcp","headers":{"Authorization":"Bearer secret-token"}}}}"#,
559        )
560        .map_err(|e| e.to_string())?;
561
562        let servers = config.into_servers(&HashMap::new(), &Vars::new()).await.map_err(|e| e.to_string())?;
563        let McpTransport::Http(config) = &servers[0].transport else {
564            return Err(format!("expected Http transport, got {:?}", servers[0].transport));
565        };
566
567        assert_eq!(config.transport.auth_header.as_deref(), Some("secret-token"));
568        Ok(())
569    }
570
571    #[tokio::test]
572    async fn into_transport_keeps_non_bearer_auth_header_verbatim() -> Result<(), String> {
573        let config = McpConfig::from_json(
574            r#"{"servers":{"weather":{"type":"http","url":"http://127.0.0.1:9000/mcp","headers":{"Authorization":"Basic dXNlcjpwYXNz"}}}}"#,
575        )
576        .map_err(|e| e.to_string())?;
577        let servers = config.into_servers(&HashMap::new(), &Vars::new()).await.map_err(|e| e.to_string())?;
578
579        let McpTransport::Http(config) = &servers[0].transport else {
580            return Err(format!("expected Http transport, got {:?}", servers[0].transport));
581        };
582        assert_eq!(config.transport.auth_header.as_deref(), Some("Basic dXNlcjpwYXNz"));
583        Ok(())
584    }
585
586    #[tokio::test]
587    async fn into_transport_expands_vars_in_auth_header() -> Result<(), String> {
588        let config = McpConfig::from_json(
589            r#"{"servers":{"weather":{"type":"http","url":"http://127.0.0.1:9000/mcp","headers":{"Authorization":"Bearer ${TOKEN}"}}}}"#,
590        )
591        .map_err(|e| e.to_string())?;
592        let vars = Vars::new().with("TOKEN", "expanded-token");
593        let servers = config.into_servers(&HashMap::new(), &vars).await.map_err(|e| e.to_string())?;
594
595        let McpTransport::Http(config) = &servers[0].transport else {
596            return Err(format!("expected Http transport, got {:?}", servers[0].transport));
597        };
598        assert_eq!(config.transport.auth_header.as_deref(), Some("expanded-token"));
599        Ok(())
600    }
601}