Skip to main content

pforge_runtime/
server.rs

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
9/// MCP Server implementation
10pub struct McpServer {
11    config: ForgeConfig,
12    registry: Arc<RwLock<HandlerRegistry>>,
13}
14
15/// Adapter to wrap pforge handlers as pmcp ToolHandler
16struct 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        // Serialize args to bytes for pforge dispatch
30        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, &params)
36            .await
37            .map_err(|e| pmcp::Error::protocol_msg(e.to_string()))?;
38
39        // Deserialize result back to Value
40        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        // Try to get actual schema from registry (may fail if lock is held)
49        let input_schema = if let Ok(guard) = self.registry.try_read() {
50            if let Some(schema) = guard.get_input_schema(&self.tool_name) {
51                // Convert RootSchema to serde_json::Value
52                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            // Fallback if lock unavailable
66            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    /// Create a new MCP server from configuration
82    pub fn new(config: ForgeConfig) -> Self {
83        Self {
84            config,
85            registry: Arc::new(RwLock::new(HandlerRegistry::new())),
86        }
87    }
88
89    /// Register all handlers from configuration
90    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                    // Native handlers will be registered by generated code
97                    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    /// Run the MCP server using pmcp protocol implementation
187    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        // Register handlers in pforge registry
196        self.register_handlers().await?;
197
198        // Build pmcp server with tool adapters
199        let mut builder = pmcp::Server::builder()
200            .name(&self.config.forge.name)
201            .version(&self.config.forge.version);
202
203        // Add tool adapters for each registered tool
204        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        // Run the server with appropriate transport
235        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                // See transport.rs for why this is not migrated to
244                // StreamableHttpTransport: it is a wire-protocol change, not a
245                // rename.
246                #[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                // see transport.rs: SSE -> StreamableHttp is a wire-protocol change
264                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    /// Get the handler registry (for testing)
299    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        // Should succeed (native handlers registered by generated code)
392        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        // Registry is accessible (test passes if no panic)
404    }
405
406    #[tokio::test]
407    async fn test_registry_returns_actual_registry() {
408        // This test catches mutation: registry() returning a new empty registry
409        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        // Get registry and verify the handler is registered
425        let registry = server.registry();
426        let reg = registry.read().await;
427
428        // The CLI handler should be registered - verify via len
429        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        // Verify the pipeline is actually registered
446        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}