1use std::collections::HashMap;
45use std::sync::Arc;
46use std::sync::atomic::{AtomicBool, Ordering};
47use std::time::Duration;
48
49use async_trait::async_trait;
50use base64::Engine;
51use tokio::sync::{Notify, RwLock, mpsc};
52use tokio::task::JoinHandle;
53
54use super::transport::ClientTransport;
55use crate::error::{Error, Result};
56use crate::protocol::{RequestId, notifications};
57
58const MCP_METHOD_HEADER: &str = "mcp-method";
59const MCP_NAME_HEADER: &str = "mcp-name";
60const MCP_PARAM_HEADER_PREFIX: &str = "mcp-param-";
61const BASE64_SENTINEL_PREFIX: &str = "=?base64?";
62const BASE64_SENTINEL_SUFFIX: &str = "?=";
63
64#[derive(Debug, Clone)]
65struct CustomHeaderMapping {
66 suffix: String,
67 property_path: Vec<String>,
68}
69
70#[cfg(feature = "oauth-client")]
71#[derive(Clone)]
72struct ScopeEscalationRuntime {
73 handler: Arc<dyn OAuthScopeEscalationHandler>,
74 state: Arc<tokio::sync::Mutex<ScopeEscalationState>>,
75 max_attempts: usize,
76}
77
78#[cfg(feature = "oauth-client")]
79struct ScopeEscalationState {
80 scopes: Vec<String>,
81 revision: usize,
82}
83
84#[cfg(feature = "oauth-client")]
85impl ScopeEscalationRuntime {
86 fn new<P>(handler: Arc<P>, config: OAuthScopeEscalationConfig) -> Self
87 where
88 P: OAuthScopeEscalationHandler,
89 {
90 Self {
91 handler,
92 state: Arc::new(tokio::sync::Mutex::new(ScopeEscalationState {
93 scopes: config.initial_scopes().to_vec(),
94 revision: 0,
95 })),
96 max_attempts: config.maximum_attempts(),
97 }
98 }
99
100 async fn respond_to_challenge(
101 &self,
102 challenge: OAuthScopeChallenge,
103 resource: &str,
104 operation: &str,
105 attempt: usize,
106 observed_revision: usize,
107 ) -> std::result::Result<ScopeEscalationDecision, OAuthClientError> {
108 let mut state = self.state.lock().await;
112 let previous_scopes = state.scopes.clone();
113 let mut requested_scopes = previous_scopes.clone();
114 for scope in &challenge.required_scopes {
115 if !requested_scopes.contains(scope) {
116 requested_scopes.push(scope.clone());
117 }
118 }
119
120 if requested_scopes == previous_scopes && state.revision > observed_revision {
121 return Ok(ScopeEscalationDecision {
122 revision: state.revision,
123 });
124 }
125 self.handler
130 .reauthorize(OAuthScopeEscalationRequest {
131 resource: resource.to_string(),
132 operation: operation.to_string(),
133 challenge,
134 previous_scopes,
135 requested_scopes: requested_scopes.clone(),
136 attempt,
137 })
138 .await?;
139
140 state.scopes = requested_scopes;
141 state.revision += 1;
142 Ok(ScopeEscalationDecision {
143 revision: state.revision,
144 })
145 }
146}
147
148#[cfg(feature = "oauth-client")]
149struct ScopeEscalationDecision {
150 revision: usize,
151}
152
153#[cfg(feature = "oauth-client")]
154use super::oauth::{
155 OAuthClientError, OAuthScopeChallenge, OAuthScopeEscalationConfig, OAuthScopeEscalationHandler,
156 OAuthScopeEscalationRequest, TokenProvider,
157};
158
159#[derive(Debug, Clone)]
174pub struct HttpClientConfig {
175 pub headers: HashMap<String, String>,
177 pub auto_sse: bool,
182 pub channel_capacity: usize,
185 pub request_timeout: Duration,
188 pub notification_timeout: Duration,
198 pub sse_reconnect: bool,
201 pub sse_reconnect_delay: Duration,
204 pub max_sse_reconnect_attempts: u32,
207 pub session_recovery: bool,
212 pub max_sse_event_size: usize,
220}
221
222pub const DEFAULT_MAX_SSE_EVENT_SIZE: usize = 16 * 1024 * 1024;
225
226impl Default for HttpClientConfig {
227 fn default() -> Self {
228 Self {
229 headers: HashMap::new(),
230 auto_sse: true,
231 channel_capacity: 256,
232 request_timeout: Duration::from_secs(30),
233 notification_timeout: Duration::from_secs(5),
234 sse_reconnect: true,
235 sse_reconnect_delay: Duration::from_secs(1),
236 max_sse_reconnect_attempts: 5,
237 session_recovery: true,
238 max_sse_event_size: DEFAULT_MAX_SSE_EVENT_SIZE,
239 }
240 }
241}
242
243impl HttpClientConfig {
244 pub fn bearer_token(mut self, token: impl Into<String>) -> Self {
246 self.headers.insert(
247 "Authorization".to_string(),
248 format!("Bearer {}", token.into()),
249 );
250 self
251 }
252
253 pub fn api_key_header(mut self, name: impl Into<String>, key: impl Into<String>) -> Self {
255 self.headers.insert(name.into(), key.into());
256 self
257 }
258
259 pub fn basic_auth(mut self, username: impl AsRef<str>, password: impl AsRef<str>) -> Self {
261 use base64::Engine;
262 let encoded = base64::engine::general_purpose::STANDARD.encode(format!(
263 "{}:{}",
264 username.as_ref(),
265 password.as_ref()
266 ));
267 self.headers
268 .insert("Authorization".to_string(), format!("Basic {}", encoded));
269 self
270 }
271
272 pub fn header(mut self, name: impl Into<String>, value: impl Into<String>) -> Self {
274 self.headers.insert(name.into(), value.into());
275 self
276 }
277}
278
279pub struct HttpClientTransport {
316 url: String,
318 client: reqwest::Client,
320 session_id: Option<String>,
322 protocol_version: Option<String>,
324 tool_header_mappings: HashMap<String, Vec<CustomHeaderMapping>>,
326 incoming_rx: mpsc::Receiver<String>,
328 incoming_tx: mpsc::Sender<String>,
331 sse_task: Option<JoinHandle<()>>,
333 request_tasks: HashMap<RequestId, JoinHandle<()>>,
339 last_event_id: Arc<RwLock<Option<String>>>,
341 sse_retry_delay: Arc<RwLock<Option<Duration>>>,
343 sse_reconnect_signal: Arc<Notify>,
345 connected: Arc<AtomicBool>,
347 config: HttpClientConfig,
349 #[cfg(feature = "oauth-client")]
351 token_provider: Option<Arc<dyn TokenProvider>>,
352 #[cfg(feature = "oauth-client")]
354 scope_escalation: Option<ScopeEscalationRuntime>,
355}
356
357impl HttpClientTransport {
358 pub fn new(url: impl Into<String>) -> Self {
371 Self::with_config(url, HttpClientConfig::default())
372 }
373
374 pub fn with_config(url: impl Into<String>, config: HttpClientConfig) -> Self {
390 let (tx, rx) = mpsc::channel(config.channel_capacity);
391 Self {
392 url: url.into(),
393 client: reqwest::Client::new(),
394 session_id: None,
395 protocol_version: None,
396 tool_header_mappings: HashMap::new(),
397 incoming_rx: rx,
398 incoming_tx: tx,
399 sse_task: None,
400 request_tasks: HashMap::new(),
401 last_event_id: Arc::new(RwLock::new(None)),
402 sse_retry_delay: Arc::new(RwLock::new(None)),
403 sse_reconnect_signal: Arc::new(Notify::new()),
404 connected: Arc::new(AtomicBool::new(true)),
405 config,
406 #[cfg(feature = "oauth-client")]
407 token_provider: None,
408 #[cfg(feature = "oauth-client")]
409 scope_escalation: None,
410 }
411 }
412
413 pub fn with_client(url: impl Into<String>, client: reqwest::Client) -> Self {
430 let config = HttpClientConfig::default();
431 let (tx, rx) = mpsc::channel(config.channel_capacity);
432 Self {
433 url: url.into(),
434 client,
435 session_id: None,
436 protocol_version: None,
437 tool_header_mappings: HashMap::new(),
438 incoming_rx: rx,
439 incoming_tx: tx,
440 sse_task: None,
441 request_tasks: HashMap::new(),
442 last_event_id: Arc::new(RwLock::new(None)),
443 sse_retry_delay: Arc::new(RwLock::new(None)),
444 sse_reconnect_signal: Arc::new(Notify::new()),
445 connected: Arc::new(AtomicBool::new(true)),
446 config,
447 #[cfg(feature = "oauth-client")]
448 token_provider: None,
449 #[cfg(feature = "oauth-client")]
450 scope_escalation: None,
451 }
452 }
453
454 pub fn bearer_token(mut self, token: impl Into<String>) -> Self {
467 self.config.headers.insert(
468 "Authorization".to_string(),
469 format!("Bearer {}", token.into()),
470 );
471 self
472 }
473
474 pub fn api_key(self, key: impl Into<String>) -> Self {
488 self.bearer_token(key)
489 }
490
491 pub fn api_key_header(mut self, name: impl Into<String>, key: impl Into<String>) -> Self {
504 self.config.headers.insert(name.into(), key.into());
505 self
506 }
507
508 pub fn basic_auth(mut self, username: impl AsRef<str>, password: impl AsRef<str>) -> Self {
522 use base64::Engine;
523 let encoded = base64::engine::general_purpose::STANDARD.encode(format!(
524 "{}:{}",
525 username.as_ref(),
526 password.as_ref()
527 ));
528 self.config
529 .headers
530 .insert("Authorization".to_string(), format!("Basic {}", encoded));
531 self
532 }
533
534 pub fn header(mut self, name: impl Into<String>, value: impl Into<String>) -> Self {
548 self.config.headers.insert(name.into(), value.into());
549 self
550 }
551
552 pub fn disable_session_recovery(mut self) -> Self {
559 self.config.session_recovery = false;
560 self
561 }
562
563 #[cfg(feature = "oauth-client")]
593 pub fn with_token_provider(mut self, provider: impl TokenProvider) -> Self {
594 self.token_provider = Some(Arc::new(provider));
595 self.scope_escalation = None;
596 self
597 }
598
599 #[cfg(feature = "oauth-client")]
612 pub fn with_scope_aware_token_provider<P>(
613 mut self,
614 provider: P,
615 config: OAuthScopeEscalationConfig,
616 ) -> Self
617 where
618 P: TokenProvider + OAuthScopeEscalationHandler,
619 {
620 let provider = Arc::new(provider);
621 self.token_provider = Some(provider.clone());
622 self.scope_escalation = Some(ScopeEscalationRuntime::new(provider, config));
623 self
624 }
625
626 fn outgoing_custom_headers(&self, parsed: &serde_json::Value) -> Vec<(String, String)> {
627 if parsed.get("method").and_then(serde_json::Value::as_str) != Some("tools/call") {
628 return Vec::new();
629 }
630 let Some(params) = parsed.get("params") else {
631 return Vec::new();
632 };
633 let Some(name) = params.get("name").and_then(serde_json::Value::as_str) else {
634 return Vec::new();
635 };
636 let Some(mappings) = self.tool_header_mappings.get(name) else {
637 return Vec::new();
638 };
639 let arguments = params.get("arguments").unwrap_or(&serde_json::Value::Null);
640
641 mappings
642 .iter()
643 .filter_map(|mapping| {
644 let value = value_at_property_path(arguments, &mapping.property_path)?;
645 if value.is_null() {
646 return None;
647 }
648 let rendered = json_value_to_header_string(value)?;
649 Some((
650 format!("{MCP_PARAM_HEADER_PREFIX}{}", mapping.suffix),
651 encode_header_value(&rendered),
652 ))
653 })
654 .collect()
655 }
656
657 fn normalize_incoming_message(&mut self, message: String) -> String {
658 if self.protocol_version.as_deref() != Some(crate::protocol::PROTOCOL_VERSION_2026_07_28) {
659 return message;
660 }
661 let Ok(mut parsed) = serde_json::from_str::<serde_json::Value>(&message) else {
662 return message;
663 };
664 let Some(tools) = parsed
665 .get_mut("result")
666 .and_then(|result| result.get_mut("tools"))
667 .and_then(serde_json::Value::as_array_mut)
668 else {
669 return message;
670 };
671
672 self.tool_header_mappings.clear();
673 tools.retain(|tool| {
674 let Some(name) = tool.get("name").and_then(serde_json::Value::as_str) else {
675 return false;
676 };
677 let Some(schema) = tool.get("inputSchema") else {
678 return false;
679 };
680 match custom_header_mappings(schema) {
681 Ok(mappings) => {
682 self.tool_header_mappings.insert(name.to_string(), mappings);
683 true
684 }
685 Err(error) => {
686 tracing::warn!(tool = %name, %error, "Excluding tool with invalid x-mcp-header annotations");
687 false
688 }
689 }
690 });
691
692 parsed.to_string()
693 }
694
695 fn start_sse_stream(&mut self) {
697 let url = self.url.clone();
698 let client = self.client.clone();
699 let session_id = self.session_id.clone().unwrap();
700 let protocol_version = self.protocol_version.clone();
701 let tx = self.incoming_tx.clone();
702 let last_event_id = self.last_event_id.clone();
703 let sse_retry_delay = self.sse_retry_delay.clone();
704 let reconnect_signal = self.sse_reconnect_signal.clone();
705 let connected = self.connected.clone();
706 let config = self.config.clone();
707 #[cfg(feature = "oauth-client")]
708 let token_provider = self.token_provider.clone();
709
710 self.sse_task = Some(tokio::spawn(async move {
711 sse_stream_loop(SseLoopParams {
712 url,
713 client,
714 session_id,
715 protocol_version,
716 tx,
717 last_event_id,
718 sse_retry_delay,
719 reconnect_signal,
720 connected,
721 config,
722 #[cfg(feature = "oauth-client")]
723 token_provider,
724 })
725 .await;
726 }));
727 }
728}
729
730fn custom_header_mappings(
731 schema: &serde_json::Value,
732) -> std::result::Result<Vec<CustomHeaderMapping>, String> {
733 fn annotation_count(value: &serde_json::Value) -> usize {
734 match value {
735 serde_json::Value::Object(object) => {
736 usize::from(object.contains_key("x-mcp-header"))
737 + object.values().map(annotation_count).sum::<usize>()
738 }
739 serde_json::Value::Array(values) => values.iter().map(annotation_count).sum::<usize>(),
740 _ => 0,
741 }
742 }
743
744 fn is_tchar(byte: u8) -> bool {
745 byte.is_ascii_alphanumeric()
746 || matches!(
747 byte,
748 b'!' | b'#'
749 | b'$'
750 | b'%'
751 | b'&'
752 | b'\''
753 | b'*'
754 | b'+'
755 | b'-'
756 | b'.'
757 | b'^'
758 | b'_'
759 | b'`'
760 | b'|'
761 | b'~'
762 )
763 }
764
765 fn primitive_header_type(schema: &serde_json::Value) -> bool {
766 match schema.get("type") {
767 Some(serde_json::Value::String(kind)) => {
768 matches!(kind.as_str(), "string" | "number" | "integer" | "boolean")
769 }
770 Some(serde_json::Value::Array(kinds)) => {
771 let mut primitive = false;
772 for kind in kinds {
773 match kind.as_str() {
774 Some("string" | "number" | "integer" | "boolean") if !primitive => {
775 primitive = true;
776 }
777 Some("null") => {}
778 _ => return false,
779 }
780 }
781 primitive
782 }
783 _ => false,
784 }
785 }
786
787 fn walk(
788 schema: &serde_json::Value,
789 path: &mut Vec<String>,
790 seen: &mut std::collections::HashSet<String>,
791 mappings: &mut Vec<CustomHeaderMapping>,
792 ) -> std::result::Result<(), String> {
793 let Some(properties) = schema
794 .get("properties")
795 .and_then(serde_json::Value::as_object)
796 else {
797 return Ok(());
798 };
799 for (property_name, property_schema) in properties {
800 path.push(property_name.clone());
801 if let Some(annotation) = property_schema.get("x-mcp-header") {
802 let suffix = annotation
803 .as_str()
804 .ok_or_else(|| format!("annotation at {} is not a string", path.join(".")))?;
805 if suffix.is_empty() || !suffix.bytes().all(is_tchar) {
806 return Err(format!(
807 "invalid header suffix {suffix:?} at {}",
808 path.join(".")
809 ));
810 }
811 if !primitive_header_type(property_schema) {
812 return Err(format!(
813 "annotation at {} is not on a primitive property",
814 path.join(".")
815 ));
816 }
817 if !seen.insert(suffix.to_ascii_lowercase()) {
818 return Err(format!("duplicate header suffix {suffix:?}"));
819 }
820 mappings.push(CustomHeaderMapping {
821 suffix: suffix.to_string(),
822 property_path: path.clone(),
823 });
824 }
825 walk(property_schema, path, seen, mappings)?;
826 path.pop();
827 }
828 Ok(())
829 }
830
831 let mut mappings = Vec::new();
832 walk(
833 schema,
834 &mut Vec::new(),
835 &mut std::collections::HashSet::new(),
836 &mut mappings,
837 )?;
838 if annotation_count(schema) != mappings.len() {
839 return Err(
840 "x-mcp-header annotation is not statically reachable through properties".to_string(),
841 );
842 }
843 Ok(mappings)
844}
845
846fn value_at_property_path<'a>(
847 root: &'a serde_json::Value,
848 path: &[String],
849) -> Option<&'a serde_json::Value> {
850 path.iter().try_fold(root, |value, key| value.get(key))
851}
852
853fn json_value_to_header_string(value: &serde_json::Value) -> Option<String> {
854 match value {
855 serde_json::Value::String(value) => Some(value.clone()),
856 serde_json::Value::Number(value) => Some(value.to_string()),
857 serde_json::Value::Bool(value) => Some(value.to_string()),
858 serde_json::Value::Null | serde_json::Value::Array(_) | serde_json::Value::Object(_) => {
859 None
860 }
861 }
862}
863
864fn encode_header_value(value: &str) -> String {
865 let unsafe_for_header =
866 value.trim() != value || value.bytes().any(|byte| !(0x20..=0x7e).contains(&byte));
867 if unsafe_for_header {
868 format!(
869 "{BASE64_SENTINEL_PREFIX}{}{BASE64_SENTINEL_SUFFIX}",
870 base64::engine::general_purpose::STANDARD.encode(value)
871 )
872 } else {
873 value.to_string()
874 }
875}
876
877#[cfg(feature = "oauth-client")]
878fn bearer_headers(token: &str) -> std::result::Result<reqwest::header::HeaderMap, String> {
879 let value = reqwest::header::HeaderValue::from_str(&format!("Bearer {token}"))
880 .map_err(|_| "token provider returned an invalid bearer token".to_string())?;
881 let mut headers = reqwest::header::HeaderMap::new();
882 headers.insert(reqwest::header::AUTHORIZATION, value);
883 Ok(headers)
887}
888
889fn is_jsonrpc_error_response(value: &serde_json::Value) -> bool {
890 value.get("error").is_some_and(serde_json::Value::is_object)
891 && value.pointer("/error/code").is_some()
892 && value.pointer("/error/message").is_some()
893}
894
895struct HttpRequestSendError {
896 message: String,
897 connection_failed: bool,
898}
899
900impl HttpRequestSendError {
901 fn request(error: reqwest::Error) -> Self {
902 Self {
903 message: format!("HTTP request failed: {error}"),
904 connection_failed: true,
905 }
906 }
907
908 #[cfg(feature = "oauth-client")]
909 fn oauth(error: OAuthClientError) -> Self {
910 Self {
911 message: error.to_string(),
912 connection_failed: false,
913 }
914 }
915}
916
917async fn send_http_request(
918 request: reqwest::RequestBuilder,
919 resource: &str,
920 operation: &str,
921 #[cfg(feature = "oauth-client")] token_provider: Option<Arc<dyn TokenProvider>>,
922 #[cfg(feature = "oauth-client")] scope_escalation: Option<ScopeEscalationRuntime>,
923 #[cfg(feature = "oauth-client")] initial_scope_revision: usize,
924) -> std::result::Result<reqwest::Response, HttpRequestSendError> {
925 #[cfg(feature = "oauth-client")]
926 let mut request = request;
927 #[cfg(not(feature = "oauth-client"))]
928 let _ = (resource, operation);
929
930 #[cfg(feature = "oauth-client")]
931 let mut observed_revision = initial_scope_revision;
932 #[cfg(feature = "oauth-client")]
933 let mut attempts = 0;
934
935 loop {
936 #[cfg(feature = "oauth-client")]
937 let retry_request = request.try_clone();
938
939 let response = request
940 .send()
941 .await
942 .map_err(HttpRequestSendError::request)?;
943
944 #[cfg(feature = "oauth-client")]
945 {
946 let challenge = if response.status() == reqwest::StatusCode::FORBIDDEN {
947 scope_challenge(response.headers())
948 } else {
949 None
950 };
951 let Some(challenge) = challenge else {
952 return Ok(response);
953 };
954 let (Some(runtime), Some(provider), Some(mut retry_request)) = (
955 scope_escalation.as_ref(),
956 token_provider.as_ref(),
957 retry_request,
958 ) else {
959 return Ok(response);
960 };
961 if attempts >= runtime.max_attempts {
962 return Ok(response);
963 }
964
965 attempts += 1;
966 let decision = runtime
967 .respond_to_challenge(challenge, resource, operation, attempts, observed_revision)
968 .await
969 .map_err(HttpRequestSendError::oauth)?;
970 observed_revision = decision.revision;
971
972 let token = provider
973 .get_token()
974 .await
975 .map_err(HttpRequestSendError::oauth)?;
976 let headers = bearer_headers(&token).map_err(|message| {
977 HttpRequestSendError::oauth(OAuthClientError::ScopeEscalation(message))
978 })?;
979 retry_request = retry_request.headers(headers);
980 request = retry_request;
981 }
982
983 #[cfg(not(feature = "oauth-client"))]
984 return Ok(response);
985 }
986}
987
988#[cfg(feature = "oauth-client")]
989fn scope_challenge(headers: &reqwest::header::HeaderMap) -> Option<OAuthScopeChallenge> {
990 headers
991 .get_all(reqwest::header::WWW_AUTHENTICATE)
992 .iter()
993 .filter_map(|value| value.to_str().ok())
994 .find_map(OAuthScopeChallenge::from_www_authenticate)
995}
996
997fn http_status_error(status: reqwest::StatusCode, headers: &reqwest::header::HeaderMap) -> String {
998 #[cfg(feature = "oauth-client")]
999 if let Some(challenge) = scope_challenge(headers) {
1000 let mut message = format!(
1001 "server returned HTTP {status}: insufficient_scope requires {}",
1002 challenge.required_scopes.join(" ")
1003 );
1004 if let Some(resource_metadata) = challenge.resource_metadata {
1005 message.push_str(&format!(" (resource metadata: {resource_metadata})"));
1006 }
1007 return message;
1008 }
1009
1010 #[cfg(not(feature = "oauth-client"))]
1011 let _ = headers;
1012 format!("server returned HTTP {status}")
1013}
1014
1015fn operation_label(parsed: Option<&serde_json::Value>) -> String {
1016 let Some(method) = parsed
1017 .and_then(|value| value.get("method"))
1018 .and_then(serde_json::Value::as_str)
1019 else {
1020 return "unknown".to_string();
1021 };
1022 let target = match method {
1023 "tools/call" | "prompts/get" => parsed
1024 .and_then(|value| value.pointer("/params/name"))
1025 .and_then(serde_json::Value::as_str),
1026 "resources/read" => parsed
1027 .and_then(|value| value.pointer("/params/uri"))
1028 .and_then(serde_json::Value::as_str),
1029 "tasks/get" | "tasks/update" | "tasks/cancel" => parsed
1030 .and_then(|value| value.pointer("/params/taskId"))
1031 .and_then(serde_json::Value::as_str),
1032 _ => None,
1033 };
1034 match target {
1035 Some(target) => format!("{method}:{target}"),
1036 None => method.to_string(),
1037 }
1038}
1039
1040#[async_trait]
1041impl ClientTransport for HttpClientTransport {
1042 async fn send(&mut self, message: &str) -> Result<()> {
1043 if !self.connected.load(Ordering::Acquire) {
1044 return Err(Error::Transport("Transport closed".to_string()));
1045 }
1046
1047 let parsed_message = serde_json::from_str::<serde_json::Value>(message).ok();
1054 let is_notification = parsed_message
1055 .as_ref()
1056 .map(|v| v.get("id").is_none())
1057 .unwrap_or(false);
1058 let method = parsed_message
1059 .as_ref()
1060 .and_then(|value| value.get("method"))
1061 .and_then(serde_json::Value::as_str);
1062 let operation = operation_label(parsed_message.as_ref());
1063 let outbound_version = parsed_message
1064 .as_ref()
1065 .and_then(|value| {
1066 value.pointer("/params/_meta/io.modelcontextprotocol~1protocolVersion")
1067 })
1068 .and_then(serde_json::Value::as_str)
1069 .map(str::to_string);
1070 let is_modern_request =
1071 outbound_version.as_deref() == Some(crate::protocol::PROTOCOL_VERSION_2026_07_28);
1072 if is_modern_request {
1073 self.protocol_version = outbound_version.clone();
1074 self.session_id = None;
1077 }
1078 let timeout = if is_notification {
1079 self.config
1080 .notification_timeout
1081 .min(self.config.request_timeout)
1082 } else {
1083 self.config.request_timeout
1084 };
1085
1086 let mut request = self
1088 .client
1089 .post(&self.url)
1090 .header("Content-Type", "application/json")
1091 .header("Accept", "application/json, text/event-stream");
1092 if method != Some("subscriptions/listen") {
1096 request = request.timeout(timeout);
1097 }
1098
1099 if !is_modern_request && let Some(ref session_id) = self.session_id {
1100 request = request.header("mcp-session-id", session_id);
1101 }
1102
1103 if let Some(version) = outbound_version.as_ref().or(self.protocol_version.as_ref()) {
1104 request = request.header("mcp-protocol-version", version);
1105 }
1106
1107 if let Some(method) = method {
1108 request = request.header(MCP_METHOD_HEADER, method);
1109 let name = match method {
1110 "tools/call" | "prompts/get" => parsed_message
1111 .as_ref()
1112 .and_then(|value| value.pointer("/params/name"))
1113 .and_then(serde_json::Value::as_str),
1114 "resources/read" => parsed_message
1115 .as_ref()
1116 .and_then(|value| value.pointer("/params/uri"))
1117 .and_then(serde_json::Value::as_str),
1118 "tasks/get" | "tasks/update" | "tasks/cancel" => parsed_message
1119 .as_ref()
1120 .and_then(|value| value.pointer("/params/taskId"))
1121 .and_then(serde_json::Value::as_str),
1122 _ => None,
1123 };
1124 if let Some(name) = name {
1125 request = request.header(MCP_NAME_HEADER, name);
1126 }
1127 }
1128
1129 if let Some(parsed) = parsed_message.as_ref() {
1130 for (name, value) in self.outgoing_custom_headers(parsed) {
1131 request = request.header(name, value);
1132 }
1133 }
1134
1135 for (key, value) in &self.config.headers {
1136 request = request.header(key.as_str(), value.as_str());
1137 }
1138
1139 #[cfg(feature = "oauth-client")]
1140 let initial_scope_revision = match &self.scope_escalation {
1141 Some(runtime) => runtime.state.lock().await.revision,
1142 None => 0,
1143 };
1144
1145 #[cfg(feature = "oauth-client")]
1147 if let Some(ref provider) = self.token_provider {
1148 let token = provider
1149 .get_token()
1150 .await
1151 .map_err(|e| Error::Transport(format!("Token provider error: {}", e)))?;
1152 request = request.headers(bearer_headers(&token).map_err(Error::Transport)?);
1153 }
1154
1155 let request = request.body(message.to_string());
1156
1157 if !is_notification && (self.session_id.is_some() || is_modern_request) {
1168 let tx = self.incoming_tx.clone();
1169 let req_id = parsed_message
1175 .as_ref()
1176 .and_then(|value| value.get("id"))
1177 .cloned();
1178 let request_id = req_id
1179 .clone()
1180 .and_then(|value| serde_json::from_value(value).ok());
1181 let is_subscription = method == Some("subscriptions/listen");
1182 let connected = self.connected.clone();
1183 let last_event_id = self.last_event_id.clone();
1184 let sse_retry_delay = self.sse_retry_delay.clone();
1185 let sse_reconnect_signal = self.sse_reconnect_signal.clone();
1186 let max_sse_event_size = self.config.max_sse_event_size;
1187 let request_resource = self.url.clone();
1188 #[cfg(feature = "oauth-client")]
1189 let token_provider = self.token_provider.clone();
1190 #[cfg(feature = "oauth-client")]
1191 let scope_escalation = self.scope_escalation.clone();
1192 self.request_tasks.retain(|_, task| !task.is_finished());
1193 let task = tokio::spawn(async move {
1194 let response_result = send_http_request(
1195 request,
1196 &request_resource,
1197 &operation,
1198 #[cfg(feature = "oauth-client")]
1199 token_provider,
1200 #[cfg(feature = "oauth-client")]
1201 scope_escalation,
1202 #[cfg(feature = "oauth-client")]
1203 initial_scope_revision,
1204 )
1205 .await;
1206 let response = match response_result {
1207 Ok(r) => r,
1208 Err(e) => {
1209 let connection_failed = e.connection_failed;
1210 tracing::error!(error = %e.message, "Background HTTP request failed");
1211 if let Some(id) = &req_id {
1212 let _ = tx.send(transport_error_frame(id, &e.message)).await;
1213 }
1214 if connection_failed {
1215 connected.store(false, Ordering::Release);
1216 }
1217 return;
1218 }
1219 };
1220
1221 let status = response.status();
1222
1223 if status == reqwest::StatusCode::ACCEPTED {
1225 return;
1226 }
1227
1228 if !status.is_success() {
1229 let status_error = http_status_error(status, response.headers());
1230 let body = response.text().await.unwrap_or_default();
1231
1232 if !body.is_empty()
1239 && let Ok(mut v) = serde_json::from_str::<serde_json::Value>(&body)
1240 && is_jsonrpc_error_response(&v)
1241 {
1242 let is_session_signal =
1243 v.pointer("/error/code").and_then(|c| c.as_i64()) == Some(-32005);
1244 if !is_session_signal
1245 && v.get("id").is_none_or(|id| id.is_null())
1246 && let Some(id) = &req_id
1247 {
1248 v["id"] = id.clone();
1249 }
1250 let _ = tx.send(v.to_string()).await;
1251 return;
1252 }
1253
1254 tracing::error!(status = %status, body = %body, "HTTP error from server");
1255 if let Some(id) = &req_id {
1256 let _ = tx.send(transport_error_frame(id, &status_error)).await;
1257 }
1258 connected.store(false, Ordering::Release);
1259 return;
1260 }
1261
1262 let is_sse = response
1264 .headers()
1265 .get("content-type")
1266 .and_then(|v| v.to_str().ok())
1267 .is_some_and(|ct| ct.contains("text/event-stream"));
1268
1269 if is_sse {
1270 let mut stream = response.bytes_stream();
1272 let mut parser = SseParser::with_limit(max_sse_event_size);
1273 let mut had_retry = false;
1274 let mut had_data = false;
1275 let mut subscription_acknowledged = false;
1276
1277 use futures::StreamExt;
1278 while let Some(result) = stream.next().await {
1279 match result {
1280 Ok(bytes) => {
1281 let text = String::from_utf8_lossy(&bytes);
1282 let events = match parser.feed(&text) {
1283 Ok(events) => events,
1284 Err(e) => {
1285 tracing::error!(error = %e, "POST SSE stream terminated");
1291 connected.store(false, Ordering::Release);
1292 return;
1293 }
1294 };
1295 for event in events {
1296 if let Some(ref id) = event.id {
1297 *last_event_id.write().await = Some(id.clone());
1298 }
1299 if let Some(retry_ms) = event.retry {
1300 *sse_retry_delay.write().await =
1301 Some(Duration::from_millis(retry_ms));
1302 had_retry = true;
1303 }
1304 if !event.data.is_empty() {
1305 had_data = true;
1306 let value =
1307 serde_json::from_str::<serde_json::Value>(&event.data);
1308 let value = match value {
1309 Ok(value) => value,
1310 Err(error) if is_subscription => {
1311 if let Some(id) = &req_id {
1312 let _ = tx
1313 .send(transport_error_frame(
1314 id,
1315 &format!(
1316 "subscription stream returned invalid JSON: {error}"
1317 ),
1318 ))
1319 .await;
1320 }
1321 return;
1322 }
1323 Err(_) => {
1324 let _ = tx.send(event.data).await;
1325 continue;
1326 }
1327 };
1328 let is_terminal =
1329 value.get("id").zip(req_id.as_ref()).is_some_and(
1330 |(actual, expected)| {
1331 json_request_ids_match(actual, expected)
1332 },
1333 ) && (value.get("result").is_some()
1334 || value.get("error").is_some());
1335
1336 if is_subscription {
1337 let violation = if is_terminal {
1338 if value.get("error").is_some() {
1339 None
1340 } else if !subscription_acknowledged {
1341 Some(
1342 "subscriptions/listen completed before acknowledgment",
1343 )
1344 } else if !value
1345 .pointer(
1346 "/result/_meta/io.modelcontextprotocol~1subscriptionId",
1347 )
1348 .zip(req_id.as_ref())
1349 .is_some_and(|(actual, expected)| {
1350 json_request_ids_match(actual, expected)
1351 })
1352 {
1353 Some(
1354 "subscriptions/listen result carried a missing or mismatched subscription ID",
1355 )
1356 } else {
1357 None
1358 }
1359 } else if value.get("method").is_some()
1360 && value.get("id").is_none()
1361 {
1362 let correlated = value
1363 .pointer(
1364 "/params/_meta/io.modelcontextprotocol~1subscriptionId",
1365 )
1366 .zip(req_id.as_ref())
1367 .is_some_and(|(actual, expected)| {
1368 json_request_ids_match(actual, expected)
1369 });
1370 let is_acknowledgment = value
1371 .get("method")
1372 .and_then(serde_json::Value::as_str)
1373 == Some(
1374 notifications::SUBSCRIPTIONS_ACKNOWLEDGED,
1375 );
1376 if !correlated {
1377 Some(
1378 "subscription notification carried a missing or mismatched subscription ID",
1379 )
1380 } else if !subscription_acknowledged
1381 && !is_acknowledgment
1382 {
1383 Some(
1384 "subscription notification arrived before acknowledgment",
1385 )
1386 } else if subscription_acknowledged
1387 && is_acknowledgment
1388 {
1389 Some(
1390 "subscription stream sent a duplicate acknowledgment",
1391 )
1392 } else {
1393 if is_acknowledgment {
1394 subscription_acknowledged = true;
1395 }
1396 None
1397 }
1398 } else {
1399 Some(
1400 "subscription stream returned an unrelated JSON-RPC message",
1401 )
1402 };
1403 if let Some(message) = violation {
1404 if let Some(id) = &req_id {
1405 let _ = tx
1406 .send(transport_error_frame(id, message))
1407 .await;
1408 }
1409 return;
1410 }
1411 }
1412 let _ = tx.send(event.data).await;
1413 if is_terminal {
1414 return;
1418 }
1419 }
1420 }
1421 }
1422 Err(e) => {
1423 tracing::warn!(error = %e, "POST SSE stream error");
1424 break;
1425 }
1426 }
1427 }
1428
1429 if !is_modern_request && had_retry && !had_data {
1434 sse_reconnect_signal.notify_one();
1435 } else {
1436 if let Some(id) = &req_id {
1441 let reason = if had_data {
1442 "server closed the response stream before the final reply"
1443 } else {
1444 "server closed the response stream without a reply"
1445 };
1446 let _ = tx.send(transport_error_frame(id, reason)).await;
1447 }
1448 }
1449 } else {
1450 match response.text().await {
1452 Ok(body) if !body.is_empty() => {
1453 let msgs = extract_json_messages(&body);
1454 if msgs.is_empty() {
1455 if let Some(id) = &req_id {
1458 let _ = tx
1459 .send(transport_error_frame(
1460 id,
1461 "server returned an unparseable response body",
1462 ))
1463 .await;
1464 }
1465 } else {
1466 for msg in msgs {
1467 let _ = tx.send(msg).await;
1468 }
1469 }
1470 }
1471 Ok(_) => {
1472 if let Some(id) = &req_id {
1475 let _ = tx
1476 .send(transport_error_frame(
1477 id,
1478 "server returned an empty response body",
1479 ))
1480 .await;
1481 }
1482 }
1483 Err(e) => {
1484 tracing::error!(error = %e, "Failed to read response body");
1485 if let Some(id) = &req_id {
1486 let _ = tx
1487 .send(transport_error_frame(
1488 id,
1489 &format!("failed to read response body: {e}"),
1490 ))
1491 .await;
1492 }
1493 connected.store(false, Ordering::Release);
1494 }
1495 }
1496 }
1497 });
1498 if let Some(request_id) = request_id {
1499 self.request_tasks.insert(request_id, task);
1500 }
1501 return Ok(());
1502 }
1503
1504 let response = send_http_request(
1508 request,
1509 &self.url,
1510 &operation,
1511 #[cfg(feature = "oauth-client")]
1512 self.token_provider.clone(),
1513 #[cfg(feature = "oauth-client")]
1514 self.scope_escalation.clone(),
1515 #[cfg(feature = "oauth-client")]
1516 initial_scope_revision,
1517 )
1518 .await
1519 .map_err(|e| Error::Transport(e.message))?;
1520
1521 let status = response.status();
1522
1523 let new_session_id = response
1525 .headers()
1526 .get("mcp-session-id")
1527 .and_then(|v| v.to_str().ok())
1528 .map(|s| s.to_string());
1529 let new_protocol_version = response
1530 .headers()
1531 .get("mcp-protocol-version")
1532 .and_then(|v| v.to_str().ok())
1533 .map(|s| s.to_string());
1534
1535 if status == reqwest::StatusCode::ACCEPTED {
1537 if !is_modern_request && let Some(sid) = new_session_id {
1539 self.session_id = Some(sid);
1540 }
1541 if let Some(pv) = new_protocol_version {
1542 self.protocol_version = Some(pv);
1543 }
1544 return Ok(());
1545 }
1546
1547 if !status.is_success() {
1548 #[cfg(feature = "oauth-client")]
1549 let status_error = http_status_error(status, response.headers());
1550 let body = response.text().await.unwrap_or_default();
1551 if is_modern_request
1552 && let Ok(mut error) = serde_json::from_str::<serde_json::Value>(&body)
1553 && is_jsonrpc_error_response(&error)
1554 {
1555 if error.get("id").is_none_or(serde_json::Value::is_null)
1556 && let Some(id) = parsed_message.as_ref().and_then(|value| value.get("id"))
1557 {
1558 error["id"] = id.clone();
1559 }
1560 self.incoming_tx
1561 .send(error.to_string())
1562 .await
1563 .map_err(|_| Error::Transport("Internal channel closed".to_string()))?;
1564 return Ok(());
1565 }
1566 if status == reqwest::StatusCode::NOT_FOUND
1571 && self.config.session_recovery
1572 && self.session_id.is_some()
1573 {
1574 return Err(Error::SessionExpired);
1575 }
1576 if status == reqwest::StatusCode::NOT_FOUND && self.session_id.is_none() {
1577 return Err(Error::Transport(format!(
1578 "HTTP 404 from {}: MCP endpoint not found (check the endpoint path; \
1579 some servers serve MCP at the root, others at /mcp)",
1580 self.url
1581 )));
1582 }
1583 #[cfg(feature = "oauth-client")]
1584 if status == reqwest::StatusCode::FORBIDDEN
1585 && status_error.contains("insufficient_scope")
1586 {
1587 return Err(Error::Transport(if body.is_empty() {
1588 status_error
1589 } else {
1590 format!("{status_error}: {body}")
1591 }));
1592 }
1593 return Err(Error::Transport(format!(
1594 "HTTP {status} from server: {body}"
1595 )));
1596 }
1597
1598 if !is_modern_request && let Some(sid) = new_session_id {
1600 let is_new_session = self.session_id.is_none();
1601 self.session_id = Some(sid);
1602
1603 if is_new_session && self.config.auto_sse {
1604 self.start_sse_stream();
1605 }
1606 }
1607 if let Some(pv) = new_protocol_version {
1608 self.protocol_version = Some(pv);
1609 }
1610
1611 let body = response
1613 .text()
1614 .await
1615 .map_err(|e| Error::Transport(format!("Failed to read response: {}", e)))?;
1616
1617 for msg in extract_json_messages(&body) {
1618 self.incoming_tx
1619 .send(msg)
1620 .await
1621 .map_err(|_| Error::Transport("Internal channel closed".to_string()))?;
1622 }
1623
1624 Ok(())
1625 }
1626
1627 async fn recv(&mut self) -> Result<Option<String>> {
1628 match self.incoming_rx.recv().await {
1629 Some(msg) => Ok(Some(self.normalize_incoming_message(msg))),
1634 None => {
1635 self.connected.store(false, Ordering::Release);
1636 Ok(None)
1637 }
1638 }
1639 }
1640
1641 fn is_connected(&self) -> bool {
1642 self.connected.load(Ordering::Acquire)
1643 }
1644
1645 async fn close(&mut self) -> Result<()> {
1646 self.connected.store(false, Ordering::Release);
1647
1648 for (_, task) in self.request_tasks.drain() {
1649 task.abort();
1650 }
1651
1652 if let Some(task) = self.sse_task.take() {
1654 task.abort();
1655 }
1656
1657 if let Some(ref session_id) = self.session_id {
1659 let mut request = self
1660 .client
1661 .delete(&self.url)
1662 .header("mcp-session-id", session_id)
1663 .timeout(Duration::from_secs(5));
1664
1665 for (key, value) in &self.config.headers {
1666 request = request.header(key.as_str(), value.as_str());
1667 }
1668
1669 #[cfg(feature = "oauth-client")]
1671 if let Some(ref provider) = self.token_provider
1672 && let Ok(token) = provider.get_token().await
1673 && let Ok(headers) = bearer_headers(&token)
1674 {
1675 request = request.headers(headers);
1676 }
1677
1678 let _ = request.send().await;
1679 }
1680
1681 self.session_id = None;
1682 Ok(())
1683 }
1684
1685 async fn reset_session(&mut self) {
1686 tracing::info!("Resetting session for re-initialization");
1687
1688 for (_, task) in self.request_tasks.drain() {
1689 task.abort();
1690 }
1691
1692 if let Some(task) = self.sse_task.take() {
1694 task.abort();
1695 }
1696
1697 self.session_id = None;
1699 self.protocol_version = None;
1700 *self.last_event_id.write().await = None;
1701 *self.sse_retry_delay.write().await = None;
1702
1703 while self.incoming_rx.try_recv().is_ok() {}
1705 }
1706
1707 fn supports_session_recovery(&self) -> bool {
1708 self.config.session_recovery
1709 }
1710
1711 async fn cancel_request(&mut self, request_id: &RequestId) -> Result<()> {
1712 if let Some(task) = self.request_tasks.remove(request_id) {
1713 task.abort();
1717 let _ = task.await;
1718 }
1719 Ok(())
1720 }
1721}
1722
1723struct SseLoopParams {
1729 url: String,
1730 client: reqwest::Client,
1731 session_id: String,
1732 protocol_version: Option<String>,
1733 tx: mpsc::Sender<String>,
1734 last_event_id: Arc<RwLock<Option<String>>>,
1735 sse_retry_delay: Arc<RwLock<Option<Duration>>>,
1736 reconnect_signal: Arc<Notify>,
1737 connected: Arc<AtomicBool>,
1738 config: HttpClientConfig,
1739 #[cfg(feature = "oauth-client")]
1740 token_provider: Option<Arc<dyn TokenProvider>>,
1741}
1742
1743async fn sse_stream_loop(params: SseLoopParams) {
1749 let SseLoopParams {
1750 url,
1751 client,
1752 session_id,
1753 protocol_version,
1754 tx,
1755 last_event_id,
1756 sse_retry_delay,
1757 reconnect_signal,
1758 connected,
1759 config,
1760 #[cfg(feature = "oauth-client")]
1761 token_provider,
1762 } = params;
1763 let mut reconnect_attempts = 0u32;
1764
1765 loop {
1766 if !connected.load(Ordering::Acquire) {
1767 break;
1768 }
1769
1770 let mut request = client
1771 .get(&url)
1772 .header("Accept", "text/event-stream")
1773 .header("mcp-session-id", &session_id);
1774
1775 if let Some(ref version) = protocol_version {
1776 request = request.header("mcp-protocol-version", version);
1777 }
1778
1779 for (key, value) in &config.headers {
1780 request = request.header(key.as_str(), value.as_str());
1781 }
1782
1783 #[cfg(feature = "oauth-client")]
1785 if let Some(ref provider) = token_provider {
1786 match provider.get_token().await {
1787 Ok(token) => match bearer_headers(&token) {
1788 Ok(headers) => request = request.headers(headers),
1789 Err(error) => {
1790 tracing::warn!(%error, "Token provider failed for SSE connection");
1791 break;
1792 }
1793 },
1794 Err(e) => {
1795 tracing::warn!(error = %e, "Token provider failed for SSE connection");
1796 break;
1797 }
1798 }
1799 }
1800
1801 if let Some(ref lei) = *last_event_id.read().await {
1803 request = request.header("Last-Event-ID", lei.clone());
1804 }
1805
1806 let response = match request.send().await {
1807 Ok(r) if r.status().is_success() => {
1808 reconnect_attempts = 0;
1809 r
1810 }
1811 Ok(r) => {
1812 tracing::warn!(status = %r.status(), "SSE connection rejected");
1813 break;
1814 }
1815 Err(e) => {
1816 tracing::warn!(error = %e, "SSE connection failed");
1817 if !config.sse_reconnect || reconnect_attempts >= config.max_sse_reconnect_attempts
1818 {
1819 break;
1820 }
1821 reconnect_attempts += 1;
1822 let delay = sse_retry_delay
1823 .read()
1824 .await
1825 .unwrap_or(config.sse_reconnect_delay);
1826 tokio::time::sleep(delay).await;
1827 continue;
1828 }
1829 };
1830
1831 let mut stream = response.bytes_stream();
1833 let mut parser = SseParser::with_limit(config.max_sse_event_size);
1834
1835 use futures::StreamExt;
1836 loop {
1837 tokio::select! {
1838 chunk = stream.next() => {
1839 match chunk {
1840 Some(Ok(bytes)) => {
1841 let text = String::from_utf8_lossy(&bytes);
1842 let events = match parser.feed(&text) {
1843 Ok(events) => events,
1844 Err(e) => {
1845 tracing::error!(error = %e, "SSE stream terminated");
1850 connected.store(false, Ordering::Release);
1851 return;
1852 }
1853 };
1854 for event in events {
1855 if let Some(ref id) = event.id {
1856 *last_event_id.write().await = Some(id.clone());
1857 }
1858 if let Some(retry_ms) = event.retry {
1859 *sse_retry_delay.write().await = Some(Duration::from_millis(retry_ms));
1860 }
1861 if !event.data.is_empty() && tx.send(event.data).await.is_err() {
1862 return; }
1864 }
1865 }
1866 Some(Err(e)) => {
1867 tracing::warn!(error = %e, "SSE stream error");
1868 break;
1869 }
1870 None => {
1871 tracing::debug!("SSE stream ended");
1872 break;
1873 }
1874 }
1875 }
1876 _ = reconnect_signal.notified() => {
1877 tracing::debug!("SSE reconnect signal received, closing current stream");
1878 break;
1879 }
1880 }
1881 }
1882
1883 if !config.sse_reconnect
1885 || !connected.load(Ordering::Acquire)
1886 || reconnect_attempts >= config.max_sse_reconnect_attempts
1887 {
1888 break;
1889 }
1890 reconnect_attempts += 1;
1891 let delay = sse_retry_delay
1892 .read()
1893 .await
1894 .unwrap_or(config.sse_reconnect_delay);
1895 tracing::info!(
1896 attempt = reconnect_attempts,
1897 max = config.max_sse_reconnect_attempts,
1898 delay_ms = delay.as_millis() as u64,
1899 "Reconnecting SSE stream"
1900 );
1901 tokio::time::sleep(delay).await;
1902 }
1903}
1904
1905fn transport_error_frame(id: &serde_json::Value, message: &str) -> String {
1923 serde_json::json!({
1924 "jsonrpc": "2.0",
1925 "id": id,
1926 "error": { "code": -32000, "message": message },
1927 })
1928 .to_string()
1929}
1930
1931fn json_request_ids_match(left: &serde_json::Value, right: &serde_json::Value) -> bool {
1932 left == right
1933 || match (left, right) {
1934 (serde_json::Value::Number(number), serde_json::Value::String(value))
1935 | (serde_json::Value::String(value), serde_json::Value::Number(number)) => number
1936 .as_i64()
1937 .is_some_and(|number| value.parse::<i64>() == Ok(number)),
1938 _ => false,
1939 }
1940}
1941
1942fn extract_json_messages(body: &str) -> Vec<String> {
1943 let trimmed = body.trim();
1944 if trimmed.is_empty() {
1945 return Vec::new();
1946 }
1947
1948 let looks_like_sse = trimmed.starts_with("event:")
1950 || trimmed.starts_with("data:")
1951 || trimmed.starts_with("id:")
1952 || trimmed.starts_with(':');
1953
1954 if looks_like_sse {
1955 let mut parser = SseParser::new();
1958 let events = parser.feed(body).unwrap_or_default();
1959 events.into_iter().map(|e| e.data).collect()
1960 } else {
1961 vec![trimmed.to_string()]
1962 }
1963}
1964
1965#[derive(Debug)]
1967struct SseEvent {
1968 id: Option<String>,
1970 data: String,
1972 retry: Option<u64>,
1974}
1975
1976struct SseParser {
1985 buffer: String,
1987 current_id: Option<String>,
1989 current_data: Vec<String>,
1990 current_retry: Option<u64>,
1991 data_len: usize,
1993 max_event_size: usize,
1995}
1996
1997impl SseParser {
1998 fn new() -> Self {
2000 Self::with_limit(usize::MAX)
2001 }
2002
2003 fn with_limit(max_event_size: usize) -> Self {
2006 Self {
2007 buffer: String::new(),
2008 current_id: None,
2009 current_data: Vec::new(),
2010 current_retry: None,
2011 data_len: 0,
2012 max_event_size,
2013 }
2014 }
2015
2016 fn feed(&mut self, text: &str) -> Result<Vec<SseEvent>> {
2022 self.buffer.push_str(text);
2023 let mut events = Vec::new();
2024
2025 while let Some(newline_pos) = self.buffer.find('\n') {
2027 let line = self.buffer[..newline_pos]
2028 .trim_end_matches('\r')
2029 .to_string();
2030 self.buffer = self.buffer[newline_pos + 1..].to_string();
2031
2032 if line.is_empty() {
2033 if !self.current_data.is_empty() || self.current_retry.is_some() {
2035 events.push(SseEvent {
2036 id: self.current_id.take(),
2037 data: self.current_data.join("\n"),
2038 retry: self.current_retry.take(),
2039 });
2040 self.current_data.clear();
2041 self.data_len = 0;
2042 }
2043 self.current_id = None;
2044 self.current_retry = None;
2045 } else if let Some(value) = line.strip_prefix("id:") {
2046 let trimmed = value.trim();
2047 if !trimmed.is_empty() {
2048 self.current_id = Some(trimmed.to_string());
2049 }
2050 } else if let Some(value) = line.strip_prefix("data:") {
2051 let data = value.trim().to_string();
2052 self.data_len += data.len();
2053 self.current_data.push(data);
2054 } else if let Some(value) = line.strip_prefix("retry:") {
2055 self.current_retry = value.trim().parse().ok();
2056 }
2057 }
2060
2061 let buffered = self.buffer.len() + self.data_len;
2065 if buffered > self.max_event_size {
2066 return Err(Error::SseEventTooLarge {
2067 size: buffered,
2068 limit: self.max_event_size,
2069 });
2070 }
2071
2072 Ok(events)
2073 }
2074}
2075
2076#[cfg(test)]
2077mod tests {
2078 use super::*;
2079
2080 #[test]
2085 fn test_parse_complete_event() {
2086 let mut parser = SseParser::new();
2087 let events = parser
2088 .feed("id: 1\nevent: message\ndata: {\"hello\":\"world\"}\n\n")
2089 .unwrap();
2090
2091 assert_eq!(events.len(), 1);
2092 assert_eq!(events[0].id, Some("1".to_string()));
2093 assert_eq!(events[0].data, "{\"hello\":\"world\"}");
2094 }
2095
2096 #[test]
2097 fn test_parse_multiple_events() {
2098 let mut parser = SseParser::new();
2099 let events = parser
2100 .feed("id: 1\ndata: first\n\nid: 2\ndata: second\n\nid: 3\ndata: third\n\n")
2101 .unwrap();
2102
2103 assert_eq!(events.len(), 3);
2104 assert_eq!(events[0].data, "first");
2105 assert_eq!(events[1].data, "second");
2106 assert_eq!(events[2].data, "third");
2107 assert_eq!(events[0].id, Some("1".to_string()));
2108 assert_eq!(events[1].id, Some("2".to_string()));
2109 assert_eq!(events[2].id, Some("3".to_string()));
2110 }
2111
2112 #[test]
2113 fn test_parse_partial_chunks() {
2114 let mut parser = SseParser::new();
2115
2116 let events = parser.feed("id: 1\nda").unwrap();
2118 assert!(events.is_empty());
2119
2120 let events = parser.feed("ta: hello\n\n").unwrap();
2122 assert_eq!(events.len(), 1);
2123 assert_eq!(events[0].id, Some("1".to_string()));
2124 assert_eq!(events[0].data, "hello");
2125 }
2126
2127 #[test]
2128 fn test_parse_multiline_data() {
2129 let mut parser = SseParser::new();
2130 let events = parser
2131 .feed("id: 1\ndata: line1\ndata: line2\ndata: line3\n\n")
2132 .unwrap();
2133
2134 assert_eq!(events.len(), 1);
2135 assert_eq!(events[0].data, "line1\nline2\nline3");
2136 }
2137
2138 #[test]
2139 fn test_parse_comment_lines() {
2140 let mut parser = SseParser::new();
2141 let events = parser.feed(": keep-alive\nid: 1\ndata: hello\n\n").unwrap();
2142
2143 assert_eq!(events.len(), 1);
2144 assert_eq!(events[0].data, "hello");
2145 }
2146
2147 #[test]
2148 fn test_parse_event_without_id() {
2149 let mut parser = SseParser::new();
2150 let events = parser.feed("data: no-id-event\n\n").unwrap();
2151
2152 assert_eq!(events.len(), 1);
2153 assert_eq!(events[0].id, None);
2154 assert_eq!(events[0].data, "no-id-event");
2155 }
2156
2157 #[test]
2158 fn test_empty_data_no_event() {
2159 let mut parser = SseParser::new();
2160 let events = parser.feed("id: 1\n\n").unwrap();
2161
2162 assert!(events.is_empty());
2164 }
2165
2166 #[test]
2167 fn test_parse_crlf_line_endings() {
2168 let mut parser = SseParser::new();
2169 let events = parser.feed("id: 1\r\ndata: crlf\r\n\r\n").unwrap();
2170
2171 assert_eq!(events.len(), 1);
2172 assert_eq!(events[0].data, "crlf");
2173 }
2174
2175 #[test]
2176 fn test_parse_json_data() {
2177 let mut parser = SseParser::new();
2178 let json = r#"{"jsonrpc":"2.0","method":"notifications/progress","params":{"token":"t1","progress":50}}"#;
2179 let input = format!("id: 42\nevent: message\ndata: {}\n\n", json);
2180 let events = parser.feed(&input).unwrap();
2181
2182 assert_eq!(events.len(), 1);
2183 assert_eq!(events[0].id, Some("42".to_string()));
2184
2185 let parsed: serde_json::Value = serde_json::from_str(&events[0].data).unwrap();
2187 assert_eq!(parsed["method"], "notifications/progress");
2188 }
2189
2190 #[test]
2191 fn test_event_exceeding_limit_is_rejected() {
2192 let mut parser = SseParser::with_limit(64);
2193
2194 let big = "data: ".to_string() + &"x".repeat(128);
2196 let err = parser.feed(&big).unwrap_err();
2197 match err {
2198 Error::SseEventTooLarge { size, limit } => {
2199 assert!(size > 64, "size {} should exceed limit", size);
2200 assert_eq!(limit, 64);
2201 }
2202 other => panic!("expected SseEventTooLarge, got {:?}", other),
2203 }
2204 }
2205
2206 #[test]
2207 fn test_accumulated_data_lines_count_toward_limit() {
2208 let mut parser = SseParser::with_limit(64);
2209
2210 let mut result = Ok(Vec::new());
2212 for _ in 0..10 {
2213 result = parser.feed("data: 0123456789\n");
2214 if result.is_err() {
2215 break;
2216 }
2217 }
2218 assert!(matches!(result, Err(Error::SseEventTooLarge { .. })));
2219 }
2220
2221 #[test]
2222 fn test_events_within_limit_pass() {
2223 let mut parser = SseParser::with_limit(64);
2224 let events = parser.feed("data: hello\n\ndata: world\n\n").unwrap();
2225 assert_eq!(events.len(), 2);
2226 }
2227
2228 #[test]
2233 fn test_default_config() {
2234 let config = HttpClientConfig::default();
2235 assert!(config.auto_sse);
2236 assert_eq!(config.channel_capacity, 256);
2237 assert_eq!(config.request_timeout, Duration::from_secs(30));
2238 assert!(config.sse_reconnect);
2239 assert_eq!(config.sse_reconnect_delay, Duration::from_secs(1));
2240 assert_eq!(config.max_sse_reconnect_attempts, 5);
2241 assert!(config.headers.is_empty());
2242 }
2243
2244 #[test]
2249 fn test_new_transport() {
2250 let transport = HttpClientTransport::new("http://localhost:3000");
2251 assert_eq!(transport.url, "http://localhost:3000");
2252 assert!(transport.session_id.is_none());
2253 assert!(transport.protocol_version.is_none());
2254 assert!(transport.is_connected());
2255 }
2256
2257 #[test]
2258 fn test_with_config() {
2259 let config = HttpClientConfig {
2260 request_timeout: Duration::from_secs(60),
2261 sse_reconnect: false,
2262 ..Default::default()
2263 };
2264 let transport = HttpClientTransport::with_config("http://example.com", config);
2265 assert_eq!(transport.url, "http://example.com");
2266 assert_eq!(transport.config.request_timeout, Duration::from_secs(60));
2267 assert!(!transport.config.sse_reconnect);
2268 }
2269
2270 #[test]
2271 fn test_with_client() {
2272 let client = reqwest::Client::new();
2273 let transport = HttpClientTransport::with_client("http://example.com", client);
2274 assert_eq!(transport.url, "http://example.com");
2275 assert!(transport.is_connected());
2276 }
2277
2278 #[test]
2283 fn test_bearer_token() {
2284 let transport =
2285 HttpClientTransport::new("http://localhost:3000").bearer_token("sk-test-token");
2286 assert_eq!(
2287 transport.config.headers.get("Authorization").unwrap(),
2288 "Bearer sk-test-token"
2289 );
2290 }
2291
2292 #[test]
2293 fn test_api_key() {
2294 let transport = HttpClientTransport::new("http://localhost:3000").api_key("sk-api-key-123");
2295 assert_eq!(
2296 transport.config.headers.get("Authorization").unwrap(),
2297 "Bearer sk-api-key-123"
2298 );
2299 }
2300
2301 #[test]
2302 fn test_api_key_header() {
2303 let transport =
2304 HttpClientTransport::new("http://localhost:3000").api_key_header("X-API-Key", "my-key");
2305 assert_eq!(transport.config.headers.get("X-API-Key").unwrap(), "my-key");
2306 assert!(!transport.config.headers.contains_key("Authorization"));
2307 }
2308
2309 #[test]
2310 fn test_basic_auth() {
2311 let transport =
2312 HttpClientTransport::new("http://localhost:3000").basic_auth("admin", "secret");
2313 let header = transport.config.headers.get("Authorization").unwrap();
2314 assert!(header.starts_with("Basic "));
2315 use base64::Engine;
2316 let decoded = base64::engine::general_purpose::STANDARD
2317 .decode(header.strip_prefix("Basic ").unwrap())
2318 .unwrap();
2319 assert_eq!(String::from_utf8(decoded).unwrap(), "admin:secret");
2320 }
2321
2322 #[test]
2323 fn test_custom_header() {
2324 let transport = HttpClientTransport::new("http://localhost:3000")
2325 .header("X-Custom", "value1")
2326 .header("X-Another", "value2");
2327 assert_eq!(transport.config.headers.get("X-Custom").unwrap(), "value1");
2328 assert_eq!(transport.config.headers.get("X-Another").unwrap(), "value2");
2329 }
2330
2331 #[test]
2332 fn test_chaining_with_config() {
2333 let config = HttpClientConfig {
2334 request_timeout: Duration::from_secs(60),
2335 ..Default::default()
2336 };
2337 let transport =
2338 HttpClientTransport::with_config("http://localhost:3000", config).bearer_token("tk");
2339 assert_eq!(transport.config.request_timeout, Duration::from_secs(60));
2340 assert_eq!(
2341 transport.config.headers.get("Authorization").unwrap(),
2342 "Bearer tk"
2343 );
2344 }
2345
2346 #[test]
2347 fn test_last_auth_wins() {
2348 let transport = HttpClientTransport::new("http://localhost:3000")
2349 .bearer_token("token1")
2350 .basic_auth("user", "pass");
2351 let header = transport.config.headers.get("Authorization").unwrap();
2352 assert!(header.starts_with("Basic "));
2353 }
2354
2355 #[test]
2356 fn test_config_bearer_token() {
2357 let config = HttpClientConfig::default().bearer_token("tk-123");
2358 assert_eq!(
2359 config.headers.get("Authorization").unwrap(),
2360 "Bearer tk-123"
2361 );
2362 }
2363
2364 #[test]
2365 fn test_config_header() {
2366 let config = HttpClientConfig::default().header("X-Foo", "bar");
2367 assert_eq!(config.headers.get("X-Foo").unwrap(), "bar");
2368 }
2369
2370 #[test]
2371 fn test_config_api_key_header() {
2372 let config = HttpClientConfig::default().api_key_header("X-Key", "secret");
2373 assert_eq!(config.headers.get("X-Key").unwrap(), "secret");
2374 }
2375
2376 #[test]
2377 fn test_config_basic_auth() {
2378 let config = HttpClientConfig::default().basic_auth("user", "pw");
2379 let header = config.headers.get("Authorization").unwrap();
2380 assert!(header.starts_with("Basic "));
2381 }
2382
2383 #[test]
2384 fn sep_2243_encodes_only_unsafe_values() {
2385 assert_eq!(encode_header_value("us west 1"), "us west 1");
2386 assert_eq!(encode_header_value(""), "");
2387 assert_eq!(encode_header_value(" padded "), "=?base64?IHBhZGRlZCA=?=");
2388 assert_eq!(
2389 encode_header_value("Hello, 世界"),
2390 "=?base64?SGVsbG8sIOS4lueVjA==?="
2391 );
2392 }
2393
2394 #[test]
2395 fn oauth_error_body_is_not_misclassified_as_jsonrpc() {
2396 assert!(!is_jsonrpc_error_response(&serde_json::json!({
2397 "error": "insufficient_scope",
2398 "error_description": "Token has insufficient scope"
2399 })));
2400 assert!(is_jsonrpc_error_response(&serde_json::json!({
2401 "jsonrpc": "2.0",
2402 "id": 1,
2403 "error": {
2404 "code": -32022,
2405 "message": "Unsupported protocol version"
2406 }
2407 })));
2408 }
2409
2410 #[test]
2411 fn sep_2243_validates_custom_header_annotations() {
2412 let mappings = custom_header_mappings(&serde_json::json!({
2413 "type": "object",
2414 "properties": {
2415 "region": {"type": "string", "x-mcp-header": "Region"},
2416 "priority": {"type": "integer", "x-mcp-header": "Priority"},
2417 "ratio": {"type": "number", "x-mcp-header": "Ratio"}
2418 }
2419 }))
2420 .unwrap();
2421 assert_eq!(mappings.len(), 3);
2422
2423 for invalid in [
2424 serde_json::json!({
2425 "type": "object",
2426 "properties": {"value": {"type": "object", "x-mcp-header": "Value"}}
2427 }),
2428 serde_json::json!({
2429 "type": "object",
2430 "properties": {
2431 "a": {"type": "string", "x-mcp-header": "Region"},
2432 "b": {"type": "string", "x-mcp-header": "region"}
2433 }
2434 }),
2435 serde_json::json!({
2436 "type": "object",
2437 "properties": {"value": {"type": "string", "x-mcp-header": "Bad Header"}}
2438 }),
2439 ] {
2440 assert!(custom_header_mappings(&invalid).is_err());
2441 }
2442 }
2443
2444 #[test]
2445 fn sep_2243_filters_invalid_tools_and_caches_valid_mappings() {
2446 let mut transport = HttpClientTransport::new("http://localhost:3000");
2447 transport.protocol_version = Some(crate::protocol::PROTOCOL_VERSION_2026_07_28.to_string());
2448 let normalized = transport.normalize_incoming_message(
2449 serde_json::json!({
2450 "jsonrpc": "2.0",
2451 "id": 1,
2452 "result": {
2453 "tools": [
2454 {
2455 "name": "valid",
2456 "inputSchema": {
2457 "type": "object",
2458 "properties": {
2459 "region": {"type": "string", "x-mcp-header": "Region"}
2460 }
2461 }
2462 },
2463 {
2464 "name": "invalid",
2465 "inputSchema": {
2466 "type": "object",
2467 "properties": {
2468 "value": {"type": "array", "x-mcp-header": "Value"}
2469 }
2470 }
2471 }
2472 ]
2473 }
2474 })
2475 .to_string(),
2476 );
2477 let parsed: serde_json::Value = serde_json::from_str(&normalized).unwrap();
2478 assert_eq!(parsed["result"]["tools"].as_array().unwrap().len(), 1);
2479 assert!(transport.tool_header_mappings.contains_key("valid"));
2480 assert!(!transport.tool_header_mappings.contains_key("invalid"));
2481 }
2482}