1use 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#[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 pub stateful: bool,
48 pub session_idle_timeout_ms: u64,
50 pub event_replay_capacity: usize,
52 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#[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#[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#[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 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 pub fn run(&self) -> io::Result<actix_web::dev::Server> {
506 self.run_on(TcpListener::bind(self.config.address)?)
507 }
508
509 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 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 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 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(¶ms, "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(¶ms, "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(¶ms, "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}