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                    #[cfg(not(feature = "http-handlers"))]
134                    {
135                        // These are bound by the pattern and only read by the
136                        // feature-enabled arm below.
137                        let _ = (endpoint, method, headers, auth, timeout_ms);
138                        return Err(Error::feature_disabled(
139                            "http-handlers",
140                            &format!("tool `{}` (type: http)", name),
141                        ));
142                    }
143
144                    #[cfg(feature = "http-handlers")]
145                    {
146                        use crate::handlers::http::{
147                            AuthConfig as HttpAuthConfig, HttpHandler,
148                            HttpMethod as HandlerHttpMethod,
149                        };
150
151                        let handler_method = match method {
152                            pforge_config::HttpMethod::Get => HandlerHttpMethod::Get,
153                            pforge_config::HttpMethod::Post => HandlerHttpMethod::Post,
154                            pforge_config::HttpMethod::Put => HandlerHttpMethod::Put,
155                            pforge_config::HttpMethod::Delete => HandlerHttpMethod::Delete,
156                            pforge_config::HttpMethod::Patch => HandlerHttpMethod::Patch,
157                        };
158
159                        let handler_auth = auth.as_ref().map(|a| match a {
160                            pforge_config::AuthConfig::Bearer { token } => HttpAuthConfig::Bearer {
161                                token: token.clone(),
162                            },
163                            pforge_config::AuthConfig::Basic { username, password } => {
164                                HttpAuthConfig::Basic {
165                                    username: username.clone(),
166                                    password: password.clone(),
167                                }
168                            }
169                            pforge_config::AuthConfig::ApiKey { key, header } => {
170                                HttpAuthConfig::ApiKey {
171                                    key: key.clone(),
172                                    header: header.clone(),
173                                }
174                            }
175                        });
176
177                        let handler = HttpHandler::new(
178                            endpoint.clone(),
179                            handler_method,
180                            headers.clone(),
181                            handler_auth,
182                            *timeout_ms,
183                        );
184                        registry.register(name, handler);
185                        eprintln!("Registered HTTP handler: {}", name);
186                    }
187                }
188                pforge_config::ToolDef::Pipeline { name, steps, .. } => {
189                    use crate::handlers::pipeline::PipelineHandlerAdapter;
190                    let handler =
191                        PipelineHandlerAdapter::from_config_steps(steps, self.registry.clone());
192                    registry.register(name, handler);
193                    eprintln!("Registered Pipeline handler: {}", name);
194                }
195            }
196        }
197
198        Ok(())
199    }
200
201    /// Run the MCP server using pmcp protocol implementation
202    pub async fn run(&self) -> Result<()> {
203        eprintln!(
204            "Starting MCP server: {} v{}",
205            self.config.forge.name, self.config.forge.version
206        );
207        eprintln!("Transport: {:?}", self.config.forge.transport);
208        eprintln!("Tools registered: {}", self.config.tools.len());
209
210        // Register handlers in pforge registry
211        self.register_handlers().await?;
212
213        // Every name we are about to ADVERTISE must be DISPATCHABLE.
214        //
215        // Before this check, `pforge new` + `pforge serve` — the exact two
216        // commands `pforge new` prints under "Next steps" — produced a server
217        // where tools/list returned `hello` and tools/call answered
218        // "Tool not found: hello". The scaffold declares a `type: native`
219        // handler that lives in the generated project's own src/, and this is
220        // the GENERIC pforge binary, which has no knowledge of that Rust. The
221        // Native arm of register_handlers registers nothing and prints a note
222        // to stderr; the builder below then added an adapter for it anyway.
223        //
224        // A stderr note is not a contract. An MCP client — usually an LLM —
225        // reads tools/list and believes it, so an advertised-but-uncallable
226        // tool surfaces at use time as a confusing protocol error rather than
227        // as a missing capability. Refusing to start says the true thing once,
228        // to the operator, at the moment they can act on it.
229        {
230            let registry = self.registry.read().await;
231            let undispatchable: Vec<&str> = self
232                .config
233                .tools
234                .iter()
235                .map(|t| match t {
236                    pforge_config::ToolDef::Native { name, .. }
237                    | pforge_config::ToolDef::Cli { name, .. }
238                    | pforge_config::ToolDef::Http { name, .. }
239                    | pforge_config::ToolDef::Pipeline { name, .. } => name.as_str(),
240                })
241                .filter(|name| !registry.has_handler(name))
242                .collect();
243
244            if !undispatchable.is_empty() {
245                return Err(Error::Handler(format!(
246                    "refusing to start: {} tool(s) are declared in the config but have no \
247                     registered handler, so they would be advertised by tools/list and fail \
248                     tools/call: {}.\n\
249                     \n\
250                     A `type: native` tool's handler is compiled INTO a server binary. The \
251                     generic `pforge serve` cannot dispatch it — build the project's own \
252                     binary (`pforge build`) and run that instead.",
253                    undispatchable.len(),
254                    undispatchable.join(", ")
255                )));
256            }
257        }
258
259        // Build pmcp server with tool adapters
260        let mut builder = pmcp::Server::builder()
261            .name(&self.config.forge.name)
262            .version(&self.config.forge.version);
263
264        // Add tool adapters for each registered tool
265        for tool in &self.config.tools {
266            let (tool_name, description) = match tool {
267                pforge_config::ToolDef::Native {
268                    name, description, ..
269                } => (name.clone(), Some(description.clone())),
270                pforge_config::ToolDef::Cli {
271                    name, description, ..
272                } => (name.clone(), Some(description.clone())),
273                pforge_config::ToolDef::Http {
274                    name, description, ..
275                } => (name.clone(), Some(description.clone())),
276                pforge_config::ToolDef::Pipeline {
277                    name, description, ..
278                } => (name.clone(), Some(description.clone())),
279            };
280
281            let adapter = PforgeToolAdapter {
282                registry: self.registry.clone(),
283                tool_name: tool_name.clone(),
284                description,
285            };
286            builder = builder.tool(&tool_name, adapter);
287        }
288
289        let server = builder
290            .build()
291            .map_err(|e| Error::Handler(format!("Failed to build MCP server: {}", e)))?;
292
293        eprintln!("MCP server ready, starting protocol loop...");
294
295        // Run the server with appropriate transport
296        match self.config.forge.transport {
297            pforge_config::TransportType::Stdio => {
298                server
299                    .run_stdio()
300                    .await
301                    .map_err(|e| Error::Handler(format!("MCP server error: {}", e)))?;
302            }
303            pforge_config::TransportType::Sse => {
304                #[cfg(not(feature = "sse"))]
305                return Err(Error::feature_disabled("sse", "transport `sse`"));
306
307                #[cfg(feature = "sse")]
308                {
309                    // See transport.rs for why this is not migrated to
310                    // StreamableHttpTransport: it is a wire-protocol change, not a
311                    // rename.
312                    #[allow(deprecated)]
313                    use pmcp::shared::{OptimizedSseConfig, OptimizedSseTransport};
314                    use std::time::Duration;
315
316                    let config = OptimizedSseConfig {
317                        url: "http://localhost:8080/sse".to_string(),
318                        connection_timeout: Duration::from_secs(30),
319                        keepalive_interval: Duration::from_secs(15),
320                        max_reconnects: 5,
321                        reconnect_delay: Duration::from_secs(1),
322                        buffer_size: 100,
323                        flush_interval: Duration::from_millis(100),
324                        enable_pooling: true,
325                        max_connections: 10,
326                        enable_compression: false,
327                    };
328                    #[allow(deprecated)]
329                    // see transport.rs: SSE -> StreamableHttp is a wire-protocol change
330                    let transport = OptimizedSseTransport::new(config);
331                    server
332                        .run(transport)
333                        .await
334                        .map_err(|e| Error::Handler(format!("MCP server error: {}", e)))?;
335                }
336            }
337            pforge_config::TransportType::WebSocket => {
338                #[cfg(not(feature = "websocket"))]
339                return Err(Error::feature_disabled(
340                    "websocket",
341                    "transport `websocket`",
342                ));
343
344                #[cfg(feature = "websocket")]
345                {
346                    #[cfg(feature = "websocket")]
347                    use pmcp::shared::{WebSocketConfig, WebSocketTransport};
348                    use std::time::Duration;
349
350                    let url = "ws://localhost:8080/ws"
351                        .parse()
352                        .map_err(|e| Error::Handler(format!("Invalid WebSocket URL: {}", e)))?;
353                    let config = WebSocketConfig {
354                        url,
355                        auto_reconnect: true,
356                        reconnect_delay: Duration::from_secs(1),
357                        max_reconnect_delay: Duration::from_secs(30),
358                        max_reconnect_attempts: Some(5),
359                        ping_interval: Some(Duration::from_secs(30)),
360                        request_timeout: Duration::from_secs(10),
361                    };
362                    let transport = WebSocketTransport::new(config);
363                    server
364                        .run(transport)
365                        .await
366                        .map_err(|e| Error::Handler(format!("MCP server error: {}", e)))?;
367                }
368            }
369        }
370
371        eprintln!("\nShutting down...");
372        Ok(())
373    }
374
375    /// Get the handler registry (for testing)
376    pub fn registry(&self) -> Arc<RwLock<HandlerRegistry>> {
377        self.registry.clone()
378    }
379}
380
381#[cfg(test)]
382mod tests {
383    use super::*;
384    use pforge_config::{ForgeMetadata, ParamSchema, ToolDef, TransportType};
385
386    fn create_test_config() -> ForgeConfig {
387        ForgeConfig {
388            forge: ForgeMetadata {
389                name: "test-server".to_string(),
390                version: "0.1.0".to_string(),
391                transport: TransportType::Stdio,
392                optimization: pforge_config::OptimizationLevel::Debug,
393            },
394            tools: vec![],
395            resources: vec![],
396            prompts: vec![],
397            state: None,
398        }
399    }
400
401    #[test]
402    fn test_server_new() {
403        let config = create_test_config();
404        let server = McpServer::new(config);
405
406        assert_eq!(server.config.forge.name, "test-server");
407        assert_eq!(server.config.forge.version, "0.1.0");
408    }
409
410    #[tokio::test]
411    async fn test_register_handlers_cli() {
412        let mut config = create_test_config();
413        config.tools.push(ToolDef::Cli {
414            name: "test_cli".to_string(),
415            description: "Test CLI handler".to_string(),
416            command: "echo".to_string(),
417            args: vec!["hello".to_string()],
418            cwd: None,
419            env: rustc_hash::FxHashMap::default(),
420            stream: false,
421            timeout_ms: None,
422        });
423
424        let server = McpServer::new(config);
425        let result = server.register_handlers().await;
426
427        assert!(result.is_ok());
428    }
429
430    fn config_with_one_http_tool() -> pforge_config::ForgeConfig {
431        let mut config = create_test_config();
432        config.tools.push(ToolDef::Http {
433            name: "test_http".to_string(),
434            description: "Test HTTP handler".to_string(),
435            endpoint: "https://api.example.com".to_string(),
436            method: pforge_config::HttpMethod::Get,
437            headers: rustc_hash::FxHashMap::default(),
438            auth: None,
439            timeout_ms: None,
440        });
441        config
442    }
443
444    #[cfg(feature = "http-handlers")]
445    #[tokio::test]
446    async fn test_register_handlers_http() {
447        let server = McpServer::new(config_with_one_http_tool());
448        let result = server.register_handlers().await;
449
450        assert!(result.is_ok());
451    }
452
453    // The counter-case matters as much as the case above: an http tool in the
454    // config must be REJECTED, loudly, by a binary built without http-handlers.
455    // Registering nothing and returning Ok would leave a server that starts
456    // clean and then answers "tool not found" for a tool the operator can see
457    // in their own forge.yaml.
458    #[cfg(not(feature = "http-handlers"))]
459    #[tokio::test]
460    async fn test_register_http_tool_without_feature_is_a_loud_error() {
461        let server = McpServer::new(config_with_one_http_tool());
462        let msg = server
463            .register_handlers()
464            .await
465            .expect_err("an http tool must not silently vanish")
466            .to_string();
467
468        assert!(
469            msg.contains("test_http"),
470            "the error must name the offending tool, got: {msg}"
471        );
472        assert!(
473            msg.contains("http-handlers") && msg.contains("--features"),
474            "the error must name the feature and how to enable it, got: {msg}"
475        );
476    }
477
478    #[tokio::test]
479    async fn test_register_handlers_native() {
480        let mut config = create_test_config();
481        config.tools.push(ToolDef::Native {
482            name: "test_native".to_string(),
483            description: "Test native handler".to_string(),
484            handler: pforge_config::HandlerRef {
485                path: "handlers::test::TestHandler".to_string(),
486                inline: None,
487            },
488            params: ParamSchema {
489                fields: rustc_hash::FxHashMap::default(),
490            },
491            timeout_ms: Some(5000),
492        });
493
494        let server = McpServer::new(config);
495        let result = server.register_handlers().await;
496
497        // Should succeed (native handlers registered by generated code)
498        assert!(result.is_ok());
499    }
500
501    #[tokio::test]
502    async fn test_registry_access() {
503        let config = create_test_config();
504        let server = McpServer::new(config);
505
506        let registry = server.registry();
507        let _lock = registry.read().await;
508
509        // Registry is accessible (test passes if no panic)
510    }
511
512    #[tokio::test]
513    async fn test_registry_returns_actual_registry() {
514        // This test catches mutation: registry() returning a new empty registry
515        let mut config = create_test_config();
516        config.tools.push(ToolDef::Cli {
517            name: "test_cli".to_string(),
518            description: "Test CLI".to_string(),
519            command: "echo".to_string(),
520            args: vec!["test".to_string()],
521            cwd: None,
522            env: rustc_hash::FxHashMap::default(),
523            stream: false,
524            timeout_ms: None,
525        });
526
527        let server = McpServer::new(config);
528        server.register_handlers().await.unwrap();
529
530        // Get registry and verify the handler is registered
531        let registry = server.registry();
532        let reg = registry.read().await;
533
534        // The CLI handler should be registered - verify via len
535        assert_eq!(reg.len(), 1, "Registry should contain registered handler");
536    }
537
538    #[tokio::test]
539    async fn test_register_handlers_pipeline() {
540        let mut config = create_test_config();
541        config.tools.push(ToolDef::Pipeline {
542            name: "test_pipeline".to_string(),
543            description: "Test pipeline handler".to_string(),
544            steps: vec![],
545        });
546
547        let server = McpServer::new(config);
548        let result = server.register_handlers().await;
549        assert!(result.is_ok());
550
551        // Verify the pipeline is actually registered
552        let registry = server.registry();
553        let reg = registry.read().await;
554        assert_eq!(reg.len(), 1, "Pipeline handler should be registered");
555    }
556
557    #[tokio::test]
558    async fn test_server_with_multiple_tools() {
559        let mut config = create_test_config();
560
561        config.tools.push(ToolDef::Cli {
562            name: "cli1".to_string(),
563            description: "CLI 1".to_string(),
564            command: "echo".to_string(),
565            args: vec![],
566            cwd: None,
567            env: rustc_hash::FxHashMap::default(),
568            stream: false,
569            timeout_ms: None,
570        });
571
572        // Only add the http tool where it can actually be served; the point of
573        // this test is that MIXED tool kinds register together, not http itself.
574        #[cfg(feature = "http-handlers")]
575        config.tools.push(ToolDef::Http {
576            name: "http1".to_string(),
577            description: "HTTP 1".to_string(),
578            endpoint: "https://example.com".to_string(),
579            method: pforge_config::HttpMethod::Get,
580            headers: rustc_hash::FxHashMap::default(),
581            auth: None,
582            timeout_ms: None,
583        });
584
585        #[cfg(not(feature = "http-handlers"))]
586        config.tools.push(ToolDef::Pipeline {
587            name: "pipeline1".to_string(),
588            description: "Pipeline 1".to_string(),
589            steps: vec![],
590        });
591
592        let server = McpServer::new(config);
593        assert_eq!(server.config.tools.len(), 2);
594
595        let result = server.register_handlers().await;
596        assert!(result.is_ok());
597    }
598}