Skip to main content

rust_zero_mcp/
lib.rs

1//! Model Context Protocol server support.
2//!
3//! This crate starts with the stateless form of the 2025-03-26 Streamable HTTP
4//! transport. A server can be mounted in any Actix application and shares its
5//! normal middleware, listener, and graceful-shutdown lifecycle.
6
7use actix_web::{
8    http::{header, StatusCode},
9    web, App, HttpRequest, HttpResponse, HttpServer,
10};
11use futures::{
12    future::{AbortHandle, Abortable, BoxFuture},
13    Stream,
14};
15use serde::{Deserialize, Serialize};
16use serde_json::{json, Value};
17use std::{
18    collections::{BTreeMap, HashMap, VecDeque},
19    fmt,
20    future::Future,
21    io,
22    net::{SocketAddr, TcpListener},
23    pin::Pin,
24    sync::{
25        atomic::{AtomicBool, AtomicU64, AtomicUsize, Ordering},
26        Arc, Mutex, RwLock,
27    },
28    time::{Duration, Instant},
29};
30use tokio::sync::broadcast;
31use uuid::Uuid;
32
33pub const LATEST_PROTOCOL_VERSION: &str = "2025-03-26";
34
35/// Configuration for a Streamable HTTP endpoint.
36#[derive(Clone, Debug, Deserialize, Serialize)]
37#[serde(default)]
38pub struct McpServerConfig {
39    pub address: SocketAddr,
40    pub workers: usize,
41    pub shutdown_timeout_ms: u64,
42    pub endpoint: String,
43    pub name: String,
44    pub version: String,
45    pub message_timeout_ms: u64,
46    /// Enables MCP sessions and the GET/DELETE transport methods.
47    pub stateful: bool,
48    /// Idle sessions are rejected and lazily removed after this interval.
49    pub session_idle_timeout_ms: u64,
50    /// Number of SSE events retained per session for `Last-Event-ID` replay.
51    pub event_replay_capacity: usize,
52    /// Permitted browser origins. An empty list rejects requests carrying an
53    /// `Origin` header while allowing non-browser clients.
54    pub allowed_origins: Vec<String>,
55}
56
57impl Default for McpServerConfig {
58    fn default() -> Self {
59        Self {
60            address: "127.0.0.1:8081".parse().unwrap(),
61            workers: 1,
62            shutdown_timeout_ms: 30_000,
63            endpoint: "/mcp".into(),
64            name: "rust-zero-mcp".into(),
65            version: "1.0.0".into(),
66            message_timeout_ms: 30_000,
67            stateful: false,
68            session_idle_timeout_ms: 30 * 60 * 1_000,
69            event_replay_capacity: 256,
70            allowed_origins: Vec::new(),
71        }
72    }
73}
74
75impl McpServerConfig {
76    pub fn validate(&self) -> Result<(), McpConfigError> {
77        if self.workers == 0 {
78            return Err(McpConfigError("workers must be positive"));
79        }
80        if self.shutdown_timeout_ms == 0 {
81            return Err(McpConfigError("shutdown_timeout_ms must be positive"));
82        }
83        if !self.endpoint.starts_with('/') || self.endpoint.contains('?') {
84            return Err(McpConfigError("endpoint must be an absolute path"));
85        }
86        if self.name.trim().is_empty() {
87            return Err(McpConfigError("name must not be empty"));
88        }
89        if self.version.trim().is_empty() {
90            return Err(McpConfigError("version must not be empty"));
91        }
92        if self.message_timeout_ms == 0 {
93            return Err(McpConfigError("message_timeout_ms must be positive"));
94        }
95        if self.stateful && self.session_idle_timeout_ms == 0 {
96            return Err(McpConfigError("session_idle_timeout_ms must be positive"));
97        }
98        if self.stateful && self.event_replay_capacity == 0 {
99            return Err(McpConfigError("event_replay_capacity must be positive"));
100        }
101        Ok(())
102    }
103}
104
105#[derive(Clone, Copy, Debug, Eq, PartialEq)]
106pub struct McpConfigError(&'static str);
107
108impl fmt::Display for McpConfigError {
109    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
110        f.write_str(self.0)
111    }
112}
113
114impl std::error::Error for McpConfigError {}
115
116/// Selected request metadata copied at the HTTP boundary for use by handlers.
117#[derive(Clone, Debug, Default, Eq, PartialEq)]
118pub struct RequestMetadata {
119    pub headers: BTreeMap<String, String>,
120    pub query: BTreeMap<String, String>,
121    pub path: BTreeMap<String, String>,
122}
123
124impl RequestMetadata {
125    pub fn header(&self, name: &str) -> Option<&str> {
126        self.headers
127            .get(&name.to_ascii_lowercase())
128            .map(String::as_str)
129    }
130
131    pub fn query(&self, name: &str) -> Option<&str> {
132        self.query.get(name).map(String::as_str)
133    }
134
135    pub fn path(&self, name: &str) -> Option<&str> {
136        self.path.get(name).map(String::as_str)
137    }
138
139    fn from_request(request: &HttpRequest) -> Self {
140        let headers = request
141            .headers()
142            .iter()
143            .filter_map(|(name, value)| {
144                value
145                    .to_str()
146                    .ok()
147                    .map(|value| (name.as_str().to_ascii_lowercase(), value.to_owned()))
148            })
149            .collect();
150        let query = web::Query::<BTreeMap<String, String>>::from_query(request.query_string())
151            .map(|query| query.into_inner())
152            .unwrap_or_default();
153        let path = request
154            .match_info()
155            .iter()
156            .map(|(name, value)| (name.to_owned(), value.to_owned()))
157            .collect();
158        Self {
159            headers,
160            query,
161            path,
162        }
163    }
164}
165
166#[derive(Clone, Debug, Deserialize, Serialize)]
167#[serde(rename_all = "camelCase")]
168pub struct Tool {
169    pub name: String,
170    #[serde(skip_serializing_if = "Option::is_none")]
171    pub description: Option<String>,
172    pub input_schema: Value,
173}
174
175impl Tool {
176    pub fn new(name: impl Into<String>, input_schema: Value) -> Self {
177        Self {
178            name: name.into(),
179            description: None,
180            input_schema,
181        }
182    }
183
184    pub fn with_description(mut self, description: impl Into<String>) -> Self {
185        self.description = Some(description.into());
186        self
187    }
188}
189
190#[derive(Clone, Debug, Deserialize, Serialize)]
191#[serde(rename_all = "camelCase")]
192pub struct Resource {
193    pub uri: String,
194    pub name: String,
195    #[serde(skip_serializing_if = "Option::is_none")]
196    pub description: Option<String>,
197    #[serde(skip_serializing_if = "Option::is_none")]
198    pub mime_type: Option<String>,
199}
200
201#[derive(Clone, Debug, Deserialize, Serialize)]
202pub struct Prompt {
203    pub name: String,
204    #[serde(skip_serializing_if = "Option::is_none")]
205    pub description: Option<String>,
206    #[serde(default, skip_serializing_if = "Vec::is_empty")]
207    pub arguments: Vec<PromptArgument>,
208}
209
210#[derive(Clone, Debug, Deserialize, Serialize)]
211pub struct PromptArgument {
212    pub name: String,
213    #[serde(skip_serializing_if = "Option::is_none")]
214    pub description: Option<String>,
215    #[serde(default)]
216    pub required: bool,
217}
218
219/// A handler failure returned as a protocol-compliant JSON-RPC error.
220#[derive(Clone, Debug)]
221pub struct McpError {
222    pub code: i64,
223    pub message: String,
224    pub data: Option<Value>,
225}
226
227impl McpError {
228    pub fn new(code: i64, message: impl Into<String>) -> Self {
229        Self {
230            code,
231            message: message.into(),
232            data: None,
233        }
234    }
235
236    pub fn invalid_params(message: impl Into<String>) -> Self {
237        Self::new(-32602, message)
238    }
239
240    pub fn with_data(mut self, data: Value) -> Self {
241        self.data = Some(data);
242        self
243    }
244}
245
246type Handler = Arc<
247    dyn Fn(RequestMetadata, Value) -> BoxFuture<'static, Result<Value, McpError>> + Send + Sync,
248>;
249
250#[derive(Clone)]
251struct Registered<T> {
252    definition: T,
253    handler: Handler,
254}
255
256#[derive(Default)]
257struct Registry {
258    tools: RwLock<BTreeMap<String, Registered<Tool>>>,
259    resources: RwLock<BTreeMap<String, Registered<Resource>>>,
260    prompts: RwLock<BTreeMap<String, Registered<Prompt>>>,
261}
262
263const SESSION_HEADER: &str = "mcp-session-id";
264const LAST_EVENT_ID_HEADER: &str = "last-event-id";
265
266#[derive(Clone)]
267enum SessionEvent {
268    Message(StoredEvent),
269    Terminated,
270}
271
272#[derive(Clone)]
273struct StoredEvent {
274    id: u64,
275    payload: Value,
276}
277
278struct Session {
279    id: String,
280    last_access: Mutex<Instant>,
281    next_event_id: AtomicU64,
282    events: Mutex<VecDeque<StoredEvent>>,
283    sender: broadcast::Sender<SessionEvent>,
284    requests: Mutex<HashMap<String, AbortHandle>>,
285    terminated: AtomicBool,
286    active_streams: AtomicUsize,
287    replay_capacity: usize,
288}
289
290impl Session {
291    fn new(replay_capacity: usize) -> Arc<Self> {
292        let (sender, _) = broadcast::channel(replay_capacity.max(16));
293        Arc::new(Self {
294            id: Uuid::new_v4().to_string(),
295            last_access: Mutex::new(Instant::now()),
296            next_event_id: AtomicU64::new(1),
297            events: Mutex::new(VecDeque::with_capacity(replay_capacity)),
298            sender,
299            requests: Mutex::new(HashMap::new()),
300            terminated: AtomicBool::new(false),
301            active_streams: AtomicUsize::new(0),
302            replay_capacity,
303        })
304    }
305
306    fn touch(&self) {
307        *self.last_access.lock().unwrap() = Instant::now();
308    }
309
310    fn is_expired(&self, idle_timeout: Duration) -> bool {
311        self.terminated.load(Ordering::Acquire)
312            || (self.active_streams.load(Ordering::Acquire) == 0
313                && self.requests.lock().unwrap().is_empty()
314                && self.last_access.lock().unwrap().elapsed() >= idle_timeout)
315    }
316
317    fn publish(&self, payload: Value) -> StoredEvent {
318        let event = StoredEvent {
319            id: self.next_event_id.fetch_add(1, Ordering::Relaxed),
320            payload,
321        };
322        let mut events = self.events.lock().unwrap();
323        if events.len() == self.replay_capacity {
324            events.pop_front();
325        }
326        events.push_back(event.clone());
327        drop(events);
328        let _ = self.sender.send(SessionEvent::Message(event.clone()));
329        event
330    }
331
332    fn cancel(&self, id: &Value) -> bool {
333        let key = request_key(id);
334        self.requests
335            .lock()
336            .unwrap()
337            .get(&key)
338            .map(|handle| {
339                handle.abort();
340                true
341            })
342            .unwrap_or(false)
343    }
344
345    fn terminate(&self) {
346        if !self.terminated.swap(true, Ordering::AcqRel) {
347            for handle in self.requests.lock().unwrap().values() {
348                handle.abort();
349            }
350            let _ = self.sender.send(SessionEvent::Terminated);
351        }
352    }
353
354    fn close_event_streams(&self) {
355        let _ = self.sender.send(SessionEvent::Terminated);
356    }
357
358    fn event_stream(
359        self: &Arc<Self>,
360        after: u64,
361    ) -> Pin<Box<dyn Stream<Item = Result<web::Bytes, actix_web::Error>>>> {
362        struct State {
363            session: Arc<Session>,
364            _active: ActiveStream,
365            replay: VecDeque<StoredEvent>,
366            receiver: broadcast::Receiver<SessionEvent>,
367            last_id: u64,
368        }
369
370        let receiver = self.sender.subscribe();
371        let replay = self
372            .events
373            .lock()
374            .unwrap()
375            .iter()
376            .filter(|event| event.id > after)
377            .cloned()
378            .collect();
379        self.active_streams.fetch_add(1, Ordering::AcqRel);
380        let state = State {
381            session: self.clone(),
382            _active: ActiveStream(self.clone()),
383            replay,
384            receiver,
385            last_id: after,
386        };
387        Box::pin(futures::stream::unfold(state, |mut state| async move {
388            loop {
389                if let Some(event) = state.replay.pop_front() {
390                    state.last_id = event.id;
391                    return Some((Ok(sse_event(&event)), state));
392                }
393                match state.receiver.recv().await {
394                    Ok(SessionEvent::Message(event)) if event.id > state.last_id => {
395                        state.last_id = event.id;
396                        return Some((Ok(sse_event(&event)), state));
397                    }
398                    Ok(SessionEvent::Message(_)) => continue,
399                    Ok(SessionEvent::Terminated) | Err(broadcast::error::RecvError::Closed) => {
400                        return None;
401                    }
402                    Err(broadcast::error::RecvError::Lagged(_)) => {
403                        state.replay = state
404                            .session
405                            .events
406                            .lock()
407                            .unwrap()
408                            .iter()
409                            .filter(|event| event.id > state.last_id)
410                            .cloned()
411                            .collect();
412                    }
413                }
414            }
415        }))
416    }
417}
418
419struct ActiveStream(Arc<Session>);
420
421impl Drop for ActiveStream {
422    fn drop(&mut self) {
423        self.0.active_streams.fetch_sub(1, Ordering::AcqRel);
424        self.0.touch();
425    }
426}
427
428#[derive(Default)]
429struct Sessions {
430    values: RwLock<HashMap<String, Arc<Session>>>,
431}
432
433/// Cloneable MCP protocol core and Actix Streamable HTTP handler.
434#[derive(Clone)]
435pub struct McpServer {
436    config: McpServerConfig,
437    registry: Arc<Registry>,
438    sessions: Arc<Sessions>,
439}
440
441impl McpServer {
442    pub fn new(config: McpServerConfig) -> Result<Self, McpConfigError> {
443        config.validate()?;
444        Ok(Self {
445            config,
446            registry: Arc::new(Registry::default()),
447            sessions: Arc::new(Sessions::default()),
448        })
449    }
450
451    pub fn add_tool<F, Fut>(&self, tool: Tool, handler: F)
452    where
453        F: Fn(RequestMetadata, Value) -> Fut + Send + Sync + 'static,
454        Fut: std::future::Future<Output = Result<Value, McpError>> + Send + 'static,
455    {
456        self.registry.tools.write().unwrap().insert(
457            tool.name.clone(),
458            Registered {
459                definition: tool,
460                handler: Arc::new(move |metadata, params| Box::pin(handler(metadata, params))),
461            },
462        );
463    }
464
465    pub fn add_resource<F, Fut>(&self, resource: Resource, handler: F)
466    where
467        F: Fn(RequestMetadata, Value) -> Fut + Send + Sync + 'static,
468        Fut: std::future::Future<Output = Result<Value, McpError>> + Send + 'static,
469    {
470        self.registry.resources.write().unwrap().insert(
471            resource.uri.clone(),
472            Registered {
473                definition: resource,
474                handler: Arc::new(move |metadata, params| Box::pin(handler(metadata, params))),
475            },
476        );
477    }
478
479    pub fn add_prompt<F, Fut>(&self, prompt: Prompt, handler: F)
480    where
481        F: Fn(RequestMetadata, Value) -> Fut + Send + Sync + 'static,
482        Fut: std::future::Future<Output = Result<Value, McpError>> + Send + 'static,
483    {
484        self.registry.prompts.write().unwrap().insert(
485            prompt.name.clone(),
486            Registered {
487                definition: prompt,
488                handler: Arc::new(move |metadata, params| Box::pin(handler(metadata, params))),
489            },
490        );
491    }
492
493    /// Mounts the configured MCP endpoint on an Actix application.
494    pub fn configure(&self, service: &mut web::ServiceConfig) {
495        service.service(
496            web::resource(self.config.endpoint.clone())
497                .app_data(web::Data::new(self.clone()))
498                .route(web::post().to(Self::http_post))
499                .route(web::get().to(Self::http_get))
500                .route(web::delete().to(Self::http_delete)),
501        );
502    }
503
504    /// Binds the configured address and starts a standalone MCP HTTP server.
505    pub fn run(&self) -> io::Result<actix_web::dev::Server> {
506        self.run_on(TcpListener::bind(self.config.address)?)
507    }
508
509    /// Starts a standalone MCP HTTP server on an existing listener.
510    pub fn run_on(&self, listener: TcpListener) -> io::Result<actix_web::dev::Server> {
511        let server = self.clone();
512        let workers = self.config.workers;
513        let shutdown_seconds = self.config.shutdown_timeout_ms.div_ceil(1_000);
514        HttpServer::new(move || {
515            let server = server.clone();
516            App::new().configure(move |service| server.configure(service))
517        })
518        .workers(workers)
519        .shutdown_timeout(shutdown_seconds)
520        .listen(listener)
521        .map(HttpServer::run)
522    }
523
524    /// Serves until the supplied signal resolves and then gracefully drains requests.
525    pub async fn serve_until<F>(&self, shutdown: F) -> io::Result<()>
526    where
527        F: Future<Output = ()>,
528    {
529        let transport = self.run()?;
530        self.drain_on_signal(transport, shutdown).await
531    }
532
533    /// Listener-based variant of [`McpServer::serve_until`].
534    pub async fn serve_on_until<F>(&self, listener: TcpListener, shutdown: F) -> io::Result<()>
535    where
536        F: Future<Output = ()>,
537    {
538        let transport = self.run_on(listener)?;
539        self.drain_on_signal(transport, shutdown).await
540    }
541
542    async fn drain_on_signal<F>(
543        &self,
544        transport: actix_web::dev::Server,
545        shutdown: F,
546    ) -> io::Result<()>
547    where
548        F: Future<Output = ()>,
549    {
550        use futures::future::{select, Either};
551
552        let handle = transport.handle();
553        match select(Box::pin(transport), Box::pin(shutdown)).await {
554            Either::Left((result, _)) => result,
555            Either::Right(((), transport)) => {
556                self.close_event_streams();
557                let (_, result) = futures::future::join(handle.stop(true), transport).await;
558                self.terminate_sessions();
559                result
560            }
561        }
562    }
563
564    fn close_event_streams(&self) {
565        for session in self.sessions.values.read().unwrap().values() {
566            session.close_event_streams();
567        }
568    }
569
570    fn terminate_sessions(&self) {
571        let mut sessions = self.sessions.values.write().unwrap();
572        for session in sessions.values() {
573            session.terminate();
574        }
575        sessions.clear();
576    }
577
578    fn create_session(&self) -> Arc<Session> {
579        self.remove_expired_sessions();
580        let session = Session::new(self.config.event_replay_capacity);
581        self.sessions
582            .values
583            .write()
584            .unwrap()
585            .insert(session.id.clone(), session.clone());
586        session
587    }
588
589    fn remove_expired_sessions(&self) {
590        let idle_timeout = Duration::from_millis(self.config.session_idle_timeout_ms);
591        let mut sessions = self.sessions.values.write().unwrap();
592        sessions.retain(|_, session| {
593            let retain = !session.is_expired(idle_timeout);
594            if !retain {
595                session.terminate();
596            }
597            retain
598        });
599    }
600
601    fn request_session(&self, request: &HttpRequest) -> Result<Arc<Session>, HttpResponse> {
602        if !self.config.stateful {
603            return Err(HttpResponse::MethodNotAllowed().finish());
604        }
605        self.remove_expired_sessions();
606        let Some(id) = request
607            .headers()
608            .get(SESSION_HEADER)
609            .and_then(|value| value.to_str().ok())
610        else {
611            return Err(HttpResponse::BadRequest().json(error_response(
612                Value::Null,
613                McpError::new(-32600, "Mcp-Session-Id header is required"),
614            )));
615        };
616        let session = self.sessions.values.read().unwrap().get(id).cloned();
617        match session {
618            Some(session) => {
619                session.touch();
620                Ok(session)
621            }
622            None => Err(HttpResponse::NotFound().json(error_response(
623                Value::Null,
624                McpError::new(-32002, "session not found or expired"),
625            ))),
626        }
627    }
628
629    async fn http_get(server: web::Data<Self>, request: HttpRequest) -> HttpResponse {
630        if let Some(response) = server.reject_origin(&request) {
631            return response;
632        }
633        let accepts_sse = request
634            .headers()
635            .get(header::ACCEPT)
636            .and_then(|value| value.to_str().ok())
637            .is_some_and(|value| value.contains("text/event-stream"));
638        if !accepts_sse {
639            return HttpResponse::NotAcceptable().finish();
640        }
641        let session = match server.request_session(&request) {
642            Ok(session) => session,
643            Err(response) => return response,
644        };
645        let after = match request.headers().get(LAST_EVENT_ID_HEADER) {
646            Some(value) => match value
647                .to_str()
648                .ok()
649                .and_then(|value| value.parse::<u64>().ok())
650            {
651                Some(value) => value,
652                None => {
653                    return HttpResponse::BadRequest().json(error_response(
654                        Value::Null,
655                        McpError::invalid_params("Last-Event-ID must be an unsigned integer"),
656                    ));
657                }
658            },
659            None => session
660                .next_event_id
661                .load(Ordering::Acquire)
662                .saturating_sub(1),
663        };
664        HttpResponse::Ok()
665            .insert_header((header::CONTENT_TYPE, "text/event-stream"))
666            .insert_header((header::CACHE_CONTROL, "no-cache"))
667            .insert_header((SESSION_HEADER, session.id.clone()))
668            .streaming(session.event_stream(after))
669    }
670
671    async fn http_delete(server: web::Data<Self>, request: HttpRequest) -> HttpResponse {
672        if let Some(response) = server.reject_origin(&request) {
673            return response;
674        }
675        let session = match server.request_session(&request) {
676            Ok(session) => session,
677            Err(response) => return response,
678        };
679        server.sessions.values.write().unwrap().remove(&session.id);
680        session.terminate();
681        HttpResponse::NoContent().finish()
682    }
683
684    fn reject_origin(&self, request: &HttpRequest) -> Option<HttpResponse> {
685        let origin = request.headers().get(header::ORIGIN)?;
686        let allowed = origin.to_str().ok().is_some_and(|origin| {
687            self.config
688                .allowed_origins
689                .iter()
690                .any(|item| item == origin)
691        });
692        (!allowed).then(|| {
693            HttpResponse::Forbidden().json(error_response(
694                Value::Null,
695                McpError::new(-32000, "origin is not allowed"),
696            ))
697        })
698    }
699
700    async fn http_post(
701        server: web::Data<Self>,
702        request: HttpRequest,
703        body: web::Bytes,
704    ) -> HttpResponse {
705        if let Some(response) = server.reject_origin(&request) {
706            return response;
707        }
708
709        let content_type_ok = request
710            .headers()
711            .get(header::CONTENT_TYPE)
712            .and_then(|value| value.to_str().ok())
713            .is_some_and(|value| value.split(';').next() == Some("application/json"));
714        if !content_type_ok {
715            return HttpResponse::UnsupportedMediaType().json(error_response(
716                Value::Null,
717                McpError::new(-32600, "Content-Type must be application/json"),
718            ));
719        }
720
721        let message: JsonRpcRequest = match serde_json::from_slice(&body) {
722            Ok(message) => message,
723            Err(error) => {
724                return HttpResponse::BadRequest().json(error_response(
725                    Value::Null,
726                    McpError::new(-32700, "parse error")
727                        .with_data(json!({"detail": error.to_string()})),
728                ));
729            }
730        };
731        if message.jsonrpc != "2.0" || message.method.is_empty() {
732            return HttpResponse::BadRequest().json(error_response(
733                message.id.unwrap_or(Value::Null),
734                McpError::new(-32600, "invalid JSON-RPC request"),
735            ));
736        }
737
738        let session = if server.config.stateful {
739            if message.method == "initialize" {
740                Some(server.create_session())
741            } else {
742                match server.request_session(&request) {
743                    Ok(session) => Some(session),
744                    Err(response) => return response,
745                }
746            }
747        } else {
748            None
749        };
750
751        // Notifications deliberately have no JSON-RPC response. Cancellation
752        // is still dispatched so it can abort a concurrent request.
753        if message.id.is_none() {
754            if message.method == "notifications/cancelled" {
755                if let (Some(session), Some(request_id)) =
756                    (session.as_ref(), message.params.get("requestId"))
757                {
758                    session.cancel(request_id);
759                }
760            }
761            return HttpResponse::Accepted().finish();
762        }
763        let id = message.id.unwrap();
764        let metadata = RequestMetadata::from_request(&request);
765        let (abort_handle, abort_registration) = AbortHandle::new_pair();
766        if let Some(session) = session.as_ref() {
767            session
768                .requests
769                .lock()
770                .unwrap()
771                .insert(request_key(&id), abort_handle);
772        }
773        let outcome = actix_web::rt::time::timeout(
774            Duration::from_millis(server.config.message_timeout_ms),
775            Abortable::new(
776                server.dispatch(&message.method, metadata, message.params),
777                abort_registration,
778            ),
779        )
780        .await;
781        if let Some(session) = session.as_ref() {
782            session.requests.lock().unwrap().remove(&request_key(&id));
783            session.touch();
784        }
785        let response = match outcome {
786            Ok(Ok(Ok(result))) => json!({"jsonrpc": "2.0", "id": id, "result": result}),
787            Ok(Ok(Err(error))) => error_response(id, error),
788            Ok(Err(_)) => error_response(id, McpError::new(-32800, "request cancelled")),
789            Err(_) => error_response(id, McpError::new(-32001, "request timed out")),
790        };
791
792        let accepts_json = request
793            .headers()
794            .get(header::ACCEPT)
795            .and_then(|value| value.to_str().ok())
796            .is_none_or(|value| value.contains("application/json") || value.contains("*/*"));
797        let session_header = session.as_ref().map(|session| session.id.clone());
798        if accepts_json {
799            let mut builder = HttpResponse::Ok();
800            if let Some(session_id) = session_header {
801                builder.insert_header((SESSION_HEADER, session_id));
802            }
803            builder.json(response)
804        } else if request
805            .headers()
806            .get(header::ACCEPT)
807            .and_then(|value| value.to_str().ok())
808            .is_some_and(|value| value.contains("text/event-stream"))
809        {
810            let event = session
811                .as_ref()
812                .map(|session| session.publish(response.clone()));
813            let mut builder = HttpResponse::Ok();
814            if let Some(session_id) = session_header {
815                builder.insert_header((SESSION_HEADER, session_id));
816            }
817            builder
818                .insert_header((header::CONTENT_TYPE, "text/event-stream"))
819                .insert_header((header::CACHE_CONTROL, "no-cache"))
820                .body(event.map_or_else(
821                    || format!("event: message\ndata: {response}\n\n"),
822                    |event| String::from_utf8_lossy(&sse_event(&event)).into_owned(),
823                ))
824        } else {
825            HttpResponse::build(StatusCode::NOT_ACCEPTABLE).finish()
826        }
827    }
828
829    async fn dispatch(
830        &self,
831        method: &str,
832        metadata: RequestMetadata,
833        params: Value,
834    ) -> Result<Value, McpError> {
835        match method {
836            "initialize" => Ok(json!({
837                "protocolVersion": LATEST_PROTOCOL_VERSION,
838                "capabilities": {
839                    "tools": {"listChanged": false},
840                    "resources": {"subscribe": false, "listChanged": false},
841                    "prompts": {"listChanged": false}
842                },
843                "serverInfo": {"name": self.config.name, "version": self.config.version}
844            })),
845            "ping" => Ok(json!({})),
846            "tools/list" => {
847                let tools = self
848                    .registry
849                    .tools
850                    .read()
851                    .unwrap()
852                    .values()
853                    .map(|entry| entry.definition.clone())
854                    .collect::<Vec<_>>();
855                Ok(json!({"tools": tools}))
856            }
857            "resources/list" => {
858                let resources = self
859                    .registry
860                    .resources
861                    .read()
862                    .unwrap()
863                    .values()
864                    .map(|entry| entry.definition.clone())
865                    .collect::<Vec<_>>();
866                Ok(json!({"resources": resources}))
867            }
868            "prompts/list" => {
869                let prompts = self
870                    .registry
871                    .prompts
872                    .read()
873                    .unwrap()
874                    .values()
875                    .map(|entry| entry.definition.clone())
876                    .collect::<Vec<_>>();
877                Ok(json!({"prompts": prompts}))
878            }
879            "tools/call" => {
880                let name = required_string(&params, "name")?;
881                let arguments = params
882                    .get("arguments")
883                    .cloned()
884                    .unwrap_or_else(|| json!({}));
885                let handler = self
886                    .registry
887                    .tools
888                    .read()
889                    .unwrap()
890                    .get(name)
891                    .map(|entry| entry.handler.clone())
892                    .ok_or_else(|| McpError::new(-32602, format!("unknown tool: {name}")))?;
893                handler(metadata, arguments).await
894            }
895            "resources/read" => {
896                let uri = required_string(&params, "uri")?;
897                let handler = self
898                    .registry
899                    .resources
900                    .read()
901                    .unwrap()
902                    .get(uri)
903                    .map(|entry| entry.handler.clone())
904                    .ok_or_else(|| McpError::new(-32602, format!("unknown resource: {uri}")))?;
905                handler(metadata, params).await
906            }
907            "prompts/get" => {
908                let name = required_string(&params, "name")?;
909                let handler = self
910                    .registry
911                    .prompts
912                    .read()
913                    .unwrap()
914                    .get(name)
915                    .map(|entry| entry.handler.clone())
916                    .ok_or_else(|| McpError::new(-32602, format!("unknown prompt: {name}")))?;
917                handler(metadata, params).await
918            }
919            _ => Err(McpError::new(-32601, "method not found")),
920        }
921    }
922}
923
924#[derive(Deserialize)]
925struct JsonRpcRequest {
926    jsonrpc: String,
927    #[serde(default)]
928    id: Option<Value>,
929    method: String,
930    #[serde(default = "empty_object")]
931    params: Value,
932}
933
934fn empty_object() -> Value {
935    json!({})
936}
937
938fn required_string<'a>(params: &'a Value, key: &str) -> Result<&'a str, McpError> {
939    params
940        .get(key)
941        .and_then(Value::as_str)
942        .ok_or_else(|| McpError::invalid_params(format!("{key} must be a string")))
943}
944
945fn error_response(id: Value, error: McpError) -> Value {
946    let mut payload = json!({"code": error.code, "message": error.message});
947    if let Some(data) = error.data {
948        payload["data"] = data;
949    }
950    json!({"jsonrpc": "2.0", "id": id, "error": payload})
951}
952
953fn request_key(id: &Value) -> String {
954    serde_json::to_string(id).unwrap_or_else(|_| "null".into())
955}
956
957fn sse_event(event: &StoredEvent) -> web::Bytes {
958    web::Bytes::from(format!(
959        "id: {}\nevent: message\ndata: {}\n\n",
960        event.id, event.payload
961    ))
962}
963
964#[cfg(test)]
965mod tests {
966    use super::*;
967    use actix_web::{body::MessageBody, http::header, test, App};
968    use futures::future::poll_fn;
969
970    fn server() -> McpServer {
971        let server = McpServer::new(McpServerConfig {
972            allowed_origins: vec!["https://client.example".into()],
973            ..McpServerConfig::default()
974        })
975        .unwrap();
976        server.add_tool(
977            Tool::new("echo", json!({"type": "object"})).with_description("echo input"),
978            |metadata, arguments| async move {
979                Ok(json!({
980                    "content": [{"type": "text", "text": arguments["text"]}],
981                    "_meta": {"tenant": metadata.header("x-tenant")}
982                }))
983            },
984        );
985        server
986    }
987
988    #[actix_web::test]
989    async fn initializes_and_advertises_capabilities() {
990        let server = server();
991        let app = test::init_service(App::new().configure(|cfg| server.configure(cfg))).await;
992        let request = test::TestRequest::post()
993            .uri("/mcp")
994            .insert_header((header::CONTENT_TYPE, "application/json"))
995            .set_payload(r#"{"jsonrpc":"2.0","id":1,"method":"initialize","params":{}}"#)
996            .to_request();
997        let response: Value = test::call_and_read_body_json(&app, request).await;
998        assert_eq!(
999            response["result"]["protocolVersion"],
1000            LATEST_PROTOCOL_VERSION
1001        );
1002        assert_eq!(response["result"]["serverInfo"]["name"], "rust-zero-mcp");
1003        assert!(response["result"]["capabilities"]["tools"].is_object());
1004    }
1005
1006    #[actix_web::test]
1007    async fn lists_and_calls_tools_with_request_metadata() {
1008        let server = server();
1009        let app = test::init_service(App::new().configure(|cfg| server.configure(cfg))).await;
1010        let list = test::TestRequest::post()
1011            .uri("/mcp")
1012            .insert_header((header::CONTENT_TYPE, "application/json"))
1013            .set_json(json!({"jsonrpc":"2.0","id":1,"method":"tools/list"}))
1014            .to_request();
1015        let listed: Value = test::call_and_read_body_json(&app, list).await;
1016        assert_eq!(listed["result"]["tools"][0]["name"], "echo");
1017
1018        let call = test::TestRequest::post()
1019            .uri("/mcp?trace=abc")
1020            .insert_header((header::CONTENT_TYPE, "application/json"))
1021            .insert_header(("x-tenant", "acme"))
1022            .set_json(json!({
1023                "jsonrpc":"2.0", "id":"call-1", "method":"tools/call",
1024                "params":{"name":"echo", "arguments":{"text":"hello"}}
1025            }))
1026            .to_request();
1027        let called: Value = test::call_and_read_body_json(&app, call).await;
1028        assert_eq!(called["result"]["content"][0]["text"], "hello");
1029        assert_eq!(called["result"]["_meta"]["tenant"], "acme");
1030    }
1031
1032    #[actix_web::test]
1033    async fn dispatches_resources_and_prompts() {
1034        let server = server();
1035        server.add_resource(
1036            Resource {
1037                uri: "file:///guide.md".into(),
1038                name: "guide".into(),
1039                description: Some("project guide".into()),
1040                mime_type: Some("text/markdown".into()),
1041            },
1042            |_, params| async move {
1043                Ok(json!({"contents": [{
1044                    "uri": params["uri"], "mimeType": "text/markdown", "text": "guide"
1045                }]}))
1046            },
1047        );
1048        server.add_prompt(
1049            Prompt {
1050                name: "review".into(),
1051                description: Some("review code".into()),
1052                arguments: vec![PromptArgument {
1053                    name: "code".into(),
1054                    description: None,
1055                    required: true,
1056                }],
1057            },
1058            |_, params| async move {
1059                Ok(json!({"messages": [{
1060                    "role": "user",
1061                    "content": {"type": "text", "text": params["arguments"]["code"]}
1062                }]}))
1063            },
1064        );
1065        let app = test::init_service(App::new().configure(|cfg| server.configure(cfg))).await;
1066
1067        let resource = test::TestRequest::post()
1068            .uri("/mcp")
1069            .insert_header((header::CONTENT_TYPE, "application/json"))
1070            .set_json(json!({
1071                "jsonrpc":"2.0", "id":1, "method":"resources/read",
1072                "params":{"uri":"file:///guide.md"}
1073            }))
1074            .to_request();
1075        let resource: Value = test::call_and_read_body_json(&app, resource).await;
1076        assert_eq!(resource["result"]["contents"][0]["text"], "guide");
1077
1078        let prompt = test::TestRequest::post()
1079            .uri("/mcp")
1080            .insert_header((header::CONTENT_TYPE, "application/json"))
1081            .set_json(json!({
1082                "jsonrpc":"2.0", "id":2, "method":"prompts/get",
1083                "params":{"name":"review", "arguments":{"code":"fn main() {}"}}
1084            }))
1085            .to_request();
1086        let prompt: Value = test::call_and_read_body_json(&app, prompt).await;
1087        assert_eq!(
1088            prompt["result"]["messages"][0]["content"]["text"],
1089            "fn main() {}"
1090        );
1091    }
1092
1093    #[actix_web::test]
1094    async fn projects_header_query_and_path_metadata() {
1095        let server = McpServer::new(McpServerConfig {
1096            endpoint: "/mcp/{scope}".into(),
1097            ..McpServerConfig::default()
1098        })
1099        .unwrap();
1100        server.add_tool(Tool::new("metadata", json!({})), |metadata, _| async move {
1101            Ok(json!({"content": [{
1102                "type": "text",
1103                "text": format!(
1104                    "{}/{}/{}",
1105                    metadata.header("x-tenant").unwrap_or_default(),
1106                    metadata.query("trace").unwrap_or_default(),
1107                    metadata.path("scope").unwrap_or_default()
1108                )
1109            }]}))
1110        });
1111        let app = test::init_service(App::new().configure(|cfg| server.configure(cfg))).await;
1112        let request = test::TestRequest::post()
1113            .uri("/mcp/admin?trace=abc")
1114            .insert_header((header::CONTENT_TYPE, "application/json"))
1115            .insert_header(("x-tenant", "acme"))
1116            .set_json(json!({
1117                "jsonrpc":"2.0", "id":1, "method":"tools/call",
1118                "params":{"name":"metadata"}
1119            }))
1120            .to_request();
1121        let response: Value = test::call_and_read_body_json(&app, request).await;
1122        assert_eq!(response["result"]["content"][0]["text"], "acme/abc/admin");
1123    }
1124
1125    #[actix_web::test]
1126    async fn returns_protocol_errors_and_accepts_notifications() {
1127        let server = server();
1128        let app = test::init_service(App::new().configure(|cfg| server.configure(cfg))).await;
1129        let missing = test::TestRequest::post()
1130            .uri("/mcp")
1131            .insert_header((header::CONTENT_TYPE, "application/json"))
1132            .set_json(json!({"jsonrpc":"2.0","id":7,"method":"missing"}))
1133            .to_request();
1134        let body: Value = test::call_and_read_body_json(&app, missing).await;
1135        assert_eq!(body["error"]["code"], -32601);
1136
1137        let notification = test::TestRequest::post()
1138            .uri("/mcp")
1139            .insert_header((header::CONTENT_TYPE, "application/json"))
1140            .set_json(json!({"jsonrpc":"2.0","method":"notifications/initialized"}))
1141            .to_request();
1142        let response = test::call_service(&app, notification).await;
1143        assert_eq!(response.status(), StatusCode::ACCEPTED);
1144    }
1145
1146    #[actix_web::test]
1147    async fn supports_sse_response_and_origin_protection() {
1148        let server = server();
1149        let app = test::init_service(App::new().configure(|cfg| server.configure(cfg))).await;
1150        let sse = test::TestRequest::post()
1151            .uri("/mcp")
1152            .insert_header((header::CONTENT_TYPE, "application/json"))
1153            .insert_header((header::ACCEPT, "text/event-stream"))
1154            .insert_header((header::ORIGIN, "https://client.example"))
1155            .set_json(json!({"jsonrpc":"2.0","id":1,"method":"ping"}))
1156            .to_request();
1157        let response = test::call_service(&app, sse).await;
1158        assert_eq!(response.status(), StatusCode::OK);
1159        assert_eq!(
1160            response.headers().get(header::CONTENT_TYPE).unwrap(),
1161            "text/event-stream"
1162        );
1163        let body = test::read_body(response).await;
1164        assert!(std::str::from_utf8(&body)
1165            .unwrap()
1166            .starts_with("event: message\ndata: "));
1167
1168        let rejected = test::TestRequest::post()
1169            .uri("/mcp")
1170            .insert_header((header::CONTENT_TYPE, "application/json"))
1171            .insert_header((header::ORIGIN, "https://evil.example"))
1172            .set_json(json!({"jsonrpc":"2.0","id":1,"method":"ping"}))
1173            .to_request();
1174        let response = test::call_service(&app, rejected).await;
1175        assert_eq!(response.status(), StatusCode::FORBIDDEN);
1176    }
1177
1178    fn stateful_server() -> McpServer {
1179        McpServer::new(McpServerConfig {
1180            stateful: true,
1181            event_replay_capacity: 8,
1182            ..McpServerConfig::default()
1183        })
1184        .unwrap()
1185    }
1186
1187    async fn initialize_session<S>(app: &S) -> String
1188    where
1189        S: actix_web::dev::Service<
1190            actix_http::Request,
1191            Response = actix_web::dev::ServiceResponse,
1192            Error = actix_web::Error,
1193        >,
1194    {
1195        let request = test::TestRequest::post()
1196            .uri("/mcp")
1197            .insert_header((header::CONTENT_TYPE, "application/json"))
1198            .set_json(json!({"jsonrpc":"2.0","id":1,"method":"initialize"}))
1199            .to_request();
1200        let response = test::call_service(app, request).await;
1201        assert_eq!(response.status(), StatusCode::OK);
1202        response
1203            .headers()
1204            .get(SESSION_HEADER)
1205            .unwrap()
1206            .to_str()
1207            .unwrap()
1208            .to_owned()
1209    }
1210
1211    #[actix_web::test]
1212    async fn stateful_sessions_are_required_and_can_be_terminated() {
1213        let server = stateful_server();
1214        let app = test::init_service(App::new().configure(|cfg| server.configure(cfg))).await;
1215        let session_id = initialize_session(&app).await;
1216
1217        let missing = test::TestRequest::post()
1218            .uri("/mcp")
1219            .insert_header((header::CONTENT_TYPE, "application/json"))
1220            .set_json(json!({"jsonrpc":"2.0","id":2,"method":"ping"}))
1221            .to_request();
1222        assert_eq!(
1223            test::call_service(&app, missing).await.status(),
1224            StatusCode::BAD_REQUEST
1225        );
1226
1227        let delete = test::TestRequest::delete()
1228            .uri("/mcp")
1229            .insert_header((SESSION_HEADER, session_id.clone()))
1230            .to_request();
1231        assert_eq!(
1232            test::call_service(&app, delete).await.status(),
1233            StatusCode::NO_CONTENT
1234        );
1235
1236        let expired = test::TestRequest::post()
1237            .uri("/mcp")
1238            .insert_header((header::CONTENT_TYPE, "application/json"))
1239            .insert_header((SESSION_HEADER, session_id))
1240            .set_json(json!({"jsonrpc":"2.0","id":3,"method":"ping"}))
1241            .to_request();
1242        assert_eq!(
1243            test::call_service(&app, expired).await.status(),
1244            StatusCode::NOT_FOUND
1245        );
1246    }
1247
1248    #[actix_web::test]
1249    async fn get_stream_replays_events_after_cursor_on_reconnect() {
1250        let server = stateful_server();
1251        let app = test::init_service(App::new().configure(|cfg| server.configure(cfg))).await;
1252        let session_id = initialize_session(&app).await;
1253
1254        for id in [10, 11] {
1255            let request = test::TestRequest::post()
1256                .uri("/mcp")
1257                .insert_header((header::CONTENT_TYPE, "application/json"))
1258                .insert_header((header::ACCEPT, "text/event-stream"))
1259                .insert_header((SESSION_HEADER, session_id.clone()))
1260                .set_json(json!({"jsonrpc":"2.0","id":id,"method":"ping"}))
1261                .to_request();
1262            assert_eq!(
1263                test::call_service(&app, request).await.status(),
1264                StatusCode::OK
1265            );
1266        }
1267
1268        let reconnect = test::TestRequest::get()
1269            .uri("/mcp")
1270            .insert_header((header::ACCEPT, "text/event-stream"))
1271            .insert_header((SESSION_HEADER, session_id))
1272            .insert_header((LAST_EVENT_ID_HEADER, "1"))
1273            .to_request();
1274        let response = test::call_service(&app, reconnect).await;
1275        assert_eq!(response.status(), StatusCode::OK);
1276        let mut body = response.into_body();
1277        let chunk = poll_fn(|cx| Pin::new(&mut body).poll_next(cx))
1278            .await
1279            .unwrap()
1280            .unwrap();
1281        let chunk = std::str::from_utf8(&chunk).unwrap();
1282        assert!(chunk.starts_with("id: 2\nevent: message\n"));
1283        assert!(chunk.contains(r#""id":11"#));
1284    }
1285
1286    #[actix_web::test]
1287    async fn cancellation_notification_aborts_an_in_flight_request() {
1288        let server = stateful_server();
1289        let started = Arc::new(tokio::sync::Notify::new());
1290        server.add_tool(Tool::new("wait", json!({})), {
1291            let started = started.clone();
1292            move |_, _| {
1293                let started = started.clone();
1294                async move {
1295                    started.notify_one();
1296                    futures::future::pending::<Result<Value, McpError>>().await
1297                }
1298            }
1299        });
1300        let app = test::init_service(App::new().configure(|cfg| server.configure(cfg))).await;
1301        let session_id = initialize_session(&app).await;
1302
1303        let call = test::TestRequest::post()
1304            .uri("/mcp")
1305            .insert_header((header::CONTENT_TYPE, "application/json"))
1306            .insert_header((SESSION_HEADER, session_id.clone()))
1307            .set_json(json!({
1308                "jsonrpc":"2.0", "id":"slow", "method":"tools/call",
1309                "params":{"name":"wait"}
1310            }))
1311            .to_request();
1312        let cancel = async {
1313            started.notified().await;
1314            let request = test::TestRequest::post()
1315                .uri("/mcp")
1316                .insert_header((header::CONTENT_TYPE, "application/json"))
1317                .insert_header((SESSION_HEADER, session_id))
1318                .set_json(json!({
1319                    "jsonrpc":"2.0", "method":"notifications/cancelled",
1320                    "params":{"requestId":"slow", "reason":"client disconnected"}
1321                }))
1322                .to_request();
1323            test::call_service(&app, request).await
1324        };
1325        let (response, cancellation) = futures::join!(test::call_service(&app, call), cancel);
1326        assert_eq!(cancellation.status(), StatusCode::ACCEPTED);
1327        let body: Value = test::read_body_json(response).await;
1328        assert_eq!(body["error"]["code"], -32800);
1329    }
1330
1331    #[actix_web::test]
1332    async fn shutdown_gracefully_drains_an_in_flight_tool_call() {
1333        let listener = TcpListener::bind("127.0.0.1:0").unwrap();
1334        let address = listener.local_addr().unwrap();
1335        let started = Arc::new(tokio::sync::Notify::new());
1336        let release = Arc::new(tokio::sync::Notify::new());
1337        let server = McpServer::new(McpServerConfig {
1338            message_timeout_ms: 2_000,
1339            shutdown_timeout_ms: 2_000,
1340            ..McpServerConfig::default()
1341        })
1342        .unwrap();
1343        server.add_tool(Tool::new("slow", json!({})), {
1344            let started = started.clone();
1345            let release = release.clone();
1346            move |_, _| {
1347                let started = started.clone();
1348                let release = release.clone();
1349                async move {
1350                    started.notify_one();
1351                    release.notified().await;
1352                    Ok(json!({"content": [{"type": "text", "text": "finished"}]}))
1353                }
1354            }
1355        });
1356        let (shutdown_sender, shutdown_receiver) = tokio::sync::oneshot::channel();
1357        let server_task = actix_web::rt::spawn(async move {
1358            server
1359                .serve_on_until(listener, async move {
1360                    let _ = shutdown_receiver.await;
1361                })
1362                .await
1363        });
1364        let request_task = actix_web::rt::spawn(async move {
1365            reqwest::Client::new()
1366                .post(format!("http://{address}/mcp"))
1367                .json(&json!({
1368                    "jsonrpc":"2.0", "id":1, "method":"tools/call",
1369                    "params":{"name":"slow"}
1370                }))
1371                .send()
1372                .await
1373        });
1374        actix_web::rt::time::timeout(Duration::from_secs(1), started.notified())
1375            .await
1376            .unwrap();
1377        shutdown_sender.send(()).unwrap();
1378        actix_web::rt::time::sleep(Duration::from_millis(20)).await;
1379        assert!(!request_task.is_finished());
1380
1381        release.notify_one();
1382        let response = request_task.await.unwrap().unwrap();
1383        assert_eq!(response.status(), reqwest::StatusCode::OK);
1384        assert_eq!(
1385            response.json::<Value>().await.unwrap()["result"]["content"][0]["text"],
1386            "finished"
1387        );
1388        server_task.await.unwrap().unwrap();
1389    }
1390}