1- use std:: collections:: HashMap ;
1+ use std:: { borrow :: Cow , collections:: HashMap , sync :: Arc } ;
22
3- use reqwest:: header:: HeaderMap ;
3+ use futures:: { StreamExt , stream:: BoxStream } ;
4+ use http:: { HeaderName , HeaderValue , header:: WWW_AUTHENTICATE } ;
5+ use reqwest:: header:: { HeaderMap , ACCEPT } ;
6+ use rmcp:: {
7+ model:: { ClientJsonRpcMessage , JsonRpcMessage , ServerJsonRpcMessage } ,
8+ transport:: {
9+ common:: http_header:: {
10+ EVENT_STREAM_MIME_TYPE , HEADER_LAST_EVENT_ID , HEADER_SESSION_ID , JSON_MIME_TYPE ,
11+ } ,
12+ streamable_http_client:: * ,
13+ } ,
14+ } ;
15+ use sse_stream:: { Sse , SseStream } ;
416
5- type ReqwestHttpTransport = rmcp:: transport:: StreamableHttpClientTransport < reqwest:: Client > ;
17+ /// Newtype wrapper around `reqwest::Client` so we can implement the foreign
18+ /// `StreamableHttpClient` trait (orphan rule).
19+ #[ derive( Clone , Debug , Default ) ]
20+ pub struct McpHttpClient ( pub reqwest:: Client ) ;
21+
22+ pub type ReqwestHttpTransport = rmcp:: transport:: StreamableHttpClientTransport < McpHttpClient > ;
623
724/// Builds a `HeaderMap` from a `HashMap<String, String>` of user-provided headers.
825///
@@ -27,3 +44,247 @@ pub fn build_client_with_headers(
2744 ) )
2845 } )
2946}
47+
48+ /// Reserved headers that must not be overridden by custom headers.
49+ /// Matches the validation logic in rmcp's `validate_custom_header`.
50+ const RESERVED_HEADERS : & [ & str ] = & [
51+ "accept" ,
52+ "content-type" ,
53+ "mcp-session-id" ,
54+ "last-event-id" ,
55+ "authorization" ,
56+ "host" ,
57+ "origin" ,
58+ ] ;
59+
60+ /// Applies custom headers to a request builder, rejecting reserved headers.
61+ fn apply_custom_headers (
62+ mut builder : reqwest:: RequestBuilder ,
63+ custom_headers : HashMap < HeaderName , HeaderValue > ,
64+ ) -> Result < reqwest:: RequestBuilder , StreamableHttpError < reqwest:: Error > > {
65+ for ( name, value) in custom_headers {
66+ let name_lower = name. as_str ( ) . to_lowercase ( ) ;
67+ if RESERVED_HEADERS . contains ( & name_lower. as_str ( ) ) {
68+ return Err ( StreamableHttpError :: ReservedHeaderConflict ( name. to_string ( ) ) ) ;
69+ }
70+ builder = builder. header ( name, value) ;
71+ }
72+ Ok ( builder)
73+ }
74+
75+ /// Extracts the scope value from a WWW-Authenticate header.
76+ fn extract_scope ( header : & str ) -> Option < String > {
77+ header. split ( ',' ) . find_map ( |part| {
78+ let part = part. trim ( ) ;
79+ if let Some ( rest) = part. strip_prefix ( "scope=" ) {
80+ Some ( rest. trim_matches ( '"' ) . to_string ( ) )
81+ } else {
82+ None
83+ }
84+ } )
85+ }
86+
87+ /// Attempts to parse `body` as a JSON-RPC error message.
88+ fn parse_json_rpc_error ( body : & str ) -> Option < ServerJsonRpcMessage > {
89+ match serde_json:: from_str :: < ServerJsonRpcMessage > ( body) {
90+ Ok ( message @ JsonRpcMessage :: Error ( _) ) => Some ( message) ,
91+ _ => None ,
92+ }
93+ }
94+
95+ /// Implement `StreamableHttpClient` for our newtype wrapper around reqwest 0.12's `Client`.
96+ ///
97+ /// rmcp 1.6.0 ships its own impl for reqwest 0.13, but warp uses reqwest 0.12.
98+ /// This provides the equivalent implementation against the 0.12 API.
99+ impl StreamableHttpClient for McpHttpClient {
100+ type Error = reqwest:: Error ;
101+
102+ async fn get_stream (
103+ & self ,
104+ uri : Arc < str > ,
105+ session_id : Arc < str > ,
106+ last_event_id : Option < String > ,
107+ auth_token : Option < String > ,
108+ custom_headers : HashMap < HeaderName , HeaderValue > ,
109+ ) -> Result < BoxStream < ' static , Result < Sse , SseError > > , StreamableHttpError < Self :: Error > > {
110+ let mut request_builder = self . 0
111+ . get ( uri. as_ref ( ) )
112+ . header ( ACCEPT , [ EVENT_STREAM_MIME_TYPE , JSON_MIME_TYPE ] . join ( ", " ) )
113+ . header ( HEADER_SESSION_ID , session_id. as_ref ( ) ) ;
114+ if let Some ( last_event_id) = last_event_id {
115+ request_builder = request_builder. header ( HEADER_LAST_EVENT_ID , last_event_id) ;
116+ }
117+ if let Some ( auth_header) = auth_token {
118+ request_builder = request_builder. bearer_auth ( auth_header) ;
119+ }
120+ request_builder = apply_custom_headers ( request_builder, custom_headers) ?;
121+ let response = request_builder
122+ . send ( )
123+ . await
124+ . map_err ( StreamableHttpError :: Client ) ?;
125+ if response. status ( ) == reqwest:: StatusCode :: METHOD_NOT_ALLOWED {
126+ return Err ( StreamableHttpError :: ServerDoesNotSupportSse ) ;
127+ }
128+ let response = response
129+ . error_for_status ( )
130+ . map_err ( StreamableHttpError :: Client ) ?;
131+ match response. headers ( ) . get ( reqwest:: header:: CONTENT_TYPE ) {
132+ Some ( ct) => {
133+ if !ct. as_bytes ( ) . starts_with ( EVENT_STREAM_MIME_TYPE . as_bytes ( ) )
134+ && !ct. as_bytes ( ) . starts_with ( JSON_MIME_TYPE . as_bytes ( ) )
135+ {
136+ return Err ( StreamableHttpError :: UnexpectedContentType ( Some (
137+ String :: from_utf8_lossy ( ct. as_bytes ( ) ) . to_string ( ) ,
138+ ) ) ) ;
139+ }
140+ }
141+ None => {
142+ return Err ( StreamableHttpError :: UnexpectedContentType ( None ) ) ;
143+ }
144+ }
145+ let event_stream = SseStream :: from_byte_stream ( response. bytes_stream ( ) ) . boxed ( ) ;
146+ Ok ( event_stream)
147+ }
148+
149+ async fn delete_session (
150+ & self ,
151+ uri : Arc < str > ,
152+ session : Arc < str > ,
153+ auth_token : Option < String > ,
154+ custom_headers : HashMap < HeaderName , HeaderValue > ,
155+ ) -> Result < ( ) , StreamableHttpError < Self :: Error > > {
156+ let mut request_builder = self . 0 . delete ( uri. as_ref ( ) ) ;
157+ if let Some ( auth_header) = auth_token {
158+ request_builder = request_builder. bearer_auth ( auth_header) ;
159+ }
160+ request_builder = request_builder. header ( HEADER_SESSION_ID , session. as_ref ( ) ) ;
161+ request_builder = apply_custom_headers ( request_builder, custom_headers) ?;
162+ let response = request_builder
163+ . send ( )
164+ . await
165+ . map_err ( StreamableHttpError :: Client ) ?;
166+ if response. status ( ) == reqwest:: StatusCode :: METHOD_NOT_ALLOWED {
167+ tracing:: debug!( "this server doesn't support deleting session" ) ;
168+ return Ok ( ( ) ) ;
169+ }
170+ let _response = response
171+ . error_for_status ( )
172+ . map_err ( StreamableHttpError :: Client ) ?;
173+ Ok ( ( ) )
174+ }
175+
176+ async fn post_message (
177+ & self ,
178+ uri : Arc < str > ,
179+ message : ClientJsonRpcMessage ,
180+ session_id : Option < Arc < str > > ,
181+ auth_token : Option < String > ,
182+ custom_headers : HashMap < HeaderName , HeaderValue > ,
183+ ) -> Result < StreamableHttpPostResponse , StreamableHttpError < Self :: Error > > {
184+ let mut request = self . 0
185+ . post ( uri. as_ref ( ) )
186+ . header ( ACCEPT , [ EVENT_STREAM_MIME_TYPE , JSON_MIME_TYPE ] . join ( ", " ) ) ;
187+ if let Some ( auth_header) = auth_token {
188+ request = request. bearer_auth ( auth_header) ;
189+ }
190+ request = apply_custom_headers ( request, custom_headers) ?;
191+ let session_was_attached = session_id. is_some ( ) ;
192+ if let Some ( session_id) = session_id {
193+ request = request. header ( HEADER_SESSION_ID , session_id. as_ref ( ) ) ;
194+ }
195+ let response = request
196+ . json ( & message)
197+ . send ( )
198+ . await
199+ . map_err ( StreamableHttpError :: Client ) ?;
200+ if response. status ( ) == reqwest:: StatusCode :: UNAUTHORIZED {
201+ if let Some ( header) = response. headers ( ) . get ( WWW_AUTHENTICATE ) {
202+ let header = header
203+ . to_str ( )
204+ . map_err ( |_| {
205+ StreamableHttpError :: UnexpectedServerResponse ( Cow :: from (
206+ "invalid www-authenticate header value" ,
207+ ) )
208+ } ) ?
209+ . to_string ( ) ;
210+ return Err ( StreamableHttpError :: AuthRequired (
211+ AuthRequiredError :: new ( header) ,
212+ ) ) ;
213+ }
214+ }
215+ if response. status ( ) == reqwest:: StatusCode :: FORBIDDEN {
216+ if let Some ( header) = response. headers ( ) . get ( WWW_AUTHENTICATE ) {
217+ let header_str = header. to_str ( ) . map_err ( |_| {
218+ StreamableHttpError :: UnexpectedServerResponse ( Cow :: from (
219+ "invalid www-authenticate header value" ,
220+ ) )
221+ } ) ?;
222+ return Err ( StreamableHttpError :: InsufficientScope (
223+ InsufficientScopeError :: new ( header_str. to_string ( ) , extract_scope ( header_str) ) ,
224+ ) ) ;
225+ }
226+ }
227+ let status = response. status ( ) ;
228+ if matches ! (
229+ status,
230+ reqwest:: StatusCode :: ACCEPTED | reqwest:: StatusCode :: NO_CONTENT
231+ ) {
232+ return Ok ( StreamableHttpPostResponse :: Accepted ) ;
233+ }
234+ if status == reqwest:: StatusCode :: NOT_FOUND && session_was_attached {
235+ return Err ( StreamableHttpError :: SessionExpired ) ;
236+ }
237+ let content_type = response
238+ . headers ( )
239+ . get ( reqwest:: header:: CONTENT_TYPE )
240+ . map ( |ct| String :: from_utf8_lossy ( ct. as_bytes ( ) ) . to_string ( ) ) ;
241+ let session_id = response
242+ . headers ( )
243+ . get ( HEADER_SESSION_ID )
244+ . and_then ( |v| v. to_str ( ) . ok ( ) )
245+ . map ( |s| s. to_string ( ) ) ;
246+ if !status. is_success ( ) {
247+ let body = response
248+ . text ( )
249+ . await
250+ . unwrap_or_else ( |_| "<failed to read response body>" . to_owned ( ) ) ;
251+ if content_type
252+ . as_deref ( )
253+ . is_some_and ( |ct| ct. as_bytes ( ) . starts_with ( JSON_MIME_TYPE . as_bytes ( ) ) )
254+ {
255+ match parse_json_rpc_error ( & body) {
256+ Some ( message) => {
257+ return Ok ( StreamableHttpPostResponse :: Json ( message, session_id) ) ;
258+ }
259+ None => tracing:: warn!(
260+ "HTTP {status}: could not parse JSON body as a JSON-RPC error"
261+ ) ,
262+ }
263+ }
264+ return Err ( StreamableHttpError :: UnexpectedServerResponse ( Cow :: Owned (
265+ format ! ( "HTTP {status}: {body}" ) ,
266+ ) ) ) ;
267+ }
268+ match content_type. as_deref ( ) {
269+ Some ( ct) if ct. as_bytes ( ) . starts_with ( EVENT_STREAM_MIME_TYPE . as_bytes ( ) ) => {
270+ let event_stream = SseStream :: from_byte_stream ( response. bytes_stream ( ) ) . boxed ( ) ;
271+ Ok ( StreamableHttpPostResponse :: Sse ( event_stream, session_id) )
272+ }
273+ Some ( ct) if ct. as_bytes ( ) . starts_with ( JSON_MIME_TYPE . as_bytes ( ) ) => {
274+ match response. json :: < ServerJsonRpcMessage > ( ) . await {
275+ Ok ( message) => Ok ( StreamableHttpPostResponse :: Json ( message, session_id) ) ,
276+ Err ( e) => {
277+ tracing:: warn!(
278+ "could not parse JSON response as ServerJsonRpcMessage, treating as accepted: {e}"
279+ ) ;
280+ Ok ( StreamableHttpPostResponse :: Accepted )
281+ }
282+ }
283+ }
284+ _ => {
285+ tracing:: error!( "unexpected content type: {:?}" , content_type) ;
286+ Err ( StreamableHttpError :: UnexpectedContentType ( content_type) )
287+ }
288+ }
289+ }
290+ }
0 commit comments