1use std::path::Path;
2
3use crate::error::{Result, TuffError};
4
5pub fn validate_config(mcp_config_path: &Path) -> Result<()> {
6 read_config(mcp_config_path).map(|_| ())
7}
8
9pub fn register_tool(
10 mcp_config_path: &Path,
11 tool_id: &str,
12 command: &str,
13 args: &[String],
14) -> Result<()> {
15 register_server(
16 mcp_config_path,
17 tool_id,
18 serde_json::json!({"command": command, "args": args}),
19 true,
20 )
21}
22
23pub fn register_server(
33 mcp_config_path: &Path,
34 server_id: &str,
35 entry: serde_json::Value,
36 allow_overwrite: bool,
37) -> Result<()> {
38 let mut config = read_config(mcp_config_path)?;
39 let config_object = config
40 .as_object_mut()
41 .expect("read_config returns a JSON object");
42 let servers = config_object
43 .entry("mcpServers")
44 .or_insert_with(|| serde_json::json!({}))
45 .as_object_mut()
46 .expect("read_config validates the mcpServers object");
47
48 if !allow_overwrite && servers.contains_key(server_id) {
49 return Err(TuffError::new(format!(
50 "refusing to overwrite untracked MCP server '{}' in {}; remove it by hand or \
51 choose a different capability id",
52 server_id,
53 mcp_config_path.display()
54 )));
55 }
56
57 servers.insert(server_id.to_string(), entry);
58
59 if let Some(parent) = mcp_config_path.parent() {
60 std::fs::create_dir_all(parent)?;
61 }
62 write_config(mcp_config_path, &config)
63}
64
65pub fn has_server(mcp_config_path: &Path, server_id: &str) -> Result<bool> {
69 if !mcp_config_path.exists() {
70 return Ok(false);
71 }
72 let config = read_config(mcp_config_path)?;
73 Ok(config
74 .get("mcpServers")
75 .and_then(serde_json::Value::as_object)
76 .is_some_and(|servers| servers.contains_key(server_id)))
77}
78
79pub fn remove_tool(mcp_config_path: &Path, tool_id: &str) -> Result<()> {
80 if !mcp_config_path.exists() {
81 return Ok(());
82 }
83
84 let mut config = read_config(mcp_config_path)?;
85 let Some(servers) = config
86 .as_object_mut()
87 .and_then(|object| object.get_mut("mcpServers"))
88 .and_then(serde_json::Value::as_object_mut)
89 else {
90 return Ok(());
91 };
92 if servers.remove(tool_id).is_none() {
93 return Ok(());
94 }
95
96 write_config(mcp_config_path, &config)
97}
98
99fn read_config(mcp_config_path: &Path) -> Result<serde_json::Value> {
100 let config = if mcp_config_path.exists() {
101 let raw = std::fs::read_to_string(mcp_config_path)?;
102 if raw.trim().is_empty() {
103 serde_json::json!({})
104 } else {
105 serde_json::from_str(&raw).map_err(|error| {
106 TuffError::new(format!(
107 "invalid MCP config at {}: {error}",
108 mcp_config_path.display()
109 ))
110 })?
111 }
112 } else {
113 serde_json::json!({})
114 };
115
116 let object = config.as_object().ok_or_else(|| {
117 TuffError::new(format!(
118 "invalid MCP config at {}: root must be a JSON object",
119 mcp_config_path.display()
120 ))
121 })?;
122 if object
123 .get("mcpServers")
124 .is_some_and(|servers| !servers.is_object())
125 {
126 return Err(TuffError::new(format!(
127 "invalid MCP config at {}: field 'mcpServers' must be a JSON object",
128 mcp_config_path.display()
129 )));
130 }
131
132 Ok(config)
133}
134
135fn write_config(mcp_config_path: &Path, config: &serde_json::Value) -> Result<()> {
136 std::fs::write(
137 mcp_config_path,
138 serde_json::to_string_pretty(config)? + "\n",
139 )?;
140 Ok(())
141}
142
143#[cfg(test)]
144mod tests {
145 use super::*;
146
147 #[test]
148 fn register_rejects_malformed_json_without_changing_it() {
149 let temp = tempfile::tempdir().expect("tempdir");
150 let path = temp.path().join("mcp.json");
151 let original = b"{ not-json\n";
152 std::fs::write(&path, original).expect("write config");
153
154 let error = register_tool(&path, "demo", "python", &[]).expect_err("invalid config");
155
156 assert!(error.to_string().contains("invalid MCP config"));
157 assert_eq!(std::fs::read(&path).expect("read config"), original);
158 }
159
160 #[test]
161 fn remove_rejects_malformed_json_without_changing_it() {
162 let temp = tempfile::tempdir().expect("tempdir");
163 let path = temp.path().join("mcp.json");
164 let original = b"[invalid";
165 std::fs::write(&path, original).expect("write config");
166
167 let error = remove_tool(&path, "demo").expect_err("invalid config");
168
169 assert!(error.to_string().contains("invalid MCP config"));
170 assert_eq!(std::fs::read(&path).expect("read config"), original);
171 }
172
173 #[test]
174 fn register_preserves_unrelated_fields() {
175 let temp = tempfile::tempdir().expect("tempdir");
176 let path = temp.path().join("mcp.json");
177 std::fs::write(
178 &path,
179 r#"{"custom":{"enabled":true},"mcpServers":{"existing":{"command":"node"}}}"#,
180 )
181 .expect("write config");
182
183 register_tool(&path, "demo", "python", &["server.py".to_string()]).expect("register tool");
184
185 let config: serde_json::Value =
186 serde_json::from_slice(&std::fs::read(path).expect("read config"))
187 .expect("parse config");
188 assert_eq!(config["custom"]["enabled"], true);
189 assert_eq!(config["mcpServers"]["existing"]["command"], "node");
190 assert_eq!(config["mcpServers"]["demo"]["command"], "python");
191 }
192
193 #[test]
194 fn validation_rejects_non_object_mcp_servers() {
195 let temp = tempfile::tempdir().expect("tempdir");
196 let path = temp.path().join("mcp.json");
197 std::fs::write(&path, r#"{"mcpServers":[]}"#).expect("write config");
198
199 let error = validate_config(&path).expect_err("invalid mcpServers");
200
201 assert!(
202 error
203 .to_string()
204 .contains("'mcpServers' must be a JSON object")
205 );
206 }
207
208 #[test]
209 fn register_server_writes_entry_and_preserves_neighbours() {
210 let tmp = tempfile::TempDir::new().unwrap();
211 let path = tmp.path().join("mcp.json");
212 std::fs::write(
213 &path,
214 "{\"custom\":true,\"mcpServers\":{\"other\":{\"command\":\"x\"}}}",
215 )
216 .unwrap();
217
218 register_server(
219 &path,
220 "github",
221 serde_json::json!({"command": "npx", "args": ["-y", "srv"], "env": {"T": "${T}"}}),
222 false,
223 )
224 .unwrap();
225
226 let config: serde_json::Value =
227 serde_json::from_str(&std::fs::read_to_string(&path).unwrap()).unwrap();
228 assert_eq!(config["custom"], true);
229 assert_eq!(config["mcpServers"]["other"]["command"], "x");
230 assert_eq!(config["mcpServers"]["github"]["env"]["T"], "${T}");
231 assert!(has_server(&path, "github").unwrap());
232 assert!(!has_server(&path, "missing").unwrap());
233 }
234
235 #[test]
236 fn register_server_refuses_untracked_collision_unless_overwrite_allowed() {
237 let tmp = tempfile::TempDir::new().unwrap();
238 let path = tmp.path().join("mcp.json");
239 let original = "{\"mcpServers\":{\"github\":{\"command\":\"hand\"}}}";
240 std::fs::write(&path, original).unwrap();
241
242 let error = register_server(
243 &path,
244 "github",
245 serde_json::json!({"command": "npx"}),
246 false,
247 )
248 .unwrap_err()
249 .to_string();
250 assert!(
251 error.contains("refusing to overwrite untracked MCP server"),
252 "{error}"
253 );
254 assert_eq!(std::fs::read_to_string(&path).unwrap(), original);
255
256 register_server(&path, "github", serde_json::json!({"command": "npx"}), true).unwrap();
257 let config: serde_json::Value =
258 serde_json::from_str(&std::fs::read_to_string(&path).unwrap()).unwrap();
259 assert_eq!(config["mcpServers"]["github"]["command"], "npx");
260 }
261}