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
917#[cfg(not(feature = "oauth-client"))]
918async fn send_http_request(
919 request: reqwest::RequestBuilder,
920 resource: &str,
921 operation: &str,
922) -> std::result::Result<reqwest::Response, HttpRequestSendError> {
923 let _ = (resource, operation);
924 request.send().await.map_err(HttpRequestSendError::request)
925}
926
927#[cfg(feature = "oauth-client")]
928async fn send_http_request(
929 mut request: reqwest::RequestBuilder,
930 resource: &str,
931 operation: &str,
932 token_provider: Option<Arc<dyn TokenProvider>>,
933 scope_escalation: Option<ScopeEscalationRuntime>,
934 initial_scope_revision: usize,
935) -> std::result::Result<reqwest::Response, HttpRequestSendError> {
936 let mut observed_revision = initial_scope_revision;
937 let mut attempts = 0;
938
939 loop {
940 let retry_request = request.try_clone();
941
942 let response = request
943 .send()
944 .await
945 .map_err(HttpRequestSendError::request)?;
946
947 let challenge = if response.status() == reqwest::StatusCode::FORBIDDEN {
948 scope_challenge(response.headers())
949 } else {
950 None
951 };
952 let Some(challenge) = challenge else {
953 return Ok(response);
954 };
955 let (Some(runtime), Some(provider), Some(mut retry_request)) = (
956 scope_escalation.as_ref(),
957 token_provider.as_ref(),
958 retry_request,
959 ) else {
960 return Ok(response);
961 };
962 if attempts >= runtime.max_attempts {
963 return Ok(response);
964 }
965
966 attempts += 1;
967 let decision = runtime
968 .respond_to_challenge(challenge, resource, operation, attempts, observed_revision)
969 .await
970 .map_err(HttpRequestSendError::oauth)?;
971 observed_revision = decision.revision;
972
973 let token = provider
974 .get_token()
975 .await
976 .map_err(HttpRequestSendError::oauth)?;
977 let headers = bearer_headers(&token).map_err(|message| {
978 HttpRequestSendError::oauth(OAuthClientError::ScopeEscalation(message))
979 })?;
980 retry_request = retry_request.headers(headers);
981 request = retry_request;
982 }
983}
984
985#[cfg(feature = "oauth-client")]
986fn scope_challenge(headers: &reqwest::header::HeaderMap) -> Option<OAuthScopeChallenge> {
987 headers
988 .get_all(reqwest::header::WWW_AUTHENTICATE)
989 .iter()
990 .filter_map(|value| value.to_str().ok())
991 .find_map(OAuthScopeChallenge::from_www_authenticate)
992}
993
994fn http_status_error(status: reqwest::StatusCode, headers: &reqwest::header::HeaderMap) -> String {
995 #[cfg(feature = "oauth-client")]
996 if let Some(challenge) = scope_challenge(headers) {
997 let mut message = format!(
998 "server returned HTTP {status}: insufficient_scope requires {}",
999 challenge.required_scopes.join(" ")
1000 );
1001 if let Some(resource_metadata) = challenge.resource_metadata {
1002 message.push_str(&format!(" (resource metadata: {resource_metadata})"));
1003 }
1004 return message;
1005 }
1006
1007 #[cfg(not(feature = "oauth-client"))]
1008 let _ = headers;
1009 format!("server returned HTTP {status}")
1010}
1011
1012fn operation_label(parsed: Option<&serde_json::Value>) -> String {
1013 let Some(method) = parsed
1014 .and_then(|value| value.get("method"))
1015 .and_then(serde_json::Value::as_str)
1016 else {
1017 return "unknown".to_string();
1018 };
1019 let target = match method {
1020 "tools/call" | "prompts/get" => parsed
1021 .and_then(|value| value.pointer("/params/name"))
1022 .and_then(serde_json::Value::as_str),
1023 "resources/read" => parsed
1024 .and_then(|value| value.pointer("/params/uri"))
1025 .and_then(serde_json::Value::as_str),
1026 "tasks/get" | "tasks/update" | "tasks/cancel" => parsed
1027 .and_then(|value| value.pointer("/params/taskId"))
1028 .and_then(serde_json::Value::as_str),
1029 _ => None,
1030 };
1031 match target {
1032 Some(target) => format!("{method}:{target}"),
1033 None => method.to_string(),
1034 }
1035}
1036
1037#[async_trait]
1038impl ClientTransport for HttpClientTransport {
1039 async fn send(&mut self, message: &str) -> Result<()> {
1040 if !self.connected.load(Ordering::Acquire) {
1041 return Err(Error::Transport("Transport closed".to_string()));
1042 }
1043
1044 let parsed_message = serde_json::from_str::<serde_json::Value>(message).ok();
1051 let is_notification = parsed_message
1052 .as_ref()
1053 .map(|v| v.get("id").is_none())
1054 .unwrap_or(false);
1055 let method = parsed_message
1056 .as_ref()
1057 .and_then(|value| value.get("method"))
1058 .and_then(serde_json::Value::as_str);
1059 let operation = operation_label(parsed_message.as_ref());
1060 let outbound_version = parsed_message
1061 .as_ref()
1062 .and_then(|value| {
1063 value.pointer("/params/_meta/io.modelcontextprotocol~1protocolVersion")
1064 })
1065 .and_then(serde_json::Value::as_str)
1066 .map(str::to_string);
1067 let is_modern_request =
1068 outbound_version.as_deref() == Some(crate::protocol::PROTOCOL_VERSION_2026_07_28);
1069 if is_modern_request {
1070 self.protocol_version = outbound_version.clone();
1071 self.session_id = None;
1074 }
1075 let timeout = if is_notification {
1076 self.config
1077 .notification_timeout
1078 .min(self.config.request_timeout)
1079 } else {
1080 self.config.request_timeout
1081 };
1082
1083 let mut request = self
1085 .client
1086 .post(&self.url)
1087 .header("Content-Type", "application/json")
1088 .header("Accept", "application/json, text/event-stream");
1089 if method != Some("subscriptions/listen") {
1093 request = request.timeout(timeout);
1094 }
1095
1096 if !is_modern_request && let Some(ref session_id) = self.session_id {
1097 request = request.header("mcp-session-id", session_id);
1098 }
1099
1100 if let Some(version) = outbound_version.as_ref().or(self.protocol_version.as_ref()) {
1101 request = request.header("mcp-protocol-version", version);
1102 }
1103
1104 if let Some(method) = method {
1105 request = request.header(MCP_METHOD_HEADER, method);
1106 let name = match method {
1107 "tools/call" | "prompts/get" => parsed_message
1108 .as_ref()
1109 .and_then(|value| value.pointer("/params/name"))
1110 .and_then(serde_json::Value::as_str),
1111 "resources/read" => parsed_message
1112 .as_ref()
1113 .and_then(|value| value.pointer("/params/uri"))
1114 .and_then(serde_json::Value::as_str),
1115 "tasks/get" | "tasks/update" | "tasks/cancel" => parsed_message
1116 .as_ref()
1117 .and_then(|value| value.pointer("/params/taskId"))
1118 .and_then(serde_json::Value::as_str),
1119 _ => None,
1120 };
1121 if let Some(name) = name {
1122 request = request.header(MCP_NAME_HEADER, name);
1123 }
1124 }
1125
1126 if let Some(parsed) = parsed_message.as_ref() {
1127 for (name, value) in self.outgoing_custom_headers(parsed) {
1128 request = request.header(name, value);
1129 }
1130 }
1131
1132 for (key, value) in &self.config.headers {
1133 request = request.header(key.as_str(), value.as_str());
1134 }
1135
1136 #[cfg(feature = "oauth-client")]
1137 let initial_scope_revision = match &self.scope_escalation {
1138 Some(runtime) => runtime.state.lock().await.revision,
1139 None => 0,
1140 };
1141
1142 #[cfg(feature = "oauth-client")]
1144 if let Some(ref provider) = self.token_provider {
1145 let token = provider
1146 .get_token()
1147 .await
1148 .map_err(|e| Error::Transport(format!("Token provider error: {}", e)))?;
1149 request = request.headers(bearer_headers(&token).map_err(Error::Transport)?);
1150 }
1151
1152 let request = request.body(message.to_string());
1153
1154 if !is_notification && (self.session_id.is_some() || is_modern_request) {
1165 let tx = self.incoming_tx.clone();
1166 let req_id = parsed_message
1172 .as_ref()
1173 .and_then(|value| value.get("id"))
1174 .cloned();
1175 let request_id = req_id
1176 .clone()
1177 .and_then(|value| serde_json::from_value(value).ok());
1178 let is_subscription = method == Some("subscriptions/listen");
1179 let connected = self.connected.clone();
1180 let last_event_id = self.last_event_id.clone();
1181 let sse_retry_delay = self.sse_retry_delay.clone();
1182 let sse_reconnect_signal = self.sse_reconnect_signal.clone();
1183 let max_sse_event_size = self.config.max_sse_event_size;
1184 let request_resource = self.url.clone();
1185 #[cfg(feature = "oauth-client")]
1186 let token_provider = self.token_provider.clone();
1187 #[cfg(feature = "oauth-client")]
1188 let scope_escalation = self.scope_escalation.clone();
1189 self.request_tasks.retain(|_, task| !task.is_finished());
1190 let task = tokio::spawn(async move {
1191 let response_result = send_http_request(
1192 request,
1193 &request_resource,
1194 &operation,
1195 #[cfg(feature = "oauth-client")]
1196 token_provider,
1197 #[cfg(feature = "oauth-client")]
1198 scope_escalation,
1199 #[cfg(feature = "oauth-client")]
1200 initial_scope_revision,
1201 )
1202 .await;
1203 let response = match response_result {
1204 Ok(r) => r,
1205 Err(e) => {
1206 let connection_failed = e.connection_failed;
1207 tracing::error!(error = %e.message, "Background HTTP request failed");
1208 if let Some(id) = &req_id {
1209 let _ = tx.send(transport_error_frame(id, &e.message)).await;
1210 }
1211 if connection_failed {
1212 connected.store(false, Ordering::Release);
1213 }
1214 return;
1215 }
1216 };
1217
1218 let status = response.status();
1219
1220 if status == reqwest::StatusCode::ACCEPTED {
1222 return;
1223 }
1224
1225 if !status.is_success() {
1226 let status_error = http_status_error(status, response.headers());
1227 let body = response.text().await.unwrap_or_default();
1228
1229 if !body.is_empty()
1236 && let Ok(mut v) = serde_json::from_str::<serde_json::Value>(&body)
1237 && is_jsonrpc_error_response(&v)
1238 {
1239 let is_session_signal =
1240 v.pointer("/error/code").and_then(|c| c.as_i64()) == Some(-32005);
1241 if !is_session_signal
1242 && v.get("id").is_none_or(|id| id.is_null())
1243 && let Some(id) = &req_id
1244 {
1245 v["id"] = id.clone();
1246 }
1247 let _ = tx.send(v.to_string()).await;
1248 return;
1249 }
1250
1251 tracing::error!(status = %status, body = %body, "HTTP error from server");
1252 if let Some(id) = &req_id {
1253 let _ = tx.send(transport_error_frame(id, &status_error)).await;
1254 }
1255 connected.store(false, Ordering::Release);
1256 return;
1257 }
1258
1259 let is_sse = response
1261 .headers()
1262 .get("content-type")
1263 .and_then(|v| v.to_str().ok())
1264 .is_some_and(|ct| ct.contains("text/event-stream"));
1265
1266 if is_sse {
1267 let mut stream = response.bytes_stream();
1269 let mut parser = SseParser::with_limit(max_sse_event_size);
1270 let mut had_retry = false;
1271 let mut had_data = false;
1272 let mut subscription_acknowledged = false;
1273
1274 use futures::StreamExt;
1275 while let Some(result) = stream.next().await {
1276 match result {
1277 Ok(bytes) => {
1278 let text = String::from_utf8_lossy(&bytes);
1279 let events = match parser.feed(&text) {
1280 Ok(events) => events,
1281 Err(e) => {
1282 tracing::error!(error = %e, "POST SSE stream terminated");
1288 connected.store(false, Ordering::Release);
1289 return;
1290 }
1291 };
1292 for event in events {
1293 if let Some(ref id) = event.id {
1294 *last_event_id.write().await = Some(id.clone());
1295 }
1296 if let Some(retry_ms) = event.retry {
1297 *sse_retry_delay.write().await =
1298 Some(Duration::from_millis(retry_ms));
1299 had_retry = true;
1300 }
1301 if !event.data.is_empty() {
1302 had_data = true;
1303 let value =
1304 serde_json::from_str::<serde_json::Value>(&event.data);
1305 let value = match value {
1306 Ok(value) => value,
1307 Err(error) if is_subscription => {
1308 if let Some(id) = &req_id {
1309 let _ = tx
1310 .send(transport_error_frame(
1311 id,
1312 &format!(
1313 "subscription stream returned invalid JSON: {error}"
1314 ),
1315 ))
1316 .await;
1317 }
1318 return;
1319 }
1320 Err(_) => {
1321 let _ = tx.send(event.data).await;
1322 continue;
1323 }
1324 };
1325 let is_terminal =
1326 value.get("id").zip(req_id.as_ref()).is_some_and(
1327 |(actual, expected)| {
1328 json_request_ids_match(actual, expected)
1329 },
1330 ) && (value.get("result").is_some()
1331 || value.get("error").is_some());
1332
1333 if is_subscription {
1334 let violation = if is_terminal {
1335 if value.get("error").is_some() {
1336 None
1337 } else if !subscription_acknowledged {
1338 Some(
1339 "subscriptions/listen completed before acknowledgment",
1340 )
1341 } else if !value
1342 .pointer(
1343 "/result/_meta/io.modelcontextprotocol~1subscriptionId",
1344 )
1345 .zip(req_id.as_ref())
1346 .is_some_and(|(actual, expected)| {
1347 json_request_ids_match(actual, expected)
1348 })
1349 {
1350 Some(
1351 "subscriptions/listen result carried a missing or mismatched subscription ID",
1352 )
1353 } else {
1354 None
1355 }
1356 } else if value.get("method").is_some()
1357 && value.get("id").is_none()
1358 {
1359 let correlated = value
1360 .pointer(
1361 "/params/_meta/io.modelcontextprotocol~1subscriptionId",
1362 )
1363 .zip(req_id.as_ref())
1364 .is_some_and(|(actual, expected)| {
1365 json_request_ids_match(actual, expected)
1366 });
1367 let is_acknowledgment = value
1368 .get("method")
1369 .and_then(serde_json::Value::as_str)
1370 == Some(
1371 notifications::SUBSCRIPTIONS_ACKNOWLEDGED,
1372 );
1373 if !correlated {
1374 Some(
1375 "subscription notification carried a missing or mismatched subscription ID",
1376 )
1377 } else if !subscription_acknowledged
1378 && !is_acknowledgment
1379 {
1380 Some(
1381 "subscription notification arrived before acknowledgment",
1382 )
1383 } else if subscription_acknowledged
1384 && is_acknowledgment
1385 {
1386 Some(
1387 "subscription stream sent a duplicate acknowledgment",
1388 )
1389 } else {
1390 if is_acknowledgment {
1391 subscription_acknowledged = true;
1392 }
1393 None
1394 }
1395 } else {
1396 Some(
1397 "subscription stream returned an unrelated JSON-RPC message",
1398 )
1399 };
1400 if let Some(message) = violation {
1401 if let Some(id) = &req_id {
1402 let _ = tx
1403 .send(transport_error_frame(id, message))
1404 .await;
1405 }
1406 return;
1407 }
1408 }
1409 let _ = tx.send(event.data).await;
1410 if is_terminal {
1411 return;
1415 }
1416 }
1417 }
1418 }
1419 Err(e) => {
1420 tracing::warn!(error = %e, "POST SSE stream error");
1421 break;
1422 }
1423 }
1424 }
1425
1426 if !is_modern_request && had_retry && !had_data {
1431 sse_reconnect_signal.notify_one();
1432 } else {
1433 if let Some(id) = &req_id {
1438 let reason = if had_data {
1439 "server closed the response stream before the final reply"
1440 } else {
1441 "server closed the response stream without a reply"
1442 };
1443 let _ = tx.send(transport_error_frame(id, reason)).await;
1444 }
1445 }
1446 } else {
1447 match response.text().await {
1449 Ok(body) if !body.is_empty() => {
1450 let msgs = extract_json_messages(&body);
1451 if msgs.is_empty() {
1452 if let Some(id) = &req_id {
1455 let _ = tx
1456 .send(transport_error_frame(
1457 id,
1458 "server returned an unparseable response body",
1459 ))
1460 .await;
1461 }
1462 } else {
1463 for msg in msgs {
1464 let _ = tx.send(msg).await;
1465 }
1466 }
1467 }
1468 Ok(_) => {
1469 if let Some(id) = &req_id {
1472 let _ = tx
1473 .send(transport_error_frame(
1474 id,
1475 "server returned an empty response body",
1476 ))
1477 .await;
1478 }
1479 }
1480 Err(e) => {
1481 tracing::error!(error = %e, "Failed to read response body");
1482 if let Some(id) = &req_id {
1483 let _ = tx
1484 .send(transport_error_frame(
1485 id,
1486 &format!("failed to read response body: {e}"),
1487 ))
1488 .await;
1489 }
1490 connected.store(false, Ordering::Release);
1491 }
1492 }
1493 }
1494 });
1495 if let Some(request_id) = request_id {
1496 self.request_tasks.insert(request_id, task);
1497 }
1498 return Ok(());
1499 }
1500
1501 let response = send_http_request(
1505 request,
1506 &self.url,
1507 &operation,
1508 #[cfg(feature = "oauth-client")]
1509 self.token_provider.clone(),
1510 #[cfg(feature = "oauth-client")]
1511 self.scope_escalation.clone(),
1512 #[cfg(feature = "oauth-client")]
1513 initial_scope_revision,
1514 )
1515 .await
1516 .map_err(|e| Error::Transport(e.message))?;
1517
1518 let status = response.status();
1519
1520 let new_session_id = response
1522 .headers()
1523 .get("mcp-session-id")
1524 .and_then(|v| v.to_str().ok())
1525 .map(|s| s.to_string());
1526 let new_protocol_version = response
1527 .headers()
1528 .get("mcp-protocol-version")
1529 .and_then(|v| v.to_str().ok())
1530 .map(|s| s.to_string());
1531
1532 if status == reqwest::StatusCode::ACCEPTED {
1534 if !is_modern_request && let Some(sid) = new_session_id {
1536 self.session_id = Some(sid);
1537 }
1538 if let Some(pv) = new_protocol_version {
1539 self.protocol_version = Some(pv);
1540 }
1541 return Ok(());
1542 }
1543
1544 if !status.is_success() {
1545 #[cfg(feature = "oauth-client")]
1546 let status_error = http_status_error(status, response.headers());
1547 let body = response.text().await.unwrap_or_default();
1548 if is_modern_request
1549 && let Ok(mut error) = serde_json::from_str::<serde_json::Value>(&body)
1550 && is_jsonrpc_error_response(&error)
1551 {
1552 if error.get("id").is_none_or(serde_json::Value::is_null)
1553 && let Some(id) = parsed_message.as_ref().and_then(|value| value.get("id"))
1554 {
1555 error["id"] = id.clone();
1556 }
1557 self.incoming_tx
1558 .send(error.to_string())
1559 .await
1560 .map_err(|_| Error::Transport("Internal channel closed".to_string()))?;
1561 return Ok(());
1562 }
1563 if status == reqwest::StatusCode::NOT_FOUND
1568 && self.config.session_recovery
1569 && self.session_id.is_some()
1570 {
1571 return Err(Error::SessionExpired);
1572 }
1573 if status == reqwest::StatusCode::NOT_FOUND && self.session_id.is_none() {
1574 return Err(Error::Transport(format!(
1575 "HTTP 404 from {}: MCP endpoint not found (check the endpoint path; \
1576 some servers serve MCP at the root, others at /mcp)",
1577 self.url
1578 )));
1579 }
1580 #[cfg(feature = "oauth-client")]
1581 if status == reqwest::StatusCode::FORBIDDEN
1582 && status_error.contains("insufficient_scope")
1583 {
1584 return Err(Error::Transport(if body.is_empty() {
1585 status_error
1586 } else {
1587 format!("{status_error}: {body}")
1588 }));
1589 }
1590 return Err(Error::Transport(format!(
1591 "HTTP {status} from server: {body}"
1592 )));
1593 }
1594
1595 if !is_modern_request && let Some(sid) = new_session_id {
1597 let is_new_session = self.session_id.is_none();
1598 self.session_id = Some(sid);
1599
1600 if is_new_session && self.config.auto_sse {
1601 self.start_sse_stream();
1602 }
1603 }
1604 if let Some(pv) = new_protocol_version {
1605 self.protocol_version = Some(pv);
1606 }
1607
1608 let body = response
1610 .text()
1611 .await
1612 .map_err(|e| Error::Transport(format!("Failed to read response: {}", e)))?;
1613
1614 for msg in extract_json_messages(&body) {
1615 self.incoming_tx
1616 .send(msg)
1617 .await
1618 .map_err(|_| Error::Transport("Internal channel closed".to_string()))?;
1619 }
1620
1621 Ok(())
1622 }
1623
1624 async fn recv(&mut self) -> Result<Option<String>> {
1625 match self.incoming_rx.recv().await {
1626 Some(msg) => Ok(Some(self.normalize_incoming_message(msg))),
1631 None => {
1632 self.connected.store(false, Ordering::Release);
1633 Ok(None)
1634 }
1635 }
1636 }
1637
1638 fn is_connected(&self) -> bool {
1639 self.connected.load(Ordering::Acquire)
1640 }
1641
1642 async fn close(&mut self) -> Result<()> {
1643 self.connected.store(false, Ordering::Release);
1644
1645 for (_, task) in self.request_tasks.drain() {
1646 task.abort();
1647 }
1648
1649 if let Some(task) = self.sse_task.take() {
1651 task.abort();
1652 }
1653
1654 if let Some(ref session_id) = self.session_id {
1656 let mut request = self
1657 .client
1658 .delete(&self.url)
1659 .header("mcp-session-id", session_id)
1660 .timeout(Duration::from_secs(5));
1661
1662 for (key, value) in &self.config.headers {
1663 request = request.header(key.as_str(), value.as_str());
1664 }
1665
1666 #[cfg(feature = "oauth-client")]
1668 if let Some(ref provider) = self.token_provider
1669 && let Ok(token) = provider.get_token().await
1670 && let Ok(headers) = bearer_headers(&token)
1671 {
1672 request = request.headers(headers);
1673 }
1674
1675 let _ = request.send().await;
1676 }
1677
1678 self.session_id = None;
1679 Ok(())
1680 }
1681
1682 async fn reset_session(&mut self) {
1683 tracing::info!("Resetting session for re-initialization");
1684
1685 for (_, task) in self.request_tasks.drain() {
1686 task.abort();
1687 }
1688
1689 if let Some(task) = self.sse_task.take() {
1691 task.abort();
1692 }
1693
1694 self.session_id = None;
1696 self.protocol_version = None;
1697 *self.last_event_id.write().await = None;
1698 *self.sse_retry_delay.write().await = None;
1699
1700 while self.incoming_rx.try_recv().is_ok() {}
1702 }
1703
1704 fn supports_session_recovery(&self) -> bool {
1705 self.config.session_recovery
1706 }
1707
1708 async fn cancel_request(&mut self, request_id: &RequestId) -> Result<()> {
1709 if let Some(task) = self.request_tasks.remove(request_id) {
1710 task.abort();
1714 let _ = task.await;
1715 }
1716 Ok(())
1717 }
1718}
1719
1720struct SseLoopParams {
1726 url: String,
1727 client: reqwest::Client,
1728 session_id: String,
1729 protocol_version: Option<String>,
1730 tx: mpsc::Sender<String>,
1731 last_event_id: Arc<RwLock<Option<String>>>,
1732 sse_retry_delay: Arc<RwLock<Option<Duration>>>,
1733 reconnect_signal: Arc<Notify>,
1734 connected: Arc<AtomicBool>,
1735 config: HttpClientConfig,
1736 #[cfg(feature = "oauth-client")]
1737 token_provider: Option<Arc<dyn TokenProvider>>,
1738}
1739
1740async fn sse_stream_loop(params: SseLoopParams) {
1746 let SseLoopParams {
1747 url,
1748 client,
1749 session_id,
1750 protocol_version,
1751 tx,
1752 last_event_id,
1753 sse_retry_delay,
1754 reconnect_signal,
1755 connected,
1756 config,
1757 #[cfg(feature = "oauth-client")]
1758 token_provider,
1759 } = params;
1760 let mut reconnect_attempts = 0u32;
1761
1762 loop {
1763 if !connected.load(Ordering::Acquire) {
1764 break;
1765 }
1766
1767 let mut request = client
1768 .get(&url)
1769 .header("Accept", "text/event-stream")
1770 .header("mcp-session-id", &session_id);
1771
1772 if let Some(ref version) = protocol_version {
1773 request = request.header("mcp-protocol-version", version);
1774 }
1775
1776 for (key, value) in &config.headers {
1777 request = request.header(key.as_str(), value.as_str());
1778 }
1779
1780 #[cfg(feature = "oauth-client")]
1782 if let Some(ref provider) = token_provider {
1783 match provider.get_token().await {
1784 Ok(token) => match bearer_headers(&token) {
1785 Ok(headers) => request = request.headers(headers),
1786 Err(error) => {
1787 tracing::warn!(%error, "Token provider failed for SSE connection");
1788 break;
1789 }
1790 },
1791 Err(e) => {
1792 tracing::warn!(error = %e, "Token provider failed for SSE connection");
1793 break;
1794 }
1795 }
1796 }
1797
1798 if let Some(ref lei) = *last_event_id.read().await {
1800 request = request.header("Last-Event-ID", lei.clone());
1801 }
1802
1803 let response = match request.send().await {
1804 Ok(r) if r.status().is_success() => {
1805 reconnect_attempts = 0;
1806 r
1807 }
1808 Ok(r) => {
1809 tracing::warn!(status = %r.status(), "SSE connection rejected");
1810 break;
1811 }
1812 Err(e) => {
1813 tracing::warn!(error = %e, "SSE connection failed");
1814 if !config.sse_reconnect || reconnect_attempts >= config.max_sse_reconnect_attempts
1815 {
1816 break;
1817 }
1818 reconnect_attempts += 1;
1819 let delay = sse_retry_delay
1820 .read()
1821 .await
1822 .unwrap_or(config.sse_reconnect_delay);
1823 tokio::time::sleep(delay).await;
1824 continue;
1825 }
1826 };
1827
1828 let mut stream = response.bytes_stream();
1830 let mut parser = SseParser::with_limit(config.max_sse_event_size);
1831
1832 use futures::StreamExt;
1833 loop {
1834 tokio::select! {
1835 chunk = stream.next() => {
1836 match chunk {
1837 Some(Ok(bytes)) => {
1838 let text = String::from_utf8_lossy(&bytes);
1839 let events = match parser.feed(&text) {
1840 Ok(events) => events,
1841 Err(e) => {
1842 tracing::error!(error = %e, "SSE stream terminated");
1847 connected.store(false, Ordering::Release);
1848 return;
1849 }
1850 };
1851 for event in events {
1852 if let Some(ref id) = event.id {
1853 *last_event_id.write().await = Some(id.clone());
1854 }
1855 if let Some(retry_ms) = event.retry {
1856 *sse_retry_delay.write().await = Some(Duration::from_millis(retry_ms));
1857 }
1858 if !event.data.is_empty() && tx.send(event.data).await.is_err() {
1859 return; }
1861 }
1862 }
1863 Some(Err(e)) => {
1864 tracing::warn!(error = %e, "SSE stream error");
1865 break;
1866 }
1867 None => {
1868 tracing::debug!("SSE stream ended");
1869 break;
1870 }
1871 }
1872 }
1873 _ = reconnect_signal.notified() => {
1874 tracing::debug!("SSE reconnect signal received, closing current stream");
1875 break;
1876 }
1877 }
1878 }
1879
1880 if !config.sse_reconnect
1882 || !connected.load(Ordering::Acquire)
1883 || reconnect_attempts >= config.max_sse_reconnect_attempts
1884 {
1885 break;
1886 }
1887 reconnect_attempts += 1;
1888 let delay = sse_retry_delay
1889 .read()
1890 .await
1891 .unwrap_or(config.sse_reconnect_delay);
1892 tracing::info!(
1893 attempt = reconnect_attempts,
1894 max = config.max_sse_reconnect_attempts,
1895 delay_ms = delay.as_millis() as u64,
1896 "Reconnecting SSE stream"
1897 );
1898 tokio::time::sleep(delay).await;
1899 }
1900}
1901
1902fn transport_error_frame(id: &serde_json::Value, message: &str) -> String {
1920 serde_json::json!({
1921 "jsonrpc": "2.0",
1922 "id": id,
1923 "error": { "code": -32000, "message": message },
1924 })
1925 .to_string()
1926}
1927
1928fn json_request_ids_match(left: &serde_json::Value, right: &serde_json::Value) -> bool {
1929 left == right
1930 || match (left, right) {
1931 (serde_json::Value::Number(number), serde_json::Value::String(value))
1932 | (serde_json::Value::String(value), serde_json::Value::Number(number)) => number
1933 .as_i64()
1934 .is_some_and(|number| value.parse::<i64>() == Ok(number)),
1935 _ => false,
1936 }
1937}
1938
1939fn extract_json_messages(body: &str) -> Vec<String> {
1940 let trimmed = body.trim();
1941 if trimmed.is_empty() {
1942 return Vec::new();
1943 }
1944
1945 let looks_like_sse = trimmed.starts_with("event:")
1947 || trimmed.starts_with("data:")
1948 || trimmed.starts_with("id:")
1949 || trimmed.starts_with(':');
1950
1951 if looks_like_sse {
1952 let mut parser = SseParser::new();
1955 let events = parser.feed(body).unwrap_or_default();
1956 events.into_iter().map(|e| e.data).collect()
1957 } else {
1958 vec![trimmed.to_string()]
1959 }
1960}
1961
1962#[derive(Debug)]
1964struct SseEvent {
1965 id: Option<String>,
1967 data: String,
1969 retry: Option<u64>,
1971}
1972
1973struct SseParser {
1982 buffer: String,
1984 current_id: Option<String>,
1986 current_data: Vec<String>,
1987 current_retry: Option<u64>,
1988 data_len: usize,
1990 max_event_size: usize,
1992}
1993
1994impl SseParser {
1995 fn new() -> Self {
1997 Self::with_limit(usize::MAX)
1998 }
1999
2000 fn with_limit(max_event_size: usize) -> Self {
2003 Self {
2004 buffer: String::new(),
2005 current_id: None,
2006 current_data: Vec::new(),
2007 current_retry: None,
2008 data_len: 0,
2009 max_event_size,
2010 }
2011 }
2012
2013 fn feed(&mut self, text: &str) -> Result<Vec<SseEvent>> {
2019 self.buffer.push_str(text);
2020 let mut events = Vec::new();
2021
2022 while let Some(newline_pos) = self.buffer.find('\n') {
2024 let line = self.buffer[..newline_pos]
2025 .trim_end_matches('\r')
2026 .to_string();
2027 self.buffer = self.buffer[newline_pos + 1..].to_string();
2028
2029 if line.is_empty() {
2030 if !self.current_data.is_empty() || self.current_retry.is_some() {
2032 events.push(SseEvent {
2033 id: self.current_id.take(),
2034 data: self.current_data.join("\n"),
2035 retry: self.current_retry.take(),
2036 });
2037 self.current_data.clear();
2038 self.data_len = 0;
2039 }
2040 self.current_id = None;
2041 self.current_retry = None;
2042 } else if let Some(value) = line.strip_prefix("id:") {
2043 let trimmed = value.trim();
2044 if !trimmed.is_empty() {
2045 self.current_id = Some(trimmed.to_string());
2046 }
2047 } else if let Some(value) = line.strip_prefix("data:") {
2048 let data = value.trim().to_string();
2049 self.data_len += data.len();
2050 self.current_data.push(data);
2051 } else if let Some(value) = line.strip_prefix("retry:") {
2052 self.current_retry = value.trim().parse().ok();
2053 }
2054 }
2057
2058 let buffered = self.buffer.len() + self.data_len;
2062 if buffered > self.max_event_size {
2063 return Err(Error::SseEventTooLarge {
2064 size: buffered,
2065 limit: self.max_event_size,
2066 });
2067 }
2068
2069 Ok(events)
2070 }
2071}
2072
2073#[cfg(test)]
2074mod tests {
2075 use super::*;
2076
2077 #[test]
2082 fn test_parse_complete_event() {
2083 let mut parser = SseParser::new();
2084 let events = parser
2085 .feed("id: 1\nevent: message\ndata: {\"hello\":\"world\"}\n\n")
2086 .unwrap();
2087
2088 assert_eq!(events.len(), 1);
2089 assert_eq!(events[0].id, Some("1".to_string()));
2090 assert_eq!(events[0].data, "{\"hello\":\"world\"}");
2091 }
2092
2093 #[test]
2094 fn test_parse_multiple_events() {
2095 let mut parser = SseParser::new();
2096 let events = parser
2097 .feed("id: 1\ndata: first\n\nid: 2\ndata: second\n\nid: 3\ndata: third\n\n")
2098 .unwrap();
2099
2100 assert_eq!(events.len(), 3);
2101 assert_eq!(events[0].data, "first");
2102 assert_eq!(events[1].data, "second");
2103 assert_eq!(events[2].data, "third");
2104 assert_eq!(events[0].id, Some("1".to_string()));
2105 assert_eq!(events[1].id, Some("2".to_string()));
2106 assert_eq!(events[2].id, Some("3".to_string()));
2107 }
2108
2109 #[test]
2110 fn test_parse_partial_chunks() {
2111 let mut parser = SseParser::new();
2112
2113 let events = parser.feed("id: 1\nda").unwrap();
2115 assert!(events.is_empty());
2116
2117 let events = parser.feed("ta: hello\n\n").unwrap();
2119 assert_eq!(events.len(), 1);
2120 assert_eq!(events[0].id, Some("1".to_string()));
2121 assert_eq!(events[0].data, "hello");
2122 }
2123
2124 #[test]
2125 fn test_parse_multiline_data() {
2126 let mut parser = SseParser::new();
2127 let events = parser
2128 .feed("id: 1\ndata: line1\ndata: line2\ndata: line3\n\n")
2129 .unwrap();
2130
2131 assert_eq!(events.len(), 1);
2132 assert_eq!(events[0].data, "line1\nline2\nline3");
2133 }
2134
2135 #[test]
2136 fn test_parse_comment_lines() {
2137 let mut parser = SseParser::new();
2138 let events = parser.feed(": keep-alive\nid: 1\ndata: hello\n\n").unwrap();
2139
2140 assert_eq!(events.len(), 1);
2141 assert_eq!(events[0].data, "hello");
2142 }
2143
2144 #[test]
2145 fn test_parse_event_without_id() {
2146 let mut parser = SseParser::new();
2147 let events = parser.feed("data: no-id-event\n\n").unwrap();
2148
2149 assert_eq!(events.len(), 1);
2150 assert_eq!(events[0].id, None);
2151 assert_eq!(events[0].data, "no-id-event");
2152 }
2153
2154 #[test]
2155 fn test_empty_data_no_event() {
2156 let mut parser = SseParser::new();
2157 let events = parser.feed("id: 1\n\n").unwrap();
2158
2159 assert!(events.is_empty());
2161 }
2162
2163 #[test]
2164 fn test_parse_crlf_line_endings() {
2165 let mut parser = SseParser::new();
2166 let events = parser.feed("id: 1\r\ndata: crlf\r\n\r\n").unwrap();
2167
2168 assert_eq!(events.len(), 1);
2169 assert_eq!(events[0].data, "crlf");
2170 }
2171
2172 #[test]
2173 fn test_parse_json_data() {
2174 let mut parser = SseParser::new();
2175 let json = r#"{"jsonrpc":"2.0","method":"notifications/progress","params":{"token":"t1","progress":50}}"#;
2176 let input = format!("id: 42\nevent: message\ndata: {}\n\n", json);
2177 let events = parser.feed(&input).unwrap();
2178
2179 assert_eq!(events.len(), 1);
2180 assert_eq!(events[0].id, Some("42".to_string()));
2181
2182 let parsed: serde_json::Value = serde_json::from_str(&events[0].data).unwrap();
2184 assert_eq!(parsed["method"], "notifications/progress");
2185 }
2186
2187 #[test]
2188 fn test_event_exceeding_limit_is_rejected() {
2189 let mut parser = SseParser::with_limit(64);
2190
2191 let big = "data: ".to_string() + &"x".repeat(128);
2193 let err = parser.feed(&big).unwrap_err();
2194 match err {
2195 Error::SseEventTooLarge { size, limit } => {
2196 assert!(size > 64, "size {} should exceed limit", size);
2197 assert_eq!(limit, 64);
2198 }
2199 other => panic!("expected SseEventTooLarge, got {:?}", other),
2200 }
2201 }
2202
2203 #[test]
2204 fn test_accumulated_data_lines_count_toward_limit() {
2205 let mut parser = SseParser::with_limit(64);
2206
2207 let mut result = Ok(Vec::new());
2209 for _ in 0..10 {
2210 result = parser.feed("data: 0123456789\n");
2211 if result.is_err() {
2212 break;
2213 }
2214 }
2215 assert!(matches!(result, Err(Error::SseEventTooLarge { .. })));
2216 }
2217
2218 #[test]
2219 fn test_events_within_limit_pass() {
2220 let mut parser = SseParser::with_limit(64);
2221 let events = parser.feed("data: hello\n\ndata: world\n\n").unwrap();
2222 assert_eq!(events.len(), 2);
2223 }
2224
2225 #[test]
2230 fn test_default_config() {
2231 let config = HttpClientConfig::default();
2232 assert!(config.auto_sse);
2233 assert_eq!(config.channel_capacity, 256);
2234 assert_eq!(config.request_timeout, Duration::from_secs(30));
2235 assert!(config.sse_reconnect);
2236 assert_eq!(config.sse_reconnect_delay, Duration::from_secs(1));
2237 assert_eq!(config.max_sse_reconnect_attempts, 5);
2238 assert!(config.headers.is_empty());
2239 }
2240
2241 #[test]
2246 fn test_new_transport() {
2247 let transport = HttpClientTransport::new("http://localhost:3000");
2248 assert_eq!(transport.url, "http://localhost:3000");
2249 assert!(transport.session_id.is_none());
2250 assert!(transport.protocol_version.is_none());
2251 assert!(transport.is_connected());
2252 }
2253
2254 #[test]
2255 fn test_with_config() {
2256 let config = HttpClientConfig {
2257 request_timeout: Duration::from_secs(60),
2258 sse_reconnect: false,
2259 ..Default::default()
2260 };
2261 let transport = HttpClientTransport::with_config("http://example.com", config);
2262 assert_eq!(transport.url, "http://example.com");
2263 assert_eq!(transport.config.request_timeout, Duration::from_secs(60));
2264 assert!(!transport.config.sse_reconnect);
2265 }
2266
2267 #[test]
2268 fn test_with_client() {
2269 let client = reqwest::Client::new();
2270 let transport = HttpClientTransport::with_client("http://example.com", client);
2271 assert_eq!(transport.url, "http://example.com");
2272 assert!(transport.is_connected());
2273 }
2274
2275 #[test]
2280 fn test_bearer_token() {
2281 let transport =
2282 HttpClientTransport::new("http://localhost:3000").bearer_token("sk-test-token");
2283 assert_eq!(
2284 transport.config.headers.get("Authorization").unwrap(),
2285 "Bearer sk-test-token"
2286 );
2287 }
2288
2289 #[test]
2290 fn test_api_key() {
2291 let transport = HttpClientTransport::new("http://localhost:3000").api_key("sk-api-key-123");
2292 assert_eq!(
2293 transport.config.headers.get("Authorization").unwrap(),
2294 "Bearer sk-api-key-123"
2295 );
2296 }
2297
2298 #[test]
2299 fn test_api_key_header() {
2300 let transport =
2301 HttpClientTransport::new("http://localhost:3000").api_key_header("X-API-Key", "my-key");
2302 assert_eq!(transport.config.headers.get("X-API-Key").unwrap(), "my-key");
2303 assert!(!transport.config.headers.contains_key("Authorization"));
2304 }
2305
2306 #[test]
2307 fn test_basic_auth() {
2308 let transport =
2309 HttpClientTransport::new("http://localhost:3000").basic_auth("admin", "secret");
2310 let header = transport.config.headers.get("Authorization").unwrap();
2311 assert!(header.starts_with("Basic "));
2312 use base64::Engine;
2313 let decoded = base64::engine::general_purpose::STANDARD
2314 .decode(header.strip_prefix("Basic ").unwrap())
2315 .unwrap();
2316 assert_eq!(String::from_utf8(decoded).unwrap(), "admin:secret");
2317 }
2318
2319 #[test]
2320 fn test_custom_header() {
2321 let transport = HttpClientTransport::new("http://localhost:3000")
2322 .header("X-Custom", "value1")
2323 .header("X-Another", "value2");
2324 assert_eq!(transport.config.headers.get("X-Custom").unwrap(), "value1");
2325 assert_eq!(transport.config.headers.get("X-Another").unwrap(), "value2");
2326 }
2327
2328 #[test]
2329 fn test_chaining_with_config() {
2330 let config = HttpClientConfig {
2331 request_timeout: Duration::from_secs(60),
2332 ..Default::default()
2333 };
2334 let transport =
2335 HttpClientTransport::with_config("http://localhost:3000", config).bearer_token("tk");
2336 assert_eq!(transport.config.request_timeout, Duration::from_secs(60));
2337 assert_eq!(
2338 transport.config.headers.get("Authorization").unwrap(),
2339 "Bearer tk"
2340 );
2341 }
2342
2343 #[test]
2344 fn test_last_auth_wins() {
2345 let transport = HttpClientTransport::new("http://localhost:3000")
2346 .bearer_token("token1")
2347 .basic_auth("user", "pass");
2348 let header = transport.config.headers.get("Authorization").unwrap();
2349 assert!(header.starts_with("Basic "));
2350 }
2351
2352 #[test]
2353 fn test_config_bearer_token() {
2354 let config = HttpClientConfig::default().bearer_token("tk-123");
2355 assert_eq!(
2356 config.headers.get("Authorization").unwrap(),
2357 "Bearer tk-123"
2358 );
2359 }
2360
2361 #[test]
2362 fn test_config_header() {
2363 let config = HttpClientConfig::default().header("X-Foo", "bar");
2364 assert_eq!(config.headers.get("X-Foo").unwrap(), "bar");
2365 }
2366
2367 #[test]
2368 fn test_config_api_key_header() {
2369 let config = HttpClientConfig::default().api_key_header("X-Key", "secret");
2370 assert_eq!(config.headers.get("X-Key").unwrap(), "secret");
2371 }
2372
2373 #[test]
2374 fn test_config_basic_auth() {
2375 let config = HttpClientConfig::default().basic_auth("user", "pw");
2376 let header = config.headers.get("Authorization").unwrap();
2377 assert!(header.starts_with("Basic "));
2378 }
2379
2380 #[test]
2381 fn sep_2243_encodes_only_unsafe_values() {
2382 assert_eq!(encode_header_value("us west 1"), "us west 1");
2383 assert_eq!(encode_header_value(""), "");
2384 assert_eq!(encode_header_value(" padded "), "=?base64?IHBhZGRlZCA=?=");
2385 assert_eq!(
2386 encode_header_value("Hello, 世界"),
2387 "=?base64?SGVsbG8sIOS4lueVjA==?="
2388 );
2389 }
2390
2391 #[test]
2392 fn oauth_error_body_is_not_misclassified_as_jsonrpc() {
2393 assert!(!is_jsonrpc_error_response(&serde_json::json!({
2394 "error": "insufficient_scope",
2395 "error_description": "Token has insufficient scope"
2396 })));
2397 assert!(is_jsonrpc_error_response(&serde_json::json!({
2398 "jsonrpc": "2.0",
2399 "id": 1,
2400 "error": {
2401 "code": -32022,
2402 "message": "Unsupported protocol version"
2403 }
2404 })));
2405 }
2406
2407 #[test]
2408 fn sep_2243_validates_custom_header_annotations() {
2409 let mappings = custom_header_mappings(&serde_json::json!({
2410 "type": "object",
2411 "properties": {
2412 "region": {"type": "string", "x-mcp-header": "Region"},
2413 "priority": {"type": "integer", "x-mcp-header": "Priority"},
2414 "ratio": {"type": "number", "x-mcp-header": "Ratio"}
2415 }
2416 }))
2417 .unwrap();
2418 assert_eq!(mappings.len(), 3);
2419
2420 for invalid in [
2421 serde_json::json!({
2422 "type": "object",
2423 "properties": {"value": {"type": "object", "x-mcp-header": "Value"}}
2424 }),
2425 serde_json::json!({
2426 "type": "object",
2427 "properties": {
2428 "a": {"type": "string", "x-mcp-header": "Region"},
2429 "b": {"type": "string", "x-mcp-header": "region"}
2430 }
2431 }),
2432 serde_json::json!({
2433 "type": "object",
2434 "properties": {"value": {"type": "string", "x-mcp-header": "Bad Header"}}
2435 }),
2436 ] {
2437 assert!(custom_header_mappings(&invalid).is_err());
2438 }
2439 }
2440
2441 #[test]
2442 fn sep_2243_filters_invalid_tools_and_caches_valid_mappings() {
2443 let mut transport = HttpClientTransport::new("http://localhost:3000");
2444 transport.protocol_version = Some(crate::protocol::PROTOCOL_VERSION_2026_07_28.to_string());
2445 let normalized = transport.normalize_incoming_message(
2446 serde_json::json!({
2447 "jsonrpc": "2.0",
2448 "id": 1,
2449 "result": {
2450 "tools": [
2451 {
2452 "name": "valid",
2453 "inputSchema": {
2454 "type": "object",
2455 "properties": {
2456 "region": {"type": "string", "x-mcp-header": "Region"}
2457 }
2458 }
2459 },
2460 {
2461 "name": "invalid",
2462 "inputSchema": {
2463 "type": "object",
2464 "properties": {
2465 "value": {"type": "array", "x-mcp-header": "Value"}
2466 }
2467 }
2468 }
2469 ]
2470 }
2471 })
2472 .to_string(),
2473 );
2474 let parsed: serde_json::Value = serde_json::from_str(&normalized).unwrap();
2475 assert_eq!(parsed["result"]["tools"].as_array().unwrap().len(), 1);
2476 assert!(transport.tool_header_mappings.contains_key("valid"));
2477 assert!(!transport.tool_header_mappings.contains_key("invalid"));
2478 }
2479}