1use crate::{Error, HandlerRegistry, Result};
2use async_trait::async_trait;
3use pforge_config::ForgeConfig;
4use pmcp::server::ToolHandler;
5use serde_json::Value;
6use std::sync::Arc;
7use tokio::sync::RwLock;
8
9pub struct McpServer {
11 config: ForgeConfig,
12 registry: Arc<RwLock<HandlerRegistry>>,
13}
14
15struct PforgeToolAdapter {
17 registry: Arc<RwLock<HandlerRegistry>>,
18 tool_name: String,
19 description: Option<String>,
20}
21
22#[async_trait]
23impl ToolHandler for PforgeToolAdapter {
24 async fn handle(
25 &self,
26 args: Value,
27 _extra: pmcp::server::cancellation::RequestHandlerExtra,
28 ) -> pmcp::Result<Value> {
29 let params = serde_json::to_vec(&args)
31 .map_err(|e| pmcp::Error::protocol_msg(format!("Failed to serialize args: {}", e)))?;
32
33 let registry = self.registry.read().await;
34 let result_bytes = registry
35 .dispatch(&self.tool_name, ¶ms)
36 .await
37 .map_err(|e| pmcp::Error::protocol_msg(e.to_string()))?;
38
39 let result: Value = serde_json::from_slice(&result_bytes).map_err(|e| {
41 pmcp::Error::protocol_msg(format!("Failed to deserialize result: {}", e))
42 })?;
43
44 Ok(result)
45 }
46
47 fn metadata(&self) -> Option<pmcp::types::ToolInfo> {
48 let input_schema = if let Ok(guard) = self.registry.try_read() {
50 if let Some(schema) = guard.get_input_schema(&self.tool_name) {
51 serde_json::to_value(&schema).unwrap_or_else(|_| {
53 serde_json::json!({
54 "type": "object",
55 "properties": {}
56 })
57 })
58 } else {
59 serde_json::json!({
60 "type": "object",
61 "properties": {}
62 })
63 }
64 } else {
65 serde_json::json!({
67 "type": "object",
68 "properties": {}
69 })
70 };
71
72 Some(pmcp::types::ToolInfo::new(
73 self.tool_name.clone(),
74 self.description.clone(),
75 input_schema,
76 ))
77 }
78}
79
80impl McpServer {
81 pub fn new(config: ForgeConfig) -> Self {
83 Self {
84 config,
85 registry: Arc::new(RwLock::new(HandlerRegistry::new())),
86 }
87 }
88
89 pub async fn register_handlers(&self) -> Result<()> {
91 let mut registry = self.registry.write().await;
92
93 for tool in &self.config.tools {
94 match tool {
95 pforge_config::ToolDef::Native { name, .. } => {
96 eprintln!(
98 "Note: Native handler '{}' requires handler implementation",
99 name
100 );
101 }
102 pforge_config::ToolDef::Cli {
103 name,
104 command,
105 args,
106 cwd,
107 env,
108 stream,
109 timeout_ms,
110 ..
111 } => {
112 use crate::handlers::cli::CliHandler;
113 let handler = CliHandler::new(
114 command.clone(),
115 args.clone(),
116 cwd.clone(),
117 env.clone(),
118 *timeout_ms,
119 *stream,
120 );
121 registry.register(name, handler);
122 eprintln!("Registered CLI handler: {}", name);
123 }
124 pforge_config::ToolDef::Http {
125 name,
126 endpoint,
127 method,
128 headers,
129 auth,
130 timeout_ms,
131 ..
132 } => {
133 use crate::handlers::http::{
134 AuthConfig as HttpAuthConfig, HttpHandler, HttpMethod as HandlerHttpMethod,
135 };
136
137 let handler_method = match method {
138 pforge_config::HttpMethod::Get => HandlerHttpMethod::Get,
139 pforge_config::HttpMethod::Post => HandlerHttpMethod::Post,
140 pforge_config::HttpMethod::Put => HandlerHttpMethod::Put,
141 pforge_config::HttpMethod::Delete => HandlerHttpMethod::Delete,
142 pforge_config::HttpMethod::Patch => HandlerHttpMethod::Patch,
143 };
144
145 let handler_auth = auth.as_ref().map(|a| match a {
146 pforge_config::AuthConfig::Bearer { token } => HttpAuthConfig::Bearer {
147 token: token.clone(),
148 },
149 pforge_config::AuthConfig::Basic { username, password } => {
150 HttpAuthConfig::Basic {
151 username: username.clone(),
152 password: password.clone(),
153 }
154 }
155 pforge_config::AuthConfig::ApiKey { key, header } => {
156 HttpAuthConfig::ApiKey {
157 key: key.clone(),
158 header: header.clone(),
159 }
160 }
161 });
162
163 let handler = HttpHandler::new(
164 endpoint.clone(),
165 handler_method,
166 headers.clone(),
167 handler_auth,
168 *timeout_ms,
169 );
170 registry.register(name, handler);
171 eprintln!("Registered HTTP handler: {}", name);
172 }
173 pforge_config::ToolDef::Pipeline { name, steps, .. } => {
174 use crate::handlers::pipeline::PipelineHandlerAdapter;
175 let handler =
176 PipelineHandlerAdapter::from_config_steps(steps, self.registry.clone());
177 registry.register(name, handler);
178 eprintln!("Registered Pipeline handler: {}", name);
179 }
180 }
181 }
182
183 Ok(())
184 }
185
186 pub async fn run(&self) -> Result<()> {
188 eprintln!(
189 "Starting MCP server: {} v{}",
190 self.config.forge.name, self.config.forge.version
191 );
192 eprintln!("Transport: {:?}", self.config.forge.transport);
193 eprintln!("Tools registered: {}", self.config.tools.len());
194
195 self.register_handlers().await?;
197
198 let mut builder = pmcp::Server::builder()
200 .name(&self.config.forge.name)
201 .version(&self.config.forge.version);
202
203 for tool in &self.config.tools {
205 let (tool_name, description) = match tool {
206 pforge_config::ToolDef::Native {
207 name, description, ..
208 } => (name.clone(), Some(description.clone())),
209 pforge_config::ToolDef::Cli {
210 name, description, ..
211 } => (name.clone(), Some(description.clone())),
212 pforge_config::ToolDef::Http {
213 name, description, ..
214 } => (name.clone(), Some(description.clone())),
215 pforge_config::ToolDef::Pipeline {
216 name, description, ..
217 } => (name.clone(), Some(description.clone())),
218 };
219
220 let adapter = PforgeToolAdapter {
221 registry: self.registry.clone(),
222 tool_name: tool_name.clone(),
223 description,
224 };
225 builder = builder.tool(&tool_name, adapter);
226 }
227
228 let server = builder
229 .build()
230 .map_err(|e| Error::Handler(format!("Failed to build MCP server: {}", e)))?;
231
232 eprintln!("MCP server ready, starting protocol loop...");
233
234 match self.config.forge.transport {
236 pforge_config::TransportType::Stdio => {
237 server
238 .run_stdio()
239 .await
240 .map_err(|e| Error::Handler(format!("MCP server error: {}", e)))?;
241 }
242 pforge_config::TransportType::Sse => {
243 #[allow(deprecated)]
247 use pmcp::shared::{OptimizedSseConfig, OptimizedSseTransport};
248 use std::time::Duration;
249
250 let config = OptimizedSseConfig {
251 url: "http://localhost:8080/sse".to_string(),
252 connection_timeout: Duration::from_secs(30),
253 keepalive_interval: Duration::from_secs(15),
254 max_reconnects: 5,
255 reconnect_delay: Duration::from_secs(1),
256 buffer_size: 100,
257 flush_interval: Duration::from_millis(100),
258 enable_pooling: true,
259 max_connections: 10,
260 enable_compression: false,
261 };
262 #[allow(deprecated)]
263 let transport = OptimizedSseTransport::new(config);
265 server
266 .run(transport)
267 .await
268 .map_err(|e| Error::Handler(format!("MCP server error: {}", e)))?;
269 }
270 pforge_config::TransportType::WebSocket => {
271 use pmcp::shared::{WebSocketConfig, WebSocketTransport};
272 use std::time::Duration;
273
274 let url = "ws://localhost:8080/ws"
275 .parse()
276 .map_err(|e| Error::Handler(format!("Invalid WebSocket URL: {}", e)))?;
277 let config = WebSocketConfig {
278 url,
279 auto_reconnect: true,
280 reconnect_delay: Duration::from_secs(1),
281 max_reconnect_delay: Duration::from_secs(30),
282 max_reconnect_attempts: Some(5),
283 ping_interval: Some(Duration::from_secs(30)),
284 request_timeout: Duration::from_secs(10),
285 };
286 let transport = WebSocketTransport::new(config);
287 server
288 .run(transport)
289 .await
290 .map_err(|e| Error::Handler(format!("MCP server error: {}", e)))?;
291 }
292 }
293
294 eprintln!("\nShutting down...");
295 Ok(())
296 }
297
298 pub fn registry(&self) -> Arc<RwLock<HandlerRegistry>> {
300 self.registry.clone()
301 }
302}
303
304#[cfg(test)]
305mod tests {
306 use super::*;
307 use pforge_config::{ForgeMetadata, ParamSchema, ToolDef, TransportType};
308
309 fn create_test_config() -> ForgeConfig {
310 ForgeConfig {
311 forge: ForgeMetadata {
312 name: "test-server".to_string(),
313 version: "0.1.0".to_string(),
314 transport: TransportType::Stdio,
315 optimization: pforge_config::OptimizationLevel::Debug,
316 },
317 tools: vec![],
318 resources: vec![],
319 prompts: vec![],
320 state: None,
321 }
322 }
323
324 #[test]
325 fn test_server_new() {
326 let config = create_test_config();
327 let server = McpServer::new(config);
328
329 assert_eq!(server.config.forge.name, "test-server");
330 assert_eq!(server.config.forge.version, "0.1.0");
331 }
332
333 #[tokio::test]
334 async fn test_register_handlers_cli() {
335 let mut config = create_test_config();
336 config.tools.push(ToolDef::Cli {
337 name: "test_cli".to_string(),
338 description: "Test CLI handler".to_string(),
339 command: "echo".to_string(),
340 args: vec!["hello".to_string()],
341 cwd: None,
342 env: rustc_hash::FxHashMap::default(),
343 stream: false,
344 timeout_ms: None,
345 });
346
347 let server = McpServer::new(config);
348 let result = server.register_handlers().await;
349
350 assert!(result.is_ok());
351 }
352
353 #[tokio::test]
354 async fn test_register_handlers_http() {
355 let mut config = create_test_config();
356 config.tools.push(ToolDef::Http {
357 name: "test_http".to_string(),
358 description: "Test HTTP handler".to_string(),
359 endpoint: "https://api.example.com".to_string(),
360 method: pforge_config::HttpMethod::Get,
361 headers: rustc_hash::FxHashMap::default(),
362 auth: None,
363 timeout_ms: None,
364 });
365
366 let server = McpServer::new(config);
367 let result = server.register_handlers().await;
368
369 assert!(result.is_ok());
370 }
371
372 #[tokio::test]
373 async fn test_register_handlers_native() {
374 let mut config = create_test_config();
375 config.tools.push(ToolDef::Native {
376 name: "test_native".to_string(),
377 description: "Test native handler".to_string(),
378 handler: pforge_config::HandlerRef {
379 path: "handlers::test::TestHandler".to_string(),
380 inline: None,
381 },
382 params: ParamSchema {
383 fields: rustc_hash::FxHashMap::default(),
384 },
385 timeout_ms: Some(5000),
386 });
387
388 let server = McpServer::new(config);
389 let result = server.register_handlers().await;
390
391 assert!(result.is_ok());
393 }
394
395 #[tokio::test]
396 async fn test_registry_access() {
397 let config = create_test_config();
398 let server = McpServer::new(config);
399
400 let registry = server.registry();
401 let _lock = registry.read().await;
402
403 }
405
406 #[tokio::test]
407 async fn test_registry_returns_actual_registry() {
408 let mut config = create_test_config();
410 config.tools.push(ToolDef::Cli {
411 name: "test_cli".to_string(),
412 description: "Test CLI".to_string(),
413 command: "echo".to_string(),
414 args: vec!["test".to_string()],
415 cwd: None,
416 env: rustc_hash::FxHashMap::default(),
417 stream: false,
418 timeout_ms: None,
419 });
420
421 let server = McpServer::new(config);
422 server.register_handlers().await.unwrap();
423
424 let registry = server.registry();
426 let reg = registry.read().await;
427
428 assert_eq!(reg.len(), 1, "Registry should contain registered handler");
430 }
431
432 #[tokio::test]
433 async fn test_register_handlers_pipeline() {
434 let mut config = create_test_config();
435 config.tools.push(ToolDef::Pipeline {
436 name: "test_pipeline".to_string(),
437 description: "Test pipeline handler".to_string(),
438 steps: vec![],
439 });
440
441 let server = McpServer::new(config);
442 let result = server.register_handlers().await;
443 assert!(result.is_ok());
444
445 let registry = server.registry();
447 let reg = registry.read().await;
448 assert_eq!(reg.len(), 1, "Pipeline handler should be registered");
449 }
450
451 #[tokio::test]
452 async fn test_server_with_multiple_tools() {
453 let mut config = create_test_config();
454
455 config.tools.push(ToolDef::Cli {
456 name: "cli1".to_string(),
457 description: "CLI 1".to_string(),
458 command: "echo".to_string(),
459 args: vec![],
460 cwd: None,
461 env: rustc_hash::FxHashMap::default(),
462 stream: false,
463 timeout_ms: None,
464 });
465
466 config.tools.push(ToolDef::Http {
467 name: "http1".to_string(),
468 description: "HTTP 1".to_string(),
469 endpoint: "https://example.com".to_string(),
470 method: pforge_config::HttpMethod::Get,
471 headers: rustc_hash::FxHashMap::default(),
472 auth: None,
473 timeout_ms: None,
474 });
475
476 let server = McpServer::new(config);
477 assert_eq!(server.config.tools.len(), 2);
478
479 let result = server.register_handlers().await;
480 assert!(result.is_ok());
481 }
482}