1use crate::rpc::{Incoming, Request};
28use crate::server::{Handler, PeerOrigin, ServeStream, SharedWriter, SubRegistry};
29use crate::wire::method;
30use serde_json::{Value, json};
31use std::io::{self, BufRead, BufReader, Read, Write};
32use std::net::{TcpListener, TcpStream};
33use std::sync::atomic::{AtomicU64, Ordering};
34use std::sync::{Arc, Mutex};
35use std::thread;
36use std::time::Duration;
37
38const SSE_KEEPALIVE: Duration = Duration::from_secs(15);
41
42const MAX_BODY: usize = 8 * 1024 * 1024;
44
45const MAX_HEAD_BYTES: usize = 64 * 1024;
51
52const MAX_HEADERS: usize = 100;
57
58#[derive(Default, Clone)]
64pub struct PeerId {
65 pub cert: bool,
67 pub subject: Option<String>,
69 pub sans: Vec<String>,
72}
73
74pub struct RequestParts<'a> {
76 pub headers: &'a [(String, String)],
78 pub peer_cert: bool,
80 pub peer_subject: Option<&'a str>,
82 pub peer_sans: &'a [String],
84}
85
86impl RequestParts<'_> {
87 pub fn header(&self, name: &str) -> Option<&str> {
89 self.headers
90 .iter()
91 .find(|(k, _)| k == name)
92 .map(|(_, v)| v.as_str())
93 }
94}
95
96pub trait HttpAuth: Send + Sync + 'static {
100 fn authenticate(&self, parts: &RequestParts) -> Option<PeerOrigin>;
101}
102
103pub struct AllowAll;
107impl HttpAuth for AllowAll {
108 fn authenticate(&self, _parts: &RequestParts) -> Option<PeerOrigin> {
109 Some(PeerOrigin::Management)
110 }
111}
112
113#[derive(Default, Clone)]
115pub struct ServeOptions {
116 pub extra_origins: Vec<String>,
123}
124
125pub enum HttpAcceptor {
130 Plain,
132 #[cfg(feature = "tls")]
134 Tls(net::tls::TlsAcceptor),
135}
136
137pub fn bind_tcp(addr: &str) -> io::Result<TcpListener> {
141 TcpListener::bind(addr)
142}
143
144#[allow(clippy::too_many_arguments)]
148pub fn spawn_accept_http(
149 listener: TcpListener,
150 acceptor: Arc<HttpAcceptor>,
151 handler: Arc<dyn Handler>,
152 auth: Arc<dyn HttpAuth>,
153 subs: SubRegistry,
154 conn_counter: Arc<AtomicU64>,
155 write_timeout: Duration,
156) -> io::Result<()> {
157 spawn_accept_http_opts(
158 listener,
159 acceptor,
160 handler,
161 auth,
162 subs,
163 conn_counter,
164 write_timeout,
165 ServeOptions::default(),
166 )
167}
168
169#[allow(clippy::too_many_arguments)]
171pub fn spawn_accept_http_opts(
172 listener: TcpListener,
173 acceptor: Arc<HttpAcceptor>,
174 handler: Arc<dyn Handler>,
175 auth: Arc<dyn HttpAuth>,
176 subs: SubRegistry,
177 conn_counter: Arc<AtomicU64>,
178 write_timeout: Duration,
179 opts: ServeOptions,
180) -> io::Result<()> {
181 let opts = Arc::new(opts);
182 thread::Builder::new()
183 .name("serve-http".into())
184 .spawn(move || {
185 for tcp in listener.incoming().flatten() {
186 let acceptor = Arc::clone(&acceptor);
187 let handler = Arc::clone(&handler);
188 let auth = Arc::clone(&auth);
189 let subs = Arc::clone(&subs);
190 let conn_counter = Arc::clone(&conn_counter);
191 let opts = Arc::clone(&opts);
192 thread::Builder::new()
193 .name("serve-http-conn".into())
194 .spawn(move || {
195 accept_and_serve(
196 tcp,
197 &acceptor,
198 &handler,
199 &auth,
200 &subs,
201 &conn_counter,
202 write_timeout,
203 &opts,
204 );
205 })
206 .ok();
207 }
208 })
209 .map(|_| ())
210}
211
212#[allow(clippy::too_many_arguments)]
213fn accept_and_serve(
214 tcp: TcpStream,
215 acceptor: &HttpAcceptor,
216 handler: &Arc<dyn Handler>,
217 auth: &Arc<dyn HttpAuth>,
218 subs: &SubRegistry,
219 conn_counter: &AtomicU64,
220 write_timeout: Duration,
221 opts: &ServeOptions,
222) {
223 let _ = tcp.set_write_timeout(Some(write_timeout));
224 let _ = tcp.set_read_timeout(Some(write_timeout));
225 match acceptor {
226 HttpAcceptor::Plain => {
227 serve_conn(
228 tcp,
229 PeerId::default(),
230 handler,
231 auth,
232 subs,
233 conn_counter,
234 opts,
235 );
236 }
237 #[cfg(feature = "tls")]
239 HttpAcceptor::Tls(tls) => {
240 if let Ok(stream) = tls.accept(tcp) {
241 let peer = peer_id(&stream);
242 serve_conn(stream, peer, handler, auth, subs, conn_counter, opts);
243 }
244 }
245 }
246}
247
248#[cfg(feature = "tls")]
252fn peer_id(stream: &net::tls::ServerTlsStream) -> PeerId {
253 match net::tls::peer_identity(stream) {
254 Some(id) => PeerId {
255 cert: true,
256 subject: id.subject_cn,
257 sans: id.sans,
258 },
259 None => PeerId::default(),
260 }
261}
262
263fn serve_conn<S: Read + Write + Send + 'static>(
266 stream: S,
267 peer: PeerId,
268 handler: &Arc<dyn Handler>,
269 auth: &Arc<dyn HttpAuth>,
270 subs: &SubRegistry,
271 conn_counter: &AtomicU64,
272 opts: &ServeOptions,
273) {
274 let mut reader = BufReader::new(stream);
275 let req = match read_request(&mut reader) {
276 Ok(req) => req,
277 Err(ReadError::HeadTooLarge) => {
281 let _ = write_simple(
282 reader.get_mut(),
283 431,
284 "Request Header Fields Too Large",
285 b"request head exceeds the header size/count limits",
286 None,
287 );
288 return;
289 }
290 Err(ReadError::Incomplete) => return, };
292
293 let cors = match check_origin(&req.headers, &opts.extra_origins) {
303 OriginCheck::NoBrowser => None,
304 OriginCheck::Allowed(o) => Some(o),
305 OriginCheck::Denied => {
306 let _ = write_simple(
307 reader.get_mut(),
308 403,
309 "Forbidden",
310 b"cross-origin request rejected",
311 None,
312 );
313 return;
314 }
315 };
316 if req.method.eq_ignore_ascii_case("OPTIONS") {
320 let _ = write_preflight(reader.get_mut(), cors.as_deref());
321 return;
322 }
323
324 let origin = {
326 let parts = RequestParts {
327 headers: &req.headers,
328 peer_cert: peer.cert,
329 peer_subject: peer.subject.as_deref(),
330 peer_sans: &peer.sans,
331 };
332 auth.authenticate(&parts)
333 };
334 let Some(origin) = origin else {
335 let _ = write_simple(reader.get_mut(), 401, "Unauthorized", b"", cors.as_deref());
336 return;
337 };
338
339 if !req.method.eq_ignore_ascii_case("POST") {
342 let _ = write_simple(
343 reader.get_mut(),
344 405,
345 "Method Not Allowed",
346 b"POST a JSON-RPC request, or POST subscriptions/listen for the SSE stream",
347 cors.as_deref(),
348 );
349 return;
350 }
351
352 let conn = conn_counter.fetch_add(1, Ordering::Relaxed);
353 handler.on_connect(origin, conn);
354
355 let incoming: Result<Incoming, _> = serde_json::from_slice(&req.body);
356 match incoming {
357 Ok(Incoming::Request(rpc_req)) if rpc_req.method == method::SUBSCRIPTIONS_LISTEN => {
358 serve_listen(
359 reader,
360 rpc_req,
361 origin,
362 conn,
363 handler,
364 subs,
365 cors.as_deref(),
366 );
367 }
368 Ok(Incoming::Request(rpc_req)) if handler.streams(&rpc_req.method) => {
371 serve_stream(reader, rpc_req, origin, conn, handler, cors.as_deref());
372 remove_and_disconnect(subs, conn, origin, handler);
373 }
374 Ok(Incoming::Request(rpc_req)) => {
375 serve_unary(
376 reader.get_mut(),
377 rpc_req,
378 origin,
379 conn,
380 handler,
381 cors.as_deref(),
382 );
383 remove_and_disconnect(subs, conn, origin, handler);
384 }
385 Ok(Incoming::Notification(_)) | Ok(Incoming::Response(_)) => {
387 let _ = write_simple(reader.get_mut(), 202, "Accepted", b"", cors.as_deref());
388 remove_and_disconnect(subs, conn, origin, handler);
389 }
390 Err(_) => {
391 let _ = write_simple(
392 reader.get_mut(),
393 400,
394 "Bad Request",
395 b"invalid JSON-RPC frame",
396 cors.as_deref(),
397 );
398 remove_and_disconnect(subs, conn, origin, handler);
399 }
400 }
401}
402
403fn serve_stream<S: Read + Write + Send + 'static>(
409 reader: BufReader<S>,
410 req: Request,
411 origin: PeerOrigin,
412 conn: u64,
413 handler: &Arc<dyn Handler>,
414 cors: Option<&str>,
415) {
416 let mut stream = reader.into_inner();
417 let head = format!(
418 "HTTP/1.1 200 OK\r\nContent-Type: text/event-stream\r\nCache-Control: no-store\r\n{}Connection: close\r\n\r\n",
419 cors_headers(cors)
420 );
421 if stream
422 .write_all(head.as_bytes())
423 .and_then(|_| stream.flush())
424 .is_err()
425 {
426 return;
427 }
428 let writer: SharedWriter = Arc::new(Mutex::new(ServeStream::Http(Box::new(stream))));
429
430 let stop = Arc::new(std::sync::atomic::AtomicBool::new(false));
433 let ka = {
434 let writer = Arc::clone(&writer);
435 let stop = Arc::clone(&stop);
436 thread::spawn(move || {
437 while !stop.load(Ordering::Relaxed) {
438 thread::sleep(SSE_KEEPALIVE);
439 if stop.load(Ordering::Relaxed) {
440 break;
441 }
442 let alive = writer
443 .lock()
444 .map(|mut w| {
445 w.write_all(b": keep-alive\n\n")
446 .and_then(|_| w.flush())
447 .is_ok()
448 })
449 .unwrap_or(false);
450 if !alive {
451 break;
452 }
453 }
454 })
455 };
456
457 let resp = handler.dispatch(req, origin, &writer, conn);
458 stop.store(true, Ordering::Relaxed);
459 if let Ok(mut w) = writer.lock() {
460 let _ = w.write_response(&resp);
461 }
462 let _ = ka.join();
463}
464
465fn serve_unary<S: Write>(
468 stream: &mut S,
469 req: Request,
470 origin: PeerOrigin,
471 conn: u64,
472 handler: &Arc<dyn Handler>,
473 cors: Option<&str>,
474) {
475 let sink: SharedWriter = Arc::new(Mutex::new(ServeStream::Http(Box::new(io::sink()))));
478 let is_initialize = req.method == method::INITIALIZE;
479 let resp = handler.dispatch(req, origin, &sink, conn);
480 let body = serde_json::to_vec(&resp).unwrap_or_default();
481 let session = if is_initialize {
482 format!("Mcp-Session-Id: {}\r\n", next_session_id())
483 } else {
484 String::new()
485 };
486 let head = format!(
487 "HTTP/1.1 200 OK\r\nContent-Type: application/json\r\n{session}{}Content-Length: {}\r\nConnection: close\r\n\r\n",
488 cors_headers(cors),
489 body.len()
490 );
491 let _ = stream.write_all(head.as_bytes());
492 let _ = stream.write_all(&body);
493 let _ = stream.flush();
494}
495
496fn serve_listen<S: Read + Write + Send + 'static>(
502 reader: BufReader<S>,
503 req: Request,
504 origin: PeerOrigin,
505 conn: u64,
506 handler: &Arc<dyn Handler>,
507 subs: &SubRegistry,
508 cors: Option<&str>,
509) {
510 let uris = listen_uris(&req);
511 let mut stream = reader.into_inner();
512 let head = format!(
514 "HTTP/1.1 200 OK\r\nContent-Type: text/event-stream\r\nCache-Control: no-store\r\n{}Connection: close\r\n\r\n",
515 cors_headers(cors)
516 );
517 if stream
518 .write_all(head.as_bytes())
519 .and_then(|_| stream.flush())
520 .is_err()
521 {
522 remove_and_disconnect(subs, conn, origin, handler);
523 return;
524 }
525
526 let writer: SharedWriter = Arc::new(Mutex::new(ServeStream::Http(Box::new(stream))));
530 for uri in &uris {
531 let sub_req = Request::new(0, method::RESOURCES_SUBSCRIBE, Some(json!({ "uri": uri })));
532 let _ = handler.dispatch(sub_req, origin, &writer, conn);
533 }
534
535 loop {
537 thread::sleep(SSE_KEEPALIVE);
538 let alive = writer
539 .lock()
540 .map(|mut w| {
541 w.write_all(b": keep-alive\n\n")
542 .and_then(|_| w.flush())
543 .is_ok()
544 })
545 .unwrap_or(false);
546 if !alive {
547 break;
548 }
549 }
550 remove_and_disconnect(subs, conn, origin, handler);
551}
552
553fn listen_uris(req: &Request) -> Vec<String> {
556 req.params
557 .as_ref()
558 .and_then(|p| p.get("notifications"))
559 .and_then(|n| n.get("resourceSubscriptions"))
560 .and_then(Value::as_array)
561 .map(|a| {
562 a.iter()
563 .filter_map(Value::as_str)
564 .map(str::to_string)
565 .collect()
566 })
567 .unwrap_or_default()
568}
569
570fn remove_and_disconnect(
571 subs: &SubRegistry,
572 conn: u64,
573 origin: PeerOrigin,
574 handler: &Arc<dyn Handler>,
575) {
576 crate::server::remove_conn_subscriptions(subs, conn);
577 handler.on_disconnect(origin, conn);
578}
579
580fn write_simple<S: Write>(
582 stream: &mut S,
583 code: u16,
584 reason: &str,
585 body: &[u8],
586 cors: Option<&str>,
587) -> io::Result<()> {
588 let head = format!(
589 "HTTP/1.1 {code} {reason}\r\nContent-Type: text/plain\r\n{}Content-Length: {}\r\nConnection: close\r\n\r\n",
590 cors_headers(cors),
591 body.len()
592 );
593 stream.write_all(head.as_bytes())?;
594 stream.write_all(body)?;
595 stream.flush()
596}
597
598fn cors_headers(origin: Option<&str>) -> String {
602 match origin {
603 Some(o) => format!(
604 "Access-Control-Allow-Origin: {o}\r\nVary: Origin\r\nAccess-Control-Expose-Headers: Mcp-Session-Id\r\n"
605 ),
606 None => String::new(),
607 }
608}
609
610fn write_preflight<S: Write>(stream: &mut S, origin: Option<&str>) -> io::Result<()> {
613 let grant = match origin {
614 Some(o) => format!(
615 "Access-Control-Allow-Origin: {o}\r\nVary: Origin\r\nAccess-Control-Allow-Methods: POST, OPTIONS\r\nAccess-Control-Allow-Headers: content-type, authorization, last-event-id, mcp-session-id\r\nAccess-Control-Max-Age: 600\r\n"
616 ),
617 None => String::new(),
618 };
619 let head =
620 format!("HTTP/1.1 204 No Content\r\n{grant}Content-Length: 0\r\nConnection: close\r\n\r\n");
621 stream.write_all(head.as_bytes())?;
622 stream.flush()
623}
624
625pub struct RawRequest {
634 pub method: String,
635 pub target: String,
636 pub headers: Vec<(String, String)>,
637 pub body: Vec<u8>,
638 pub peer_cert: bool,
640 pub peer_subject: Option<String>,
642 pub peer_sans: Vec<String>,
644}
645
646impl RawRequest {
647 pub fn header(&self, name: &str) -> Option<&str> {
649 self.headers
650 .iter()
651 .find(|(k, _)| k == name)
652 .map(|(_, v)| v.as_str())
653 }
654 pub fn path(&self) -> &str {
656 self.target.split('?').next().unwrap_or(&self.target)
657 }
658}
659
660pub struct RawResponse {
662 pub status: u16,
663 pub reason: &'static str,
664 pub content_type: &'static str,
665 pub body: Vec<u8>,
666 pub headers: Vec<(&'static str, String)>,
668}
669
670impl RawResponse {
671 pub fn json(status: u16, reason: &'static str, body: impl Into<Vec<u8>>) -> RawResponse {
673 RawResponse {
674 status,
675 reason,
676 content_type: "application/json",
677 body: body.into(),
678 headers: Vec::new(),
679 }
680 }
681 pub fn text(status: u16, reason: &'static str, body: impl Into<Vec<u8>>) -> RawResponse {
683 RawResponse {
684 status,
685 reason,
686 content_type: "text/plain",
687 body: body.into(),
688 headers: Vec::new(),
689 }
690 }
691}
692
693pub trait RawHandler: Send + Sync + 'static {
696 fn handle(&self, req: &RawRequest) -> RawResponse;
697}
698
699pub fn spawn_accept_raw(
702 listener: TcpListener,
703 acceptor: Arc<HttpAcceptor>,
704 handler: Arc<dyn RawHandler>,
705 write_timeout: Duration,
706) -> io::Result<()> {
707 thread::Builder::new()
708 .name("serve-webhook".into())
709 .spawn(move || {
710 for tcp in listener.incoming().flatten() {
711 let acceptor = Arc::clone(&acceptor);
712 let handler = Arc::clone(&handler);
713 thread::Builder::new()
714 .name("webhook-conn".into())
715 .spawn(move || {
716 let _ = tcp.set_write_timeout(Some(write_timeout));
717 let _ = tcp.set_read_timeout(Some(write_timeout));
718 match &*acceptor {
719 HttpAcceptor::Plain => serve_conn_raw(tcp, PeerId::default(), &handler),
720 #[cfg(feature = "tls")]
721 HttpAcceptor::Tls(tls) => {
722 if let Ok(stream) = tls.accept(tcp) {
723 let peer = peer_id(&stream);
724 serve_conn_raw(stream, peer, &handler);
725 }
726 }
727 }
728 })
729 .ok();
730 }
731 })
732 .map(|_| ())
733}
734
735fn serve_conn_raw<S: Read + Write + Send + 'static>(
736 stream: S,
737 peer: PeerId,
738 handler: &Arc<dyn RawHandler>,
739) {
740 let mut reader = BufReader::new(stream);
741 let req = match read_request(&mut reader) {
742 Ok(req) => req,
743 Err(ReadError::HeadTooLarge) => {
744 let _ = write_simple(
745 reader.get_mut(),
746 431,
747 "Request Header Fields Too Large",
748 b"request head exceeds the header size/count limits",
749 None,
750 );
751 return;
752 }
753 Err(ReadError::Incomplete) => return,
754 };
755 if matches!(check_origin(&req.headers, &[]), OriginCheck::Denied) {
757 let _ = write_simple(
758 reader.get_mut(),
759 403,
760 "Forbidden",
761 b"cross-origin request rejected",
762 None,
763 );
764 return;
765 }
766 let raw = RawRequest {
767 method: req.method,
768 target: req.target,
769 headers: req.headers,
770 body: req.body,
771 peer_cert: peer.cert,
772 peer_subject: peer.subject,
773 peer_sans: peer.sans,
774 };
775 let resp = handler.handle(&raw);
776 let _ = write_raw(reader.get_mut(), &resp);
777}
778
779fn write_raw<S: Write>(stream: &mut S, resp: &RawResponse) -> io::Result<()> {
780 let mut head = format!(
781 "HTTP/1.1 {} {}\r\nContent-Type: {}\r\nContent-Length: {}\r\nConnection: close\r\n",
782 resp.status,
783 resp.reason,
784 resp.content_type,
785 resp.body.len()
786 );
787 for (name, value) in &resp.headers {
788 head.push_str(name);
789 head.push_str(": ");
790 head.push_str(value);
791 head.push_str("\r\n");
792 }
793 head.push_str("\r\n");
794 stream.write_all(head.as_bytes())?;
795 stream.write_all(&resp.body)?;
796 stream.flush()
797}
798
799struct HttpRequest {
801 method: String,
802 #[allow(dead_code)]
803 target: String,
804 headers: Vec<(String, String)>,
805 body: Vec<u8>,
806}
807
808enum OriginCheck {
810 NoBrowser,
812 Allowed(String),
814 Denied,
816}
817
818fn check_origin(headers: &[(String, String)], extra: &[String]) -> OriginCheck {
822 match headers.iter().find(|(k, _)| k == "origin") {
823 None => OriginCheck::NoBrowser,
824 Some((_, origin)) => {
825 if origin_is_loopback(origin) || extra.iter().any(|e| e == origin) {
826 OriginCheck::Allowed(origin.clone())
827 } else {
828 OriginCheck::Denied
829 }
830 }
831 }
832}
833
834fn origin_is_loopback(origin: &str) -> bool {
837 let after_scheme = origin.split_once("://").map(|(_, r)| r).unwrap_or(origin);
838 let authority = after_scheme.split('/').next().unwrap_or(after_scheme);
839 let host = if let Some(v6) = authority.strip_prefix('[') {
841 v6.split(']').next().unwrap_or(v6)
842 } else {
843 authority.split(':').next().unwrap_or(authority)
844 };
845 host == "localhost" || host == "::1" || host.starts_with("127.")
846}
847
848fn next_session_id() -> String {
854 static SEQ: AtomicU64 = AtomicU64::new(0);
855 let n = SEQ.fetch_add(1, Ordering::Relaxed);
856 let millis = std::time::SystemTime::now()
857 .duration_since(std::time::UNIX_EPOCH)
858 .map(|d| d.as_millis())
859 .unwrap_or(0);
860 format!("s-{millis:x}-{n:x}")
861}
862
863enum ReadError {
865 Incomplete,
868 HeadTooLarge,
871}
872
873fn read_head_line<S: Read>(
879 reader: &mut BufReader<S>,
880 budget: &mut usize,
881 line: &mut String,
882) -> Result<usize, ReadError> {
883 let n = Read::take(&mut *reader, *budget as u64 + 1)
886 .read_line(line)
887 .map_err(|_| ReadError::Incomplete)?;
888 if n > *budget {
889 return Err(ReadError::HeadTooLarge);
890 }
891 *budget -= n;
892 Ok(n)
893}
894
895fn read_request<S: Read>(reader: &mut BufReader<S>) -> Result<HttpRequest, ReadError> {
898 let mut budget = MAX_HEAD_BYTES;
899 let mut request_line = String::new();
900 if read_head_line(reader, &mut budget, &mut request_line)? == 0 {
901 return Err(ReadError::Incomplete);
902 }
903 let mut parts = request_line.split_whitespace();
904 let method = parts.next().ok_or(ReadError::Incomplete)?.to_string();
905 let target = parts.next().ok_or(ReadError::Incomplete)?.to_string();
906
907 let mut headers = Vec::new();
908 let mut content_length = 0usize;
909 loop {
910 let mut line = String::new();
911 if read_head_line(reader, &mut budget, &mut line)? == 0 {
912 break;
913 }
914 let line = line.trim_end();
915 if line.is_empty() {
916 break; }
918 if let Some((k, v)) = line.split_once(':') {
919 if headers.len() >= MAX_HEADERS {
920 return Err(ReadError::HeadTooLarge);
921 }
922 let name = k.trim().to_ascii_lowercase();
923 let value = v.trim().to_string();
924 if name == "content-length" {
925 content_length = value.parse().unwrap_or(0);
926 }
927 headers.push((name, value));
928 }
929 }
930 if content_length > MAX_BODY {
931 return Err(ReadError::Incomplete);
932 }
933 let mut body = vec![0u8; content_length];
934 if content_length > 0 {
935 reader
936 .read_exact(&mut body)
937 .map_err(|_| ReadError::Incomplete)?;
938 }
939 Ok(HttpRequest {
940 method,
941 target,
942 headers,
943 body,
944 })
945}
946
947#[cfg(test)]
948mod tests {
949 use super::*;
950 use crate::rpc::{self, Response};
951 use crate::server::{notify_resource_updated_keep, register_subscriber};
952 use std::io::BufRead;
953
954 struct TestHandler {
958 subs: SubRegistry,
959 }
960 impl Handler for TestHandler {
961 fn dispatch(
962 &self,
963 req: Request,
964 _origin: PeerOrigin,
965 writer: &SharedWriter,
966 conn: u64,
967 ) -> Response {
968 if let Some(resp) = crate::server::lifecycle_response(
969 &req,
970 &json!({"name": "test", "version": "1"}),
971 &json!({"tools": {}, "resources": {"subscribe": true}}),
972 ) {
973 return resp;
974 }
975 match req.method.as_str() {
976 "tools/call" => Response::ok(req.id, json!({"ok": true})),
977 "resources/subscribe" => {
978 let uri = req
979 .params
980 .as_ref()
981 .and_then(|p| p["uri"].as_str())
982 .unwrap_or("");
983 if uri == "res://ok" {
985 register_subscriber(&self.subs, uri, conn, writer);
986 Response::ok(req.id, json!({}))
987 } else {
988 Response::err(req.id, rpc::RESOURCE_NOT_FOUND, "no")
989 }
990 }
991 _ => Response::err(req.id, rpc::METHOD_NOT_FOUND, "unknown"),
992 }
993 }
994 }
995
996 fn http_post(addr: &str, body: &str) -> (Vec<(String, String)>, String) {
997 let mut s = TcpStream::connect(addr).unwrap();
998 let req = format!(
999 "POST /mcp HTTP/1.1\r\nHost: x\r\nContent-Type: application/json\r\nContent-Length: {}\r\nConnection: close\r\n\r\n{body}",
1000 body.len()
1001 );
1002 s.write_all(req.as_bytes()).unwrap();
1003 s.set_read_timeout(Some(Duration::from_secs(5))).ok();
1004 let mut reader = BufReader::new(s);
1005 let mut status = String::new();
1006 reader.read_line(&mut status).unwrap();
1007 let mut headers = Vec::new();
1008 loop {
1009 let mut l = String::new();
1010 reader.read_line(&mut l).unwrap();
1011 if l.trim().is_empty() {
1012 break;
1013 }
1014 if let Some((k, v)) = l.split_once(':') {
1015 headers.push((k.trim().to_ascii_lowercase(), v.trim().to_string()));
1016 }
1017 }
1018 let mut body = String::new();
1019 reader.read_to_string(&mut body).unwrap();
1020 (headers, body)
1021 }
1022
1023 fn spawn_server() -> (String, SubRegistry) {
1024 let subs: SubRegistry = Arc::new(Mutex::new(std::collections::HashMap::new()));
1025 let listener = bind_tcp("127.0.0.1:0").unwrap();
1026 let addr = listener.local_addr().unwrap().to_string();
1027 spawn_accept_http(
1028 listener,
1029 Arc::new(HttpAcceptor::Plain),
1030 Arc::new(TestHandler {
1031 subs: Arc::clone(&subs),
1032 }),
1033 Arc::new(AllowAll),
1034 Arc::clone(&subs),
1035 Arc::new(AtomicU64::new(0)),
1036 Duration::from_secs(5),
1037 )
1038 .unwrap();
1039 (addr, subs)
1040 }
1041
1042 #[test]
1043 fn unary_post_returns_application_json() {
1044 let (addr, _subs) = spawn_server();
1045 let (headers, body) = http_post(
1046 &addr,
1047 r#"{"jsonrpc":"2.0","id":1,"method":"tools/call","params":{"name":"x"}}"#,
1048 );
1049 assert!(
1050 headers
1051 .iter()
1052 .any(|(k, v)| k == "content-type" && v.contains("application/json")),
1053 "headers: {headers:?}"
1054 );
1055 let v: Value = serde_json::from_str(&body).unwrap();
1056 assert_eq!(v["result"]["ok"], true);
1057 }
1058
1059 fn http_post_origin(addr: &str, origin: Option<&str>, body: &str) -> u16 {
1061 let mut s = TcpStream::connect(addr).unwrap();
1062 let origin_line = origin
1063 .map(|o| format!("Origin: {o}\r\n"))
1064 .unwrap_or_default();
1065 let req = format!(
1066 "POST /mcp HTTP/1.1\r\nHost: x\r\n{origin_line}Content-Type: application/json\r\nContent-Length: {}\r\nConnection: close\r\n\r\n{body}",
1067 body.len()
1068 );
1069 s.write_all(req.as_bytes()).unwrap();
1070 s.set_read_timeout(Some(Duration::from_secs(5))).ok();
1071 let mut status = String::new();
1072 BufReader::new(s).read_line(&mut status).unwrap();
1073 status
1074 .split_whitespace()
1075 .nth(1)
1076 .and_then(|c| c.parse().ok())
1077 .unwrap_or(0)
1078 }
1079
1080 #[test]
1081 fn initialize_stamps_a_unique_session_header() {
1082 let (addr, _subs) = spawn_server();
1083 let init = r#"{"jsonrpc":"2.0","id":1,"method":"initialize","params":{"protocolVersion":"2025-11-25"}}"#;
1084 let sid = |addr: &str| {
1085 let (headers, _) = http_post(addr, init);
1086 headers
1087 .into_iter()
1088 .find(|(k, _)| k == "mcp-session-id")
1089 .map(|(_, v)| v)
1090 };
1091 let a = sid(&addr).expect("initialize stamps a session id");
1092 let b = sid(&addr).expect("second initialize stamps a session id");
1093 assert_ne!(a, "srv", "the session id is not the old constant");
1094 assert_ne!(a, b, "each initialize mints a distinct session id");
1095 }
1096
1097 #[test]
1098 fn a_cross_origin_request_is_rejected_403() {
1099 let (addr, _subs) = spawn_server();
1100 let call = r#"{"jsonrpc":"2.0","id":1,"method":"tools/call","params":{"name":"x"}}"#;
1101 assert_eq!(
1103 http_post_origin(&addr, Some("https://evil.example"), call),
1104 403
1105 );
1106 assert_eq!(http_post_origin(&addr, None, call), 200);
1108 assert_eq!(
1110 http_post_origin(&addr, Some("http://localhost:3000"), call),
1111 200
1112 );
1113 assert_eq!(http_post_origin(&addr, Some("http://127.0.0.1"), call), 200);
1114 }
1115
1116 fn spawn_server_with_origins(extra: &[&str]) -> String {
1118 let subs: SubRegistry = Arc::new(Mutex::new(std::collections::HashMap::new()));
1119 let listener = bind_tcp("127.0.0.1:0").unwrap();
1120 let addr = listener.local_addr().unwrap().to_string();
1121 spawn_accept_http_opts(
1122 listener,
1123 Arc::new(HttpAcceptor::Plain),
1124 Arc::new(TestHandler {
1125 subs: Arc::clone(&subs),
1126 }),
1127 Arc::new(AllowAll),
1128 Arc::clone(&subs),
1129 Arc::new(AtomicU64::new(0)),
1130 Duration::from_secs(5),
1131 ServeOptions {
1132 extra_origins: extra.iter().map(|s| s.to_string()).collect(),
1133 },
1134 )
1135 .unwrap();
1136 addr
1137 }
1138
1139 fn http_raw(addr: &str, req: &str) -> (u16, Vec<(String, String)>) {
1141 let mut s = TcpStream::connect(addr).unwrap();
1142 s.write_all(req.as_bytes()).unwrap();
1143 s.set_read_timeout(Some(Duration::from_secs(5))).ok();
1144 let mut reader = BufReader::new(s);
1145 let mut status = String::new();
1146 reader.read_line(&mut status).unwrap();
1147 let code = status
1148 .split_whitespace()
1149 .nth(1)
1150 .and_then(|c| c.parse().ok())
1151 .unwrap_or(0);
1152 let mut headers = Vec::new();
1153 loop {
1154 let mut l = String::new();
1155 if reader.read_line(&mut l).unwrap_or(0) == 0 {
1156 break;
1157 }
1158 if l.trim().is_empty() {
1159 break;
1160 }
1161 if let Some((k, v)) = l.split_once(':') {
1162 headers.push((k.trim().to_ascii_lowercase(), v.trim().to_string()));
1163 }
1164 }
1165 (code, headers)
1166 }
1167
1168 #[test]
1169 fn a_configured_extra_origin_is_served_with_cors_headers() {
1170 let addr = spawn_server_with_origins(&["https://ui.example"]);
1171 let call = r#"{"jsonrpc":"2.0","id":1,"method":"tools/call","params":{"name":"x"}}"#;
1172 let req = format!(
1174 "POST / HTTP/1.1\r\nHost: x\r\nOrigin: https://ui.example\r\nContent-Length: {}\r\nConnection: close\r\n\r\n{call}",
1175 call.len()
1176 );
1177 let (code, headers) = http_raw(&addr, &req);
1178 assert_eq!(code, 200);
1179 assert!(
1180 headers
1181 .iter()
1182 .any(|(k, v)| { k == "access-control-allow-origin" && v == "https://ui.example" }),
1183 "CORS echo missing: {headers:?}"
1184 );
1185 let bad = format!(
1187 "POST / HTTP/1.1\r\nHost: x\r\nOrigin: https://evil.example\r\nContent-Length: {}\r\nConnection: close\r\n\r\n{call}",
1188 call.len()
1189 );
1190 assert_eq!(http_raw(&addr, &bad).0, 403);
1191 let local = format!(
1193 "POST / HTTP/1.1\r\nHost: x\r\nOrigin: http://127.0.0.1:5173\r\nContent-Length: {}\r\nConnection: close\r\n\r\n{call}",
1194 call.len()
1195 );
1196 let (code, headers) = http_raw(&addr, &local);
1197 assert_eq!(code, 200);
1198 assert!(
1199 headers
1200 .iter()
1201 .any(|(k, v)| k == "access-control-allow-origin" && v == "http://127.0.0.1:5173")
1202 );
1203 }
1204
1205 #[test]
1206 fn a_cors_preflight_options_is_answered_before_auth() {
1207 let addr = spawn_server_with_origins(&["https://ui.example"]);
1208 let req = "OPTIONS / HTTP/1.1\r\nHost: x\r\nOrigin: https://ui.example\r\nAccess-Control-Request-Method: POST\r\nContent-Length: 0\r\nConnection: close\r\n\r\n";
1209 let (code, headers) = http_raw(&addr, req);
1210 assert_eq!(code, 204);
1211 assert!(
1212 headers
1213 .iter()
1214 .any(|(k, v)| k == "access-control-allow-origin" && v == "https://ui.example")
1215 );
1216 assert!(
1217 headers
1218 .iter()
1219 .any(|(k, v)| k == "access-control-allow-headers" && v.contains("authorization"))
1220 );
1221 let bad = "OPTIONS / HTTP/1.1\r\nHost: x\r\nOrigin: https://evil.example\r\nContent-Length: 0\r\nConnection: close\r\n\r\n";
1223 assert_eq!(http_raw(&addr, bad).0, 403);
1224 }
1225
1226 #[test]
1227 fn origin_loopback_classification() {
1228 assert!(origin_is_loopback("http://localhost"));
1229 assert!(origin_is_loopback("http://localhost:8080"));
1230 assert!(origin_is_loopback("https://127.0.0.1:443"));
1231 assert!(origin_is_loopback("http://[::1]:9000"));
1232 assert!(!origin_is_loopback("https://evil.example"));
1233 assert!(!origin_is_loopback("http://169.254.1.1")); assert!(!origin_is_loopback("null")); }
1236
1237 #[test]
1238 fn subscriptions_listen_streams_a_pushed_update_as_sse() {
1239 let (addr, subs) = spawn_server();
1240 let addr2 = addr.clone();
1242 let got = Arc::new(Mutex::new(String::new()));
1243 let got2 = Arc::clone(&got);
1244 thread::spawn(move || {
1245 let mut s = TcpStream::connect(&addr2).unwrap();
1246 let body = r#"{"jsonrpc":"2.0","id":1,"method":"subscriptions/listen","params":{"notifications":{"resourceSubscriptions":["res://ok"]}}}"#;
1247 let req = format!(
1248 "POST /mcp HTTP/1.1\r\nHost: x\r\nAccept: text/event-stream\r\nContent-Length: {}\r\nConnection: close\r\n\r\n{body}",
1249 body.len()
1250 );
1251 s.write_all(req.as_bytes()).unwrap();
1252 s.set_read_timeout(Some(Duration::from_secs(5))).ok();
1253 let mut reader = BufReader::new(s);
1254 let mut line = String::new();
1255 for _ in 0..50 {
1257 line.clear();
1258 if reader.read_line(&mut line).unwrap_or(0) == 0 {
1259 break;
1260 }
1261 if line.starts_with("data:") {
1262 *got2.lock().unwrap() = line.clone();
1263 break;
1264 }
1265 }
1266 });
1267
1268 let deadline = std::time::Instant::now() + Duration::from_secs(3);
1270 while std::time::Instant::now() < deadline {
1271 if subs.lock().unwrap().contains_key("res://ok") {
1272 break;
1273 }
1274 thread::sleep(Duration::from_millis(20));
1275 }
1276 notify_resource_updated_keep(&subs, "res://ok");
1277
1278 let deadline = std::time::Instant::now() + Duration::from_secs(3);
1279 loop {
1280 if got.lock().unwrap().starts_with("data:") {
1281 break;
1282 }
1283 assert!(std::time::Instant::now() < deadline, "no SSE push observed");
1284 thread::sleep(Duration::from_millis(20));
1285 }
1286 let data = got.lock().unwrap().clone();
1287 assert!(data.contains("notifications/resources/updated"), "{data}");
1288 assert!(data.contains("res://ok"), "{data}");
1289 }
1290}