mcp_attr/
server.rs

1//! Module for implementing MCP server
2
3use std::{future::Future, sync::Arc};
4
5use jsoncall::{
6    ErrorCode, Handler, Hook, NotificationContext, Params, RequestContextAs, RequestId, Response,
7    Result, Session, SessionContext, SessionOptions, SessionResult, bail_public,
8};
9use serde::{Serialize, de::DeserializeOwned};
10use serde_json::Map;
11
12use crate::{
13    PROTOCOL_VERSION,
14    common::McpCancellationHook,
15    schema::{
16        CallToolRequestParams, CallToolResult, CancelledNotificationParams, ClientCapabilities,
17        CompleteRequestParams, CompleteResult, CreateMessageRequestParams, CreateMessageResult,
18        GetPromptRequestParams, GetPromptResult, Implementation, InitializeRequestParams,
19        InitializeResult, InitializedNotificationParams, ListPromptsRequestParams,
20        ListPromptsResult, ListResourceTemplatesRequestParams, ListResourceTemplatesResult,
21        ListResourcesRequestParams, ListResourcesResult, ListRootsRequestParams, ListRootsResult,
22        ListToolsRequestParams, ListToolsResult, PingRequestParams, ProgressNotificationParams,
23        ReadResourceRequestParams, ReadResourceResult, Root, ServerCapabilities,
24        ServerCapabilitiesPrompts, ServerCapabilitiesResources, ServerCapabilitiesTools,
25    },
26    server::errors::{prompt_not_found, tool_not_found},
27    utils::Empty,
28};
29
30pub mod errors;
31mod mcp_server_attr;
32
33pub use mcp_server_attr::mcp_server;
34
35struct McpServerHandler {
36    server: Arc<dyn DynMcpServer>,
37    initialize: Option<Arc<InitializeRequestParams>>,
38    is_initizlized: bool,
39}
40impl Handler for McpServerHandler {
41    fn hook(&self) -> Arc<dyn Hook> {
42        Arc::new(McpCancellationHook)
43    }
44    fn request(
45        &mut self,
46        method: &str,
47        params: Params,
48        cx: jsoncall::RequestContext,
49    ) -> Result<Response> {
50        match method {
51            "initialize" => return cx.handle(self.initialize(params.to()?)),
52            "ping" => return cx.handle(self.ping(params.to_opt()?)),
53            _ => {}
54        }
55        let (Some(initialize), true) = (&self.initialize, self.is_initizlized) else {
56            bail_public!(_, "Server not initialized");
57        };
58        let i = initialize.clone();
59        match method {
60            "prompts/list" => self.call_opt(params, cx, |s, p, cx| s.dyn_prompts_list(p, cx, i)),
61            "prompts/get" => self.call(params, cx, |s, p, cx| s.dyn_prompts_get(p, cx, i)),
62            "resources/list" => {
63                self.call_opt(params, cx, |s, p, cx| s.dyn_resources_list(p, cx, i))
64            }
65            "resources/templates/list" => self.call_opt(params, cx, |s, p, cx| {
66                s.dyn_resources_templates_list(p, cx, i)
67            }),
68            "resources/read" => self.call(params, cx, |s, p, cx| s.dyn_resources_read(p, cx, i)),
69            "tools/list" => self.call_opt(params, cx, |s, p, cx| s.dyn_tools_list(p, cx, i)),
70            "tools/call" => self.call(params, cx, |s, p, cx| s.dyn_tools_call(p, cx, i)),
71            "completion/complete" => {
72                self.call(params, cx, |s, p, cx| s.dyn_completion_complete(p, cx, i))
73            }
74            _ => cx.method_not_found(),
75        }
76    }
77    fn notification(
78        &mut self,
79        method: &str,
80        params: Params,
81        cx: NotificationContext,
82    ) -> Result<Response> {
83        match method {
84            "notifications/initialized" => cx.handle(self.initialized(params.to_opt()?)),
85            "notifications/cancelled" => self.notifications_cancelled(params.to()?, cx),
86            _ => cx.method_not_found(),
87        }
88    }
89}
90impl McpServerHandler {
91    pub fn new(server: impl McpServer) -> Self {
92        Self {
93            server: Arc::new(server),
94            initialize: None,
95            is_initizlized: false,
96        }
97    }
98}
99impl McpServerHandler {
100    fn initialize(&mut self, p: InitializeRequestParams) -> Result<InitializeResult> {
101        if p.protocol_version != PROTOCOL_VERSION {
102            bail_public!(ErrorCode::INVALID_PARAMS, "Unsupported protocol version");
103        }
104        self.initialize = Some(Arc::new(p));
105        Ok(self.server.initialize_result())
106    }
107    fn initialized(&mut self, _p: Option<InitializedNotificationParams>) -> Result<()> {
108        if self.initialize.is_none() {
109            bail_public!(
110                _,
111                "`initialize` request must be called before `initialized` notification"
112            );
113        }
114        self.is_initizlized = true;
115        Ok(())
116    }
117    fn ping(&self, _p: Option<PingRequestParams>) -> Result<Empty> {
118        Ok(Empty::default())
119    }
120    fn notifications_cancelled(
121        &self,
122        p: CancelledNotificationParams,
123        cx: NotificationContext,
124    ) -> Result<Response> {
125        cx.session().cancel_incoming_request(&p.request_id, None);
126        cx.handle(Ok(()))
127    }
128
129    // fn logging_set_level(&self, p: SetLevelRequestParams) -> Result<()> {
130    //     todo!()
131    // }
132
133    // fn resources_subscribe(&self, p: SubscribeRequestParams) -> Result<()> {
134    //     todo!()
135    // }
136
137    // fn resources_unsubscribe(&self, p: UnsubscribeRequestParams) -> Result<()> {
138    //     todo!()
139    // }
140
141    fn call<P, R>(
142        &self,
143        p: Params,
144        cx: jsoncall::RequestContext,
145        f: impl FnOnce(Arc<dyn DynMcpServer>, P, RequestContextAs<R>) -> Result<Response>,
146    ) -> Result<Response>
147    where
148        P: DeserializeOwned,
149        R: Serialize,
150    {
151        f(self.server.clone(), p.to()?, cx.to())
152    }
153    fn call_opt<P, R>(
154        &self,
155        p: Params,
156        cx: jsoncall::RequestContext,
157        f: impl FnOnce(Arc<dyn DynMcpServer>, P, RequestContextAs<R>) -> Result<Response>,
158    ) -> Result<Response>
159    where
160        P: DeserializeOwned + Default,
161        R: Serialize,
162    {
163        f(
164            self.server.clone(),
165            p.to_opt()?.unwrap_or_default(),
166            cx.to(),
167        )
168    }
169}
170
171trait DynMcpServer: Send + Sync + 'static {
172    fn initialize_result(&self) -> InitializeResult;
173
174    fn dyn_prompts_list(
175        self: Arc<Self>,
176        p: ListPromptsRequestParams,
177        cx: RequestContextAs<ListPromptsResult>,
178        initialize: Arc<InitializeRequestParams>,
179    ) -> Result<Response>;
180
181    fn dyn_prompts_get(
182        self: Arc<Self>,
183        p: GetPromptRequestParams,
184        cx: RequestContextAs<GetPromptResult>,
185        initialize: Arc<InitializeRequestParams>,
186    ) -> Result<Response>;
187
188    fn dyn_resources_list(
189        self: Arc<Self>,
190        p: ListResourcesRequestParams,
191        cx: RequestContextAs<ListResourcesResult>,
192        initialize: Arc<InitializeRequestParams>,
193    ) -> Result<Response>;
194
195    fn dyn_resources_read(
196        self: Arc<Self>,
197        p: ReadResourceRequestParams,
198        cx: RequestContextAs<ReadResourceResult>,
199        initialize: Arc<InitializeRequestParams>,
200    ) -> Result<Response>;
201
202    fn dyn_resources_templates_list(
203        self: Arc<Self>,
204        p: ListResourceTemplatesRequestParams,
205        cx: RequestContextAs<ListResourceTemplatesResult>,
206        initialize: Arc<InitializeRequestParams>,
207    ) -> Result<Response>;
208
209    fn dyn_tools_list(
210        self: Arc<Self>,
211        p: ListToolsRequestParams,
212        cx: RequestContextAs<ListToolsResult>,
213        initialize: Arc<InitializeRequestParams>,
214    ) -> Result<Response>;
215
216    fn dyn_tools_call(
217        self: Arc<Self>,
218        p: CallToolRequestParams,
219        cx: RequestContextAs<CallToolResult>,
220        initialize: Arc<InitializeRequestParams>,
221    ) -> Result<Response>;
222
223    fn dyn_completion_complete(
224        self: Arc<Self>,
225        p: CompleteRequestParams,
226        cx: RequestContextAs<CompleteResult>,
227        initialize: Arc<InitializeRequestParams>,
228    ) -> Result<Response>;
229}
230impl<T: McpServer> DynMcpServer for T {
231    fn initialize_result(&self) -> InitializeResult {
232        InitializeResult {
233            capabilities: self.capabilities(),
234            instructions: self.instructions(),
235            meta: Map::new(),
236            protocol_version: PROTOCOL_VERSION.to_string(),
237            server_info: self.server_info(),
238        }
239    }
240    fn dyn_prompts_list(
241        self: Arc<Self>,
242        p: ListPromptsRequestParams,
243        cx: RequestContextAs<ListPromptsResult>,
244        initialize: Arc<InitializeRequestParams>,
245    ) -> Result<Response> {
246        let mut mcp_cx = RequestContext::new(&cx, initialize);
247        cx.handle_async(async move { self.prompts_list(p, &mut mcp_cx).await })
248    }
249
250    fn dyn_prompts_get(
251        self: Arc<Self>,
252        p: GetPromptRequestParams,
253        cx: RequestContextAs<GetPromptResult>,
254        initialize: Arc<InitializeRequestParams>,
255    ) -> Result<Response> {
256        let mut mcp_cx = RequestContext::new(&cx, initialize);
257        cx.handle_async(async move { self.prompts_get(p, &mut mcp_cx).await })
258    }
259
260    fn dyn_resources_list(
261        self: Arc<Self>,
262        p: ListResourcesRequestParams,
263        cx: RequestContextAs<ListResourcesResult>,
264        initialize: Arc<InitializeRequestParams>,
265    ) -> Result<Response> {
266        let mut mcp_cx = RequestContext::new(&cx, initialize);
267        cx.handle_async(async move { self.resources_list(p, &mut mcp_cx).await })
268    }
269
270    fn dyn_resources_templates_list(
271        self: Arc<Self>,
272        p: ListResourceTemplatesRequestParams,
273        cx: RequestContextAs<ListResourceTemplatesResult>,
274        initialize: Arc<InitializeRequestParams>,
275    ) -> Result<Response> {
276        let mut mcp_cx = RequestContext::new(&cx, initialize);
277        cx.handle_async(async move { self.resources_templates_list(p, &mut mcp_cx).await })
278    }
279
280    fn dyn_resources_read(
281        self: Arc<Self>,
282        p: ReadResourceRequestParams,
283        cx: RequestContextAs<ReadResourceResult>,
284        initialize: Arc<InitializeRequestParams>,
285    ) -> Result<Response> {
286        let mut mcp_cx = RequestContext::new(&cx, initialize);
287        cx.handle_async(async move { self.resources_read(p, &mut mcp_cx).await })
288    }
289
290    fn dyn_tools_list(
291        self: Arc<Self>,
292        p: ListToolsRequestParams,
293        cx: RequestContextAs<ListToolsResult>,
294        initialize: Arc<InitializeRequestParams>,
295    ) -> Result<Response> {
296        let mut mcp_cx = RequestContext::new(&cx, initialize);
297        cx.handle_async(async move { self.tools_list(p, &mut mcp_cx).await })
298    }
299
300    fn dyn_tools_call(
301        self: Arc<Self>,
302        p: CallToolRequestParams,
303        cx: RequestContextAs<CallToolResult>,
304        initialize: Arc<InitializeRequestParams>,
305    ) -> Result<Response> {
306        let mut mcp_cx = RequestContext::new(&cx, initialize);
307        cx.handle_async(async move { self.tools_call(p, &mut mcp_cx).await })
308    }
309
310    fn dyn_completion_complete(
311        self: Arc<Self>,
312        p: CompleteRequestParams,
313        cx: RequestContextAs<CompleteResult>,
314        initialize: Arc<InitializeRequestParams>,
315    ) -> Result<Response> {
316        let mut mcp_cx = RequestContext::new(&cx, initialize);
317        cx.handle_async(async move { self.completion_complete(p, &mut mcp_cx).await })
318    }
319}
320
321/// Trait for implementing MCP server
322pub trait McpServer: Send + Sync + 'static {
323    /// Returns `server_info` used in the [`initialize`] request response
324    ///
325    /// [`initialize`]: https://spec.modelcontextprotocol.io/specification/2024-11-05/basic/lifecycle/#initialization
326    fn server_info(&self) -> Implementation {
327        Implementation::from_compile_time_env()
328    }
329
330    /// Returns `instructions` used in the [`initialize`] request response
331    ///
332    /// [`initialize`]: https://spec.modelcontextprotocol.io/specification/2024-11-05/basic/lifecycle/#initialization
333    fn instructions(&self) -> Option<String> {
334        None
335    }
336
337    /// Returns `capabilities` used in the [`initialize`] request response
338    ///
339    /// [`initialize`]: https://spec.modelcontextprotocol.io/specification/2024-11-05/basic/lifecycle/#initialization
340    fn capabilities(&self) -> ServerCapabilities {
341        ServerCapabilities {
342            prompts: Some(ServerCapabilitiesPrompts {
343                ..Default::default()
344            }),
345            resources: Some(ServerCapabilitiesResources {
346                ..Default::default()
347            }),
348            tools: Some(ServerCapabilitiesTools {
349                ..Default::default()
350            }),
351            ..Default::default()
352        }
353    }
354
355    /// Handles [`prompts/list`]
356    ///
357    /// [`prompts/list`]: https://spec.modelcontextprotocol.io/specification/2024-11-05/server/prompts/#listing-prompts
358    #[allow(unused_variables)]
359    fn prompts_list(
360        &self,
361        p: ListPromptsRequestParams,
362        cx: &mut RequestContext,
363    ) -> impl Future<Output = Result<ListPromptsResult>> + Send {
364        async { Ok(ListPromptsResult::default()) }
365    }
366
367    /// Handles [`prompts/get`]
368    ///
369    /// [`prompts/get`]: https://spec.modelcontextprotocol.io/specification/2024-11-05/server/prompts/#getting-a-prompt
370    #[allow(unused_variables)]
371    fn prompts_get(
372        &self,
373        p: GetPromptRequestParams,
374        cx: &mut RequestContext,
375    ) -> impl Future<Output = Result<GetPromptResult>> + Send {
376        async move { Err(prompt_not_found(&p.name)) }
377    }
378
379    /// Handles [`resources/list`]
380    ///
381    /// [`resources/list`]: https://spec.modelcontextprotocol.io/specification/2024-11-05/server/resources/#listing-resources
382    #[allow(unused_variables)]
383    fn resources_list(
384        &self,
385        p: ListResourcesRequestParams,
386        cx: &mut RequestContext,
387    ) -> impl Future<Output = Result<ListResourcesResult>> + Send {
388        async { Ok(ListResourcesResult::default()) }
389    }
390
391    /// Handles [`resources/templates/list`]
392    ///
393    /// [`resources/templates/list`]: https://spec.modelcontextprotocol.io/specification/2024-11-05/server/resources/#resource-templates
394    #[allow(unused_variables)]
395    fn resources_templates_list(
396        &self,
397        p: ListResourceTemplatesRequestParams,
398        cx: &mut RequestContext,
399    ) -> impl Future<Output = Result<ListResourceTemplatesResult>> + Send {
400        async { Ok(ListResourceTemplatesResult::default()) }
401    }
402
403    /// Handles [`resources/read`]
404    ///
405    /// [`resources/read`]: https://spec.modelcontextprotocol.io/specification/2024-11-05/server/resources/#reading-resources
406    #[allow(unused_variables)]
407    fn resources_read(
408        &self,
409        p: ReadResourceRequestParams,
410        cx: &mut RequestContext,
411    ) -> impl Future<Output = Result<ReadResourceResult>> + Send {
412        async move { bail_public!(ErrorCode::INVALID_PARAMS, "Resource `{}` not found", p.uri) }
413    }
414
415    /// Handles [`tools/list`]
416    ///
417    /// [`tools/list`]: https://spec.modelcontextprotocol.io/specification/2024-11-05/server/tools/#listing-tools
418    #[allow(unused_variables)]
419    fn tools_list(
420        &self,
421        p: ListToolsRequestParams,
422        cx: &mut RequestContext,
423    ) -> impl Future<Output = Result<ListToolsResult>> + Send {
424        async { Ok(ListToolsResult::default()) }
425    }
426
427    /// Handles [`tools/call`]
428    ///
429    /// [`tools/call`]: https://spec.modelcontextprotocol.io/specification/2024-11-05/server/tools/#calling-a-tool
430    #[allow(unused_variables)]
431    fn tools_call(
432        &self,
433        p: CallToolRequestParams,
434        cx: &mut RequestContext,
435    ) -> impl Future<Output = Result<CallToolResult>> + Send {
436        async move { Err(tool_not_found(&p.name)) }
437    }
438
439    /// Handles [`completion/complete`]
440    ///
441    /// [`completion/complete`]: https://spec.modelcontextprotocol.io/specification/2024-11-05/server/utilities/completion/#completing-a-prompt
442    #[allow(unused_variables)]
443    fn completion_complete(
444        &self,
445        p: CompleteRequestParams,
446        cx: &mut RequestContext,
447    ) -> impl Future<Output = Result<CompleteResult>> + Send {
448        async { Ok(CompleteResult::default()) }
449    }
450
451    /// Gets the JSON RPC `Handler`
452    fn into_handler(self) -> impl Handler + Send + Sync + 'static
453    where
454        Self: Sized + Send + Sync + 'static,
455    {
456        McpServerHandler::new(self)
457    }
458}
459
460/// Context for retrieving request-related information and calling client features
461pub struct RequestContext {
462    session: SessionContext,
463    id: RequestId,
464    initialize: Arc<InitializeRequestParams>,
465}
466
467impl RequestContext {
468    fn new(
469        cx: &RequestContextAs<impl Serialize>,
470        initialize: Arc<InitializeRequestParams>,
471    ) -> Self {
472        Self {
473            session: cx.session(),
474            id: cx.id().clone(),
475            initialize,
476        }
477    }
478
479    /// Gets client information
480    pub fn client_info(&self) -> &Implementation {
481        &self.initialize.client_info
482    }
483
484    /// Gets client capabilities
485    pub fn client_capabilities(&self) -> &ClientCapabilities {
486        &self.initialize.capabilities
487    }
488
489    /// Notifies progress of the request associated with this context
490    ///
491    /// See [`notifications/progress`]
492    ///
493    /// [`notifications/progress`]: https://spec.modelcontextprotocol.io/specification/2024-11-05/server/notifications/#progress-notification
494    pub fn progress(&self, progress: f64, total: Option<f64>) {
495        self.session
496            .notification(
497                "notifications/progress",
498                Some(&ProgressNotificationParams {
499                    progress,
500                    total,
501                    progress_token: self.id.clone(),
502                }),
503            )
504            .unwrap();
505    }
506
507    /// Calls [`sampling/createMessage`]
508    pub async fn sampling_create_message(
509        &self,
510        p: CreateMessageRequestParams,
511    ) -> SessionResult<CreateMessageResult> {
512        self.session
513            .request("sampling/createMessage", Some(&p))
514            .await
515    }
516
517    /// Calls [`roots/list`]
518    pub async fn roots_list(&self) -> SessionResult<Vec<Root>> {
519        let res: ListRootsResult = self
520            .session
521            .request("roots/list", Some(&ListRootsRequestParams::default()))
522            .await?;
523        Ok(res.roots)
524    }
525}
526
527/// Runs an MCP server using stdio transport
528pub async fn serve_stdio(server: impl McpServer) -> SessionResult<()> {
529    Session::from_stdio(McpServerHandler::new(server), &SessionOptions::default())
530        .wait()
531        .await
532}
533
534/// Runs an MCP server using stdio transport with specified options
535pub async fn serve_stdio_with(
536    server: impl McpServer,
537    options: &SessionOptions,
538) -> SessionResult<()> {
539    Session::from_stdio(McpServerHandler::new(server), options)
540        .wait()
541        .await
542}