prism_mcp_rs/server/
builder.rs1use std::collections::HashMap;
7
8use crate::core::{Prompt, Resource, Tool};
9use crate::protocol::types::{
10 CompletionsCapability, LoggingCapability, PromptsCapability, ResourceTemplate,
11 ResourcesCapability, SamplingCapability, ServerCapabilities, ToolsCapability,
12};
13use crate::protocol::ProtocolMode;
14use crate::server::{McpServer, ServerConfig};
15
16pub struct ServerBuilder {
32 name: Option<String>,
33 version: Option<String>,
34 capabilities: ServerCapabilities,
35 config: ServerConfig,
36 resources: HashMap<String, Resource>,
37 tools: HashMap<String, Tool>,
38 prompts: HashMap<String, Prompt>,
39 resource_templates: HashMap<String, ResourceTemplate>,
40 protocol_mode: ProtocolMode,
41}
42
43impl ServerBuilder {
44 pub fn new() -> Self {
46 Self {
47 name: None,
48 version: None,
49 capabilities: ServerCapabilities::default(),
50 config: ServerConfig::default(),
51 resources: HashMap::new(),
52 tools: HashMap::new(),
53 prompts: HashMap::new(),
54 resource_templates: HashMap::new(),
55 protocol_mode: ProtocolMode::Auto,
56 }
57 }
58
59 pub fn name<S: Into<String>>(mut self, name: S) -> Self {
61 self.name = Some(name.into());
62 self
63 }
64
65 pub fn version<S: Into<String>>(mut self, version: S) -> Self {
67 self.version = Some(version.into());
68 self
69 }
70
71 pub fn capabilities(mut self, capabilities: ServerCapabilities) -> Self {
73 self.capabilities = capabilities;
74 self
75 }
76
77 pub fn with_prompts(mut self) -> Self {
79 self.capabilities.prompts = Some(PromptsCapability {
80 list_changed: Some(true),
81 });
82 self
83 }
84
85 pub fn with_resources(mut self) -> Self {
87 self.capabilities.resources = Some(ResourcesCapability {
88 subscribe: Some(true),
89 list_changed: Some(true),
90 });
91 self
92 }
93
94 pub fn with_tools(mut self) -> Self {
96 self.capabilities.tools = Some(ToolsCapability {
97 list_changed: Some(true),
98 });
99 self
100 }
101
102 pub fn with_sampling(mut self) -> Self {
104 self.capabilities.sampling = Some(SamplingCapability::default());
105 self
106 }
107
108 pub fn with_logging(mut self) -> Self {
110 self.capabilities.logging = Some(LoggingCapability::default());
111 self
112 }
113
114 pub fn with_completions(mut self) -> Self {
116 self.capabilities.completions = Some(CompletionsCapability::default());
117 self
118 }
119
120 pub fn with_roots(self) -> Self {
122 self
125 }
126
127 pub fn with_experimental<K: Into<String>>(mut self, key: K, value: serde_json::Value) -> Self {
129 if self.capabilities.experimental.is_none() {
130 self.capabilities.experimental = Some(HashMap::new());
131 }
132 if let Some(ref mut experimental) = self.capabilities.experimental {
133 experimental.insert(key.into(), value);
134 }
135 self
136 }
137
138 pub fn config(mut self, config: ServerConfig) -> Self {
140 self.config = config;
141 self
142 }
143
144 pub fn protocol_mode(mut self, mode: ProtocolMode) -> Self {
146 self.protocol_mode = mode;
147 self
148 }
149
150 pub fn max_concurrent_requests(mut self, max: usize) -> Self {
152 self.config.max_concurrent_requests = max;
153 self
154 }
155
156 pub fn request_timeout_ms(mut self, timeout: u64) -> Self {
158 self.config.request_timeout_ms = timeout;
159 self
160 }
161
162 pub fn validate_requests(mut self, validate: bool) -> Self {
164 self.config.validate_requests = validate;
165 self
166 }
167
168 pub fn enable_logging(mut self, enable: bool) -> Self {
170 self.config.enable_logging = enable;
171 self
172 }
173
174 pub fn add_resource(mut self, resource: Resource) -> Self {
176 self.resources.insert(resource.info.uri.clone(), resource);
177 self
178 }
179
180 pub fn add_tool(mut self, tool: Tool) -> Self {
182 self.tools.insert(tool.info.name.clone(), tool);
183 self
184 }
185
186 pub fn add_prompt(mut self, prompt: Prompt) -> Self {
188 self.prompts.insert(prompt.info.name.clone(), prompt);
189 self
190 }
191
192 pub fn add_resource_template(mut self, template: ResourceTemplate) -> Self {
194 self.resource_templates
195 .insert(template.uri_template.clone(), template);
196 self
197 }
198
199 pub fn build(self) -> McpServer {
205 let name = self.name.expect("Server name is required");
206 let version = self.version.expect("Server version is required");
207
208 let mut server = McpServer::new(name, version);
209 server.set_capabilities(self.capabilities);
210 server.set_config(self.config);
211 server.set_protocol_mode(self.protocol_mode);
212
213 server.set_initial_resources(self.resources);
217 server.set_initial_tools(self.tools);
218 server.set_initial_prompts(self.prompts);
219 server.set_initial_resource_templates(self.resource_templates);
220
221 server
222 }
223
224 pub fn try_build(self) -> Result<McpServer, ServerBuilderError> {
226 let name = self.name.ok_or(ServerBuilderError::MissingName)?;
227 let version = self.version.ok_or(ServerBuilderError::MissingVersion)?;
228
229 let mut server = McpServer::new(name, version);
230 server.set_capabilities(self.capabilities);
231 server.set_config(self.config);
232 server.set_protocol_mode(self.protocol_mode);
233
234 server.set_initial_resources(self.resources);
235 server.set_initial_tools(self.tools);
236 server.set_initial_prompts(self.prompts);
237 server.set_initial_resource_templates(self.resource_templates);
238
239 Ok(server)
240 }
241}
242
243impl Default for ServerBuilder {
244 fn default() -> Self {
245 Self::new()
246 }
247}
248
249#[derive(Debug, Clone, PartialEq)]
251pub enum ServerBuilderError {
252 MissingName,
254 MissingVersion,
256}
257
258impl std::fmt::Display for ServerBuilderError {
259 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
260 match self {
261 Self::MissingName => write!(f, "Server name is required"),
262 Self::MissingVersion => write!(f, "Server version is required"),
263 }
264 }
265}
266
267impl std::error::Error for ServerBuilderError {}
268
269#[cfg(test)]
270mod tests {
271 use super::*;
272 use serde_json::json;
273
274 #[test]
275 fn test_builder_basic() {
276 let _server = ServerBuilder::new()
277 .name("test-server")
278 .version("1.0.0")
279 .build();
280
281 }
284
285 #[test]
286 fn test_builder_with_capabilities() {
287 let _server = ServerBuilder::new()
288 .name("test-server")
289 .version("1.0.0")
290 .with_prompts()
291 .with_resources()
292 .with_tools()
293 .with_sampling()
294 .with_logging()
295 .with_completions()
296 .build();
297 }
298
299 #[test]
300 fn test_builder_with_config() {
301 let _server = ServerBuilder::new()
302 .name("test-server")
303 .version("1.0.0")
304 .max_concurrent_requests(50)
305 .request_timeout_ms(60000)
306 .validate_requests(false)
307 .enable_logging(false)
308 .build();
309 }
310
311 #[test]
312 fn test_builder_with_experimental() {
313 let _server = ServerBuilder::new()
314 .name("test-server")
315 .version("1.0.0")
316 .with_experimental("custom_feature", json!(true))
317 .with_experimental("beta_mode", json!({"enabled": true}))
318 .build();
319 }
320
321 #[test]
322 #[should_panic(expected = "Server name is required")]
323 fn test_builder_missing_name() {
324 ServerBuilder::new().version("1.0.0").build();
325 }
326
327 #[test]
328 #[should_panic(expected = "Server version is required")]
329 fn test_builder_missing_version() {
330 ServerBuilder::new().name("test-server").build();
331 }
332
333 #[test]
334 fn test_try_build_success() {
335 let result = ServerBuilder::new()
336 .name("test-server")
337 .version("1.0.0")
338 .try_build();
339
340 assert!(result.is_ok());
341 }
342
343 #[test]
344 fn test_try_build_missing_name() {
345 let result = ServerBuilder::new().version("1.0.0").try_build();
346
347 assert!(matches!(result, Err(ServerBuilderError::MissingName)));
348 }
349
350 #[test]
351 fn test_try_build_missing_version() {
352 let result = ServerBuilder::new().name("test-server").try_build();
353
354 assert!(matches!(result, Err(ServerBuilderError::MissingVersion)));
355 }
356}